mirror of
https://github.com/FluuxIO/go-xmpp.git
synced 2024-11-21 01:52:01 -08:00
0227596f90
Since a transport (and a streamlogger) does not implement io.ByteReader xml.Decoder wraps it using `bufio.NewReader(transport)` so it can easily read bytes one at a time. This has the unfortuante effect of resulting in a panic if we try to parse a stanza that is larger than the default buffer size of 4096 bytes. To fix this we wrap the transport using `bufio.NewReaderSize()` with a much larger buffer size.
134 lines
3.0 KiB
Go
134 lines
3.0 KiB
Go
package xmpp
|
|
|
|
import (
|
|
"bufio"
|
|
"crypto/tls"
|
|
"encoding/xml"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"time"
|
|
|
|
"gosrc.io/xmpp/stanza"
|
|
)
|
|
|
|
// XMPPTransport implements the XMPP native TCP transport
|
|
type XMPPTransport struct {
|
|
openStatement string
|
|
Config TransportConfiguration
|
|
TLSConfig *tls.Config
|
|
decoder *xml.Decoder
|
|
conn net.Conn
|
|
readWriter io.ReadWriter
|
|
logFile io.Writer
|
|
isSecure bool
|
|
}
|
|
|
|
var componentStreamOpen = fmt.Sprintf("<?xml version='1.0'?><stream:stream to='%%s' xmlns='%s' xmlns:stream='%s'>", stanza.NSComponent, stanza.NSStream)
|
|
|
|
var clientStreamOpen = fmt.Sprintf("<?xml version='1.0'?><stream:stream to='%%s' xmlns='%s' xmlns:stream='%s' version='1.0'>", stanza.NSClient, stanza.NSStream)
|
|
|
|
func (t *XMPPTransport) Connect() (string, error) {
|
|
var err error
|
|
|
|
t.conn, err = net.DialTimeout("tcp", t.Config.Address, time.Duration(t.Config.ConnectTimeout)*time.Second)
|
|
if err != nil {
|
|
return "", NewConnError(err, true)
|
|
}
|
|
|
|
t.readWriter = newStreamLogger(t.conn, t.logFile)
|
|
t.decoder = xml.NewDecoder(bufio.NewReaderSize(t.readWriter, maxPacketSize))
|
|
t.decoder.CharsetReader = t.Config.CharsetReader
|
|
return t.StartStream()
|
|
}
|
|
|
|
func (t XMPPTransport) StartStream() (string, error) {
|
|
if _, err := fmt.Fprintf(t, t.openStatement, t.Config.Domain); err != nil {
|
|
t.Close()
|
|
return "", NewConnError(err, true)
|
|
}
|
|
|
|
sessionID, err := stanza.InitStream(t.GetDecoder())
|
|
if err != nil {
|
|
t.Close()
|
|
return "", NewConnError(err, false)
|
|
}
|
|
return sessionID, nil
|
|
}
|
|
|
|
func (t XMPPTransport) DoesStartTLS() bool {
|
|
return true
|
|
}
|
|
|
|
func (t XMPPTransport) GetDomain() string {
|
|
return t.Config.Domain
|
|
}
|
|
|
|
func (t XMPPTransport) GetDecoder() *xml.Decoder {
|
|
return t.decoder
|
|
}
|
|
|
|
func (t XMPPTransport) IsSecure() bool {
|
|
return t.isSecure
|
|
}
|
|
|
|
func (t *XMPPTransport) StartTLS() error {
|
|
if t.Config.TLSConfig == nil {
|
|
t.TLSConfig = &tls.Config{}
|
|
} else {
|
|
t.TLSConfig = t.Config.TLSConfig.Clone()
|
|
}
|
|
|
|
if t.TLSConfig.ServerName == "" {
|
|
t.TLSConfig.ServerName = t.Config.Domain
|
|
}
|
|
tlsConn := tls.Client(t.conn, t.TLSConfig)
|
|
// We convert existing connection to TLS
|
|
if err := tlsConn.Handshake(); err != nil {
|
|
return err
|
|
}
|
|
|
|
t.conn = tlsConn
|
|
t.readWriter = newStreamLogger(tlsConn, t.logFile)
|
|
t.decoder = xml.NewDecoder(bufio.NewReaderSize(t.readWriter, maxPacketSize))
|
|
t.decoder.CharsetReader = t.Config.CharsetReader
|
|
|
|
if !t.TLSConfig.InsecureSkipVerify {
|
|
if err := tlsConn.VerifyHostname(t.Config.Domain); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
t.isSecure = true
|
|
return nil
|
|
}
|
|
|
|
func (t XMPPTransport) Ping() error {
|
|
n, err := t.conn.Write([]byte("\n"))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if n != 1 {
|
|
return errors.New("Could not write ping")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (t XMPPTransport) Read(p []byte) (n int, err error) {
|
|
return t.readWriter.Read(p)
|
|
}
|
|
|
|
func (t XMPPTransport) Write(p []byte) (n int, err error) {
|
|
return t.readWriter.Write(p)
|
|
}
|
|
|
|
func (t XMPPTransport) Close() error {
|
|
_, _ = t.readWriter.Write([]byte("</stream:stream>"))
|
|
return t.conn.Close()
|
|
}
|
|
|
|
func (t *XMPPTransport) LogTraffic(logFile io.Writer) {
|
|
t.logFile = logFile
|
|
}
|