| 1 | local verse = require"verse"; |
| 2 | local base64, unbase64 = require "mime".b64, require"mime".unb64; |
| 3 | local xmlns_sasl = "urn:ietf:params:xml:ns:xmpp-sasl"; |
| 4 | |
| 5 | function verse.plugins.sasl(stream) |
| 6 | local function handle_features(features_stanza) |
| 7 | if stream.authenticated then return; end |
| 8 | stream:debug("Authenticating with SASL..."); |
| 9 | local sasl_mechanisms = features_stanza:get_child("mechanisms", xmlns_sasl); |
| 10 | if not sasl_mechanisms then return end |
| 11 | |
| 12 | local mechanisms = {}; |
| 13 | local preference = {}; |
| 14 | local offered = {}; |
| 15 | |
| 16 | for mech in sasl_mechanisms:childtags("mechanism") do |
| 17 | mech = mech:get_text(); |
| 18 | stream:debug("Server offers %s", mech); |
| 19 | offered[mech] = true; |
| 20 | if not mechanisms[mech] then |
| 21 | local name = mech:match("[^-]+"); |
| 22 | local ok, impl = pcall(require, "verse.util.sasl."..name:lower()); |
| 23 | if ok then |
| 24 | stream:debug("Loaded SASL %s module", name); |
| 25 | mechanisms[mech], preference[mech] = impl(stream, mech); |
| 26 | elseif not tostring(impl):match("not found") then |
| 27 | stream:debug("Loading failed: %s", tostring(impl)); |
| 28 | end |
| 29 | end |
| 30 | end |
| 31 | |
| 32 | local supported = {}; -- by the server |
| 33 | for mech in pairs(mechanisms) do |
| 34 | table.insert(supported, mech); |
| 35 | end |
| 36 | if not supported[1] then |
| 37 | stream:event("authentication-failure", { condition = "no-supported-sasl-mechanisms", mechanisms = offered }); |
| 38 | stream:close(); |
| 39 | return; |
| 40 | end |
| 41 | table.sort(supported, function (a, b) return preference[a] > preference[b]; end); |
| 42 | local mechanism, initial_data = supported[1]; |
| 43 | stream:debug("Selecting %s mechanism...", mechanism); |
| 44 | stream.sasl_mechanism = coroutine.wrap(mechanisms[mechanism]); |
| 45 | initial_data = stream:sasl_mechanism(mechanism); |
| 46 | local auth_stanza = verse.stanza("auth", { xmlns = xmlns_sasl, mechanism = mechanism }); |
| 47 | if initial_data then |
| 48 | auth_stanza:text(base64(initial_data)); |
| 49 | end |
| 50 | stream:send(auth_stanza); |
| 51 | return true; |
| 52 | end |
| 53 | |
| 54 | local function handle_sasl(sasl_stanza) |
| 55 | if sasl_stanza.name == "failure" then |
| 56 | local err = sasl_stanza.tags[1]; |
| 57 | local text = sasl_stanza:get_child_text("text"); |
| 58 | stream:event("authentication-failure", { condition = err.name, text = text }); |
| 59 | stream:close(); |
| 60 | return false; |
| 61 | end |
| 62 | local ok, err = stream.sasl_mechanism(sasl_stanza.name, unbase64(sasl_stanza:get_text())); |
| 63 | if not ok then |
| 64 | stream:event("authentication-failure", { condition = err }); |
| 65 | stream:close(); |
| 66 | return false; |
| 67 | elseif ok == true then |
| 68 | stream:event("authentication-success"); |
| 69 | stream.authenticated = true |
| 70 | stream:reopen(); |
| 71 | else |
| 72 | stream:send(verse.stanza("response", { xmlns = xmlns_sasl }):text(base64(ok))); |
| 73 | end |
| 74 | return true; |
| 75 | end |
| 76 | |
| 77 | stream:hook("stream-features", handle_features, 300); |
| 78 | stream:hook("stream/"..xmlns_sasl, handle_sasl); |
| 79 | |
| 80 | return true; |
| 81 | end |