Files

510 lines
12 KiB
Go

package xmpp
import (
"bytes"
"crypto/rand"
"encoding/hex"
"encoding/xml"
"fmt"
"io"
"log/slog"
"net"
"strings"
"sync"
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
)
const (
nsClient = "jabber:client"
nsStream = "http://etherx.jabber.org/streams"
nsSASL = "urn:ietf:params:xml:ns:xmpp-sasl"
nsBind = "urn:ietf:params:xml:ns:xmpp-bind"
nsSession = "urn:ietf:params:xml:ns:xmpp-session"
nsTLS = "urn:ietf:params:xml:ns:xmpp-tls"
nsIQAuth = "http://jabber.org/features/iq-auth"
nsPing = "urn:xmpp:ping"
)
// handshake states.
const (
hsNeedAuth = iota
hsNeedBind
hsNeedSession
hsNeedHelloWorld
hsReady
)
// Conn handles one plaintext XMPP connection from accept to close.
type Conn struct {
netConn net.Conn
cfg config.Config
registry *session.Registry
bus *session.Bus
clock Clock
diag session.DiagnosticSink
mu sync.Mutex
closed bool
tok *Tokenizer
hsState int
streamID string
domain string
authcid string
resource string
fullJID string
serial string
gen uint64
pingLoop *pingLoop
}
// serve runs the connection until it ends. It owns the read loop and cleanup.
func (c *Conn) serve() {
c.tok = NewTokenizer()
defer c.close()
buf := make([]byte, 4096)
for {
n, err := c.netConn.Read(buf)
if n > 0 {
evs := c.tok.Feed(buf[:n])
for _, ev := range evs {
if !c.handleEvent(ev) {
return
}
}
}
if err != nil {
if err != io.EOF && !isClosed(err) {
slog.Debug("xmpp read error", "jid", c.fullJID, "err", err)
}
c.endSession(session.ReasonTCPClose)
return
}
}
}
func isClosed(err error) bool {
if err == nil {
return false
}
if netErr, ok := err.(net.Error); ok {
return netErr.Timeout() || strings.Contains(err.Error(), "use of closed network connection")
}
return strings.Contains(err.Error(), "use of closed network connection")
}
func (c *Conn) handleEvent(ev Event) bool {
switch ev.Kind {
case StreamOpen:
return c.handleStreamOpen(ev.Attrs)
case Stanza:
if !c.isAuthStanza(ev.Data) {
logStanza("in", ev.Data)
c.emitDiag(session.DirectionIn, "stanza", ev.Data, "")
}
if c.hsState == hsReady {
return c.handlePostReadyStanza(ev.Data)
}
return c.handlePreReadyStanza(ev.Data)
case StreamClose:
c.endSession(session.ReasonStreamClose)
return false
case TokenError:
c.endSession(session.ReasonMalformedXML)
return false
}
return true
}
func (c *Conn) isAuthStanza(b []byte) bool {
space, local, err := rootSpaceLocal(b)
return err == nil && space == nsSASL && local == "auth"
}
func (c *Conn) handleStreamOpen(attrs map[string]string) bool {
domain := attrs["to"]
if domain == "" {
c.endSession(session.ReasonMalformedXML)
return false
}
c.domain = domain
if c.streamID == "" {
c.streamID = newStreamID()
}
var payload string
switch c.hsState {
case hsNeedAuth:
payload = streamOpen(c.streamID, c.domain) +
`<stream:features>` +
`<auth xmlns="http://jabber.org/features/iq-auth"/>` +
`<starttls xmlns="urn:ietf:params:xml:ns:xmpp-tls"><required/></starttls>` +
`<mechanisms xmlns="urn:ietf:params:xml:ns:xmpp-sasl"><mechanism>PLAIN</mechanism></mechanisms>` +
`</stream:features>`
case hsNeedBind:
payload = streamOpen(c.streamID, c.domain) +
`<stream:features>` +
`<bind xmlns="urn:ietf:params:xml:ns:xmpp-bind"/>` +
`<session xmlns="urn:ietf:params:xml:ns:xmpp-session"/>` +
`</stream:features>`
default:
c.endSession(session.ReasonMalformedXML)
return false
}
if err := c.writeRaw([]byte(payload)); err != nil {
c.endSession(session.ReasonTCPClose)
return false
}
return true
}
func streamOpen(id, domain string) string {
return fmt.Sprintf(`<stream:stream xmlns:stream="http://etherx.jabber.org/streams" xmlns="jabber:client" version="1.0" id="%s" from="%s">`, id, domain)
}
func (c *Conn) handlePreReadyStanza(stanza []byte) bool {
space, local, err := rootSpaceLocal(stanza)
if err != nil {
c.endSession(session.ReasonMalformedXML)
return false
}
switch local {
case "starttls":
// STARTTLS is advertised but the robot ignores it; keep reading.
return true
case "auth":
switch space {
case nsIQAuth:
// iq-auth is advertised but not supported. Fail without logging payload.
slog.Info("sasl invalid mechanism")
_ = c.writeStanza([]byte(`<failure xmlns="http://jabber.org/features/iq-auth"><not-authorized/></failure>`))
c.endSession(session.ReasonMalformedXML)
return false
case nsSASL:
return c.handleAuth(stanza)
}
case "iq":
return c.handleIQStanza(stanza)
case "presence":
return c.handleHelloWorld(stanza)
}
// Unknown well-formed pre-READY stanza: log and continue.
logStanza("in", stanza)
return true
}
func (c *Conn) handleAuth(stanza []byte) bool {
var auth saslAuth
if err := xml.Unmarshal(stanza, &auth); err != nil {
slog.Info("sasl malformed")
_ = c.writeStanza(SASLFailureXML("malformed-request"))
c.endSession(session.ReasonMalformedXML)
return false
}
if auth.Mechanism != "PLAIN" {
slog.Info("sasl invalid mechanism")
_ = c.writeStanza(SASLFailureXML("invalid-mechanism"))
c.endSession(session.ReasonMalformedXML)
return false
}
parsed, err := ParseSASLPlain(auth.Chardata)
if err != nil {
slog.Info("sasl malformed")
_ = c.writeStanza(SASLFailureXML("malformed-request"))
c.endSession(session.ReasonMalformedXML)
return false
}
c.authcid = parsed.Authcid
slog.Info("sasl authenticated", "authcid", c.authcid)
_ = c.writeStanza(SASLSuccessXML)
c.tok.ExpectNewStream()
c.hsState = hsNeedBind
return true
}
func (c *Conn) handleIQStanza(stanza []byte) bool {
// Try bind first.
var bindReq iqBind
if err := xml.Unmarshal(stanza, &bindReq); err == nil && bindReq.Bind.XMLName.Local == "bind" {
return c.handleBind(bindReq)
}
// Then session.
var sessReq iqSession
if err := xml.Unmarshal(stanza, &sessReq); err == nil && sessReq.Session.XMLName.Local == "session" {
return c.handleSession(sessReq)
}
logStanza("in", stanza)
return true
}
func (c *Conn) handleBind(req iqBind) bool {
if req.Type != "set" {
logStanza("in", []byte(fmt.Sprintf("bind iq type=%s", req.Type)))
return true
}
c.resource = req.Bind.Resource
if c.resource == "" {
c.resource = "atom"
}
c.fullJID = fmt.Sprintf("%s@%s/%s", c.authcid, c.domain, c.resource)
c.serial = c.authcid
sendFn := func(b []byte) error { return c.writeRaw(b) }
closeFn := func() { c.closeWith(nil) }
gen, replaced := c.registry.Bind(c.fullJID, sendFn, closeFn)
c.gen = gen
if replaced {
c.bus.SessionDown(session.DownEvent{
Generation: gen - 1,
JID: c.fullJID,
Serial: c.serial,
Reason: session.ReasonReplaced,
})
}
result := fmt.Sprintf(`<iq type="result" id="%s"><bind xmlns="urn:ietf:params:xml:ns:xmpp-bind"><jid>%s</jid></bind></iq>`,
escapeXMLAttr(req.ID), c.fullJID)
if err := c.writeStanza([]byte(result)); err != nil {
c.endSession(session.ReasonTCPClose)
return false
}
c.hsState = hsNeedSession
return true
}
func (c *Conn) handleSession(req iqSession) bool {
if req.Type != "set" {
logStanza("in", []byte(fmt.Sprintf("session iq type=%s", req.Type)))
return true
}
result := fmt.Sprintf(`<iq type="result" id="%s"/>`, escapeXMLAttr(req.ID))
if err := c.writeStanza([]byte(result)); err != nil {
c.endSession(session.ReasonTCPClose)
return false
}
c.hsState = hsNeedHelloWorld
return true
}
func (c *Conn) handleHelloWorld(stanza []byte) bool {
var p presenceHello
if err := xml.Unmarshal(stanza, &p); err != nil {
c.endSession(session.ReasonMalformedXML)
return false
}
if strings.TrimSpace(p.Status) != "hello world" {
// Not the READY presence; keep waiting.
logStanza("in", stanza)
return true
}
dummy := fmt.Sprintf(`<presence to="%s"> dummy </presence>`, escapeXMLAttr(c.fullJID))
if err := c.writeStanza([]byte(dummy)); err != nil {
c.endSession(session.ReasonTCPClose)
return false
}
c.hsState = hsReady
c.bus.SessionReady(session.ReadyEvent{
Generation: c.gen,
JID: c.fullJID,
Serial: c.serial,
Send: func(b []byte) error { return c.registry.Send(c.fullJID, c.gen, b) },
})
c.pingLoop = newPingLoop(c)
go c.pingLoop.run()
return true
}
func (c *Conn) handlePostReadyStanza(stanza []byte) bool {
if !c.isCurrent() {
// Generation was replaced; stop processing buffered data.
return false
}
// Bot domain ping.
var ping iqPing
if err := xml.Unmarshal(stanza, &ping); err == nil && ping.Type == "get" && ping.Ping.XMLName.Space == nsPing {
reply := fmt.Sprintf(`<iq type="result" to="%s" from="%s" id="%s"/>`,
escapeXMLAttr(ping.From), escapeXMLAttr(ping.To), escapeXMLAttr(ping.ID))
_ = c.writeStanza([]byte(reply))
return true
}
// Controller ping result.
var result iqResult
if err := xml.Unmarshal(stanza, &result); err == nil && result.Type == "result" && result.ID != "" {
if c.pingLoop != nil && c.pingLoop.result(result.ID) {
return true
}
}
// Any other stanza is forwarded to observers.
c.bus.Stanza(session.StanzaEvent{
Generation: c.gen,
JID: c.fullJID,
Serial: c.serial,
Stanza: append([]byte(nil), stanza...),
})
return true
}
func (c *Conn) isCurrent() bool {
gen, ok := c.registry.Current(c.fullJID)
return ok && gen == c.gen
}
// endSession closes the connection and, if this generation is still current,
// emits SessionDown and removes the registry slot.
func (c *Conn) endSession(reason string) {
c.close()
if c.gen == 0 {
return
}
currentGen, ok := c.registry.Current(c.fullJID)
if !ok || currentGen != c.gen {
return
}
c.bus.SessionDown(session.DownEvent{
Generation: c.gen,
JID: c.fullJID,
Serial: c.serial,
Reason: reason,
})
c.registry.Remove(c.fullJID, c.gen)
}
// closeWith is the registry close callback. It sends </stream:stream> best
// effort and then closes the socket.
func (c *Conn) closeWith(_ error) {
_ = c.writeRaw([]byte("</stream:stream>"))
c.close()
}
func (c *Conn) close() {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return
}
c.closed = true
c.mu.Unlock()
_ = c.netConn.Close()
if c.pingLoop != nil {
c.pingLoop.stop()
}
}
// writeStanza writes a stanza or feature element, logging and emitting a
// diagnostic copy. It never logs SASL material because auth is never sent out.
func (c *Conn) writeStanza(b []byte) error {
logStanza("out", b)
c.emitDiag(session.DirectionOut, "stanza", b, "")
return c.writeRaw(b)
}
// writeRaw performs a mutex-protected write without logging.
func (c *Conn) writeRaw(b []byte) error {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return net.ErrClosed
}
_, err := c.netConn.Write(b)
return err
}
func (c *Conn) emitDiag(dir, kind string, xml []byte, reason string) {
if c.diag == nil {
return
}
c.diag(session.Diagnostic{
Direction: dir,
Kind: kind,
Generation: c.gen,
XML: append([]byte(nil), xml...),
Reason: reason,
})
}
func rootSpaceLocal(b []byte) (string, string, error) {
d := xml.NewDecoder(bytes.NewReader(b))
tok, err := d.Token()
if err != nil {
return "", "", err
}
start, ok := tok.(xml.StartElement)
if !ok {
return "", "", fmt.Errorf("first token is not a start element")
}
return start.Name.Space, start.Name.Local, nil
}
func newStreamID() string {
b := make([]byte, 16)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
func escapeXMLAttr(s string) string {
s = strings.ReplaceAll(s, "&", "&amp;")
s = strings.ReplaceAll(s, "<", "&lt;")
s = strings.ReplaceAll(s, ">", "&gt;")
s = strings.ReplaceAll(s, "\"", "&quot;")
return s
}
type saslAuth struct {
XMLName xml.Name `xml:"urn:ietf:params:xml:ns:xmpp-sasl auth"`
Mechanism string `xml:"mechanism,attr"`
Chardata string `xml:",chardata"`
}
type iqBind struct {
XMLName xml.Name `xml:"iq"`
ID string `xml:"id,attr"`
Type string `xml:"type,attr"`
Bind struct {
XMLName xml.Name `xml:"bind"`
Resource string `xml:"resource"`
} `xml:"bind"`
}
type iqSession struct {
XMLName xml.Name `xml:"iq"`
ID string `xml:"id,attr"`
Type string `xml:"type,attr"`
Session struct {
XMLName xml.Name `xml:"session"`
} `xml:"session"`
}
type presenceHello struct {
XMLName xml.Name `xml:"presence"`
Status string `xml:"status"`
}
type iqPing struct {
XMLName xml.Name `xml:"iq"`
ID string `xml:"id,attr"`
Type string `xml:"type,attr"`
From string `xml:"from,attr"`
To string `xml:"to,attr"`
Ping struct {
XMLName xml.Name `xml:"ping"`
} `xml:"ping"`
}
type iqResult struct {
XMLName xml.Name `xml:"iq"`
ID string `xml:"id,attr"`
Type string `xml:"type,attr"`
From string `xml:"from,attr"`
To string `xml:"to,attr"`
}