Files

342 lines
7.0 KiB
Go

package xmpp
import (
"bytes"
"fmt"
)
// Tokenizer frames XMPP stanzas from a byte stream. It does not use an
// incremental XML decoder because the robot stream is intentionally incomplete:
// <stream:stream> stays open, a stanza can span TCP reads, and several stanzas
// can share one TCP segment.
type Tokenizer struct {
state tokState
buf []byte
pos int // current scan position within buf
depth int
stanzaStart int
err error
}
type tokState int
const (
stateNeedStream tokState = iota
stateInStream
)
// EventKind identifies the kind of token emitted by the tokenizer.
type EventKind int
const (
StreamOpen EventKind = iota
Stanza
StreamClose
TokenError
)
// Event is one token from the XMPP byte stream.
type Event struct {
Kind EventKind
Data []byte
Attrs map[string]string
}
// NewTokenizer returns an empty tokenizer in the need-stream state.
func NewTokenizer() *Tokenizer {
return &Tokenizer{state: stateNeedStream}
}
// ExpectNewStream resets the tokenizer to accept a fresh XML declaration and
// stream open without requiring a closing </stream:stream>. This is used after
// a SASL success.
func (t *Tokenizer) ExpectNewStream() {
t.state = stateNeedStream
t.depth = 0
t.stanzaStart = 0
t.pos = 0
}
// Feed adds bytes and returns any complete events produced. After a TokenError
// event, no further events are emitted until Reset/ExpectNewStream is called.
func (t *Tokenizer) Feed(p []byte) []Event {
if t.err != nil && t.state != stateNeedStream {
return nil
}
t.buf = append(t.buf, p...)
var events []Event
for {
ev, ok := t.advance()
if !ok {
break
}
events = append(events, ev)
if ev.Kind == TokenError {
break
}
}
return events
}
// Err returns the terminal framing error, if any.
func (t *Tokenizer) Err() error { return t.err }
func (t *Tokenizer) advance() (Event, bool) {
switch t.state {
case stateNeedStream:
return t.advanceNeedStream()
case stateInStream:
return t.advanceInStream()
}
return Event{}, false
}
func (t *Tokenizer) advanceNeedStream() (Event, bool) {
for {
t.pos += skipSpace(t.buf[t.pos:])
if t.pos >= len(t.buf) {
return Event{}, false
}
if bytes.HasPrefix(t.buf[t.pos:], []byte("<?xml")) {
end, ok := scanPIEnd(t.buf[t.pos:])
if !ok {
return Event{}, false
}
t.pos += end
continue
}
break
}
if !bytes.HasPrefix(t.buf[t.pos:], []byte("<stream:stream")) {
if len(t.buf[t.pos:]) < len("<stream:stream") {
return Event{}, false
}
t.err = fmt.Errorf("expected stream open, got %q", firstToken(t.buf[t.pos:]))
return Event{Kind: TokenError, Data: copyBytes(t.buf[t.pos:])}, true
}
end, ok := scanTagEnd(t.buf[t.pos:])
if !ok {
return Event{}, false
}
tag := t.buf[t.pos : t.pos+end]
attrs := parseAttrs(tag)
t.buf = t.buf[t.pos+end:]
t.pos = 0
t.state = stateInStream
t.depth = 0
t.stanzaStart = 0
return Event{Kind: StreamOpen, Data: copyBytes(tag), Attrs: attrs}, true
}
func (t *Tokenizer) advanceInStream() (Event, bool) {
for {
if t.depth == 0 {
t.pos += skipSpace(t.buf[t.pos:])
}
if t.pos >= len(t.buf) {
return Event{}, false
}
if t.buf[t.pos] != '<' {
if t.depth == 0 {
t.err = fmt.Errorf("unexpected text at stream level")
t.buf = t.buf[t.pos:]
return Event{Kind: TokenError, Data: copyBytes(t.buf)}, true
}
// Inside a stanza we skip text content until the next tag.
next := bytes.IndexByte(t.buf[t.pos:], '<')
if next < 0 {
return Event{}, false
}
t.pos += next
continue
}
end, ok := scanTagEnd(t.buf[t.pos:])
if !ok {
return Event{}, false
}
tag := t.buf[t.pos : t.pos+end]
nextPos := t.pos + end
if bytes.HasPrefix(tag, []byte("</stream:stream")) {
e := Event{Kind: StreamClose, Data: copyBytes(tag)}
t.buf = t.buf[nextPos:]
t.pos = 0
t.state = stateNeedStream
t.depth = 0
t.stanzaStart = 0
return e, true
}
if bytes.HasPrefix(tag, []byte("</")) {
t.depth--
if t.depth < 0 {
t.err = fmt.Errorf("close tag without matching open")
t.buf = t.buf[nextPos:]
t.pos = 0
return Event{Kind: TokenError, Data: copyBytes(tag)}, true
}
if t.depth == 0 {
data := copyBytes(t.buf[t.stanzaStart:nextPos])
t.buf = t.buf[nextPos:]
t.pos = 0
t.stanzaStart = 0
return Event{Kind: Stanza, Data: data}, true
}
t.pos = nextPos
continue
}
sc := isSelfClosing(tag)
if t.depth == 0 {
if sc {
data := copyBytes(t.buf[t.pos:nextPos])
t.buf = t.buf[nextPos:]
t.pos = 0
t.stanzaStart = 0
return Event{Kind: Stanza, Data: data}, true
}
t.stanzaStart = t.pos
t.depth++
t.pos = nextPos
continue
}
if !sc {
t.depth++
}
t.pos = nextPos
}
}
func skipSpace(b []byte) int {
for i, c := range b {
if c != ' ' && c != '\t' && c != '\n' && c != '\r' {
return i
}
}
return len(b)
}
// scanTagEnd returns the index just past the matching '>' for a tag that starts
// at data[0] == '<'. It is quote-aware and returns ok=false if the tag is not
// yet complete.
func scanTagEnd(data []byte) (int, bool) {
var inQuote byte
for i := 1; i < len(data); i++ {
c := data[i]
if inQuote != 0 {
if c == inQuote {
inQuote = 0
}
continue
}
if c == '"' || c == '\'' {
inQuote = c
continue
}
if c == '>' {
return i + 1, true
}
}
return 0, false
}
// scanPIEnd returns the index just past the matching '?>' for a processing
// instruction that starts with '<?'.
func scanPIEnd(data []byte) (int, bool) {
var inQuote byte
for i := 2; i < len(data); i++ {
c := data[i]
if inQuote != 0 {
if c == inQuote {
inQuote = 0
}
continue
}
if c == '"' || c == '\'' {
inQuote = c
continue
}
if c == '?' && i+1 < len(data) && data[i+1] == '>' {
return i + 2, true
}
}
return 0, false
}
func isSelfClosing(tag []byte) bool {
i := len(tag) - 2 // position before '>'
for i >= 0 && isSpace(tag[i]) {
i--
}
return i >= 0 && tag[i] == '/'
}
func isSpace(c byte) bool {
return c == ' ' || c == '\t' || c == '\n' || c == '\r'
}
func parseAttrs(tag []byte) map[string]string {
attrs := make(map[string]string)
i := 1
// skip element name
for i < len(tag) && !isSpace(tag[i]) {
i++
}
for i < len(tag) {
for i < len(tag) && isSpace(tag[i]) {
i++
}
if i >= len(tag) {
break
}
nameStart := i
for i < len(tag) && tag[i] != '=' && !isSpace(tag[i]) {
i++
}
name := string(tag[nameStart:i])
for i < len(tag) && isSpace(tag[i]) {
i++
}
if i >= len(tag) || tag[i] != '=' {
continue
}
i++ // skip '='
for i < len(tag) && isSpace(tag[i]) {
i++
}
if i >= len(tag) {
break
}
quote := tag[i]
if quote != '"' && quote != '\'' {
continue
}
i++
valStart := i
for i < len(tag) && tag[i] != quote {
i++
}
if i >= len(tag) {
break
}
attrs[name] = string(tag[valStart:i])
i++ // skip closing quote
}
return attrs
}
func copyBytes(b []byte) []byte {
out := make([]byte, len(b))
copy(out, b)
return out
}
func firstToken(b []byte) []byte {
end := bytes.IndexAny(b, " \t\n\r><")
if end < 0 {
return copyBytes(b)
}
return copyBytes(b[:end])
}