telegabber/xmpp/component.go
2026-05-27 08:37:47 -07:00

413 lines
11 KiB
Go

package xmpp
import (
"github.com/pkg/errors"
"regexp"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"dev.narayana.im/narayana/telegabber/badger"
"dev.narayana.im/narayana/telegabber/calls"
"dev.narayana.im/narayana/telegabber/calls/signaling/tgsig"
"dev.narayana.im/narayana/telegabber/calls/signaling/xmppsig"
"dev.narayana.im/narayana/telegabber/config"
"dev.narayana.im/narayana/telegabber/persistence"
"dev.narayana.im/narayana/telegabber/telegram"
"dev.narayana.im/narayana/telegabber/xmpp/gateway"
"dev.narayana.im/narayana/telegabber/xmpp/jingle"
"github.com/pion/webrtc/v4"
log "github.com/sirupsen/logrus"
"gosrc.io/xmpp"
"gosrc.io/xmpp/stanza"
)
var tgConf config.TelegramConfig
var sessions map[string]*telegram.Client
var db *persistence.SessionsYamlDB
// componentEverConnected gates SaveSessions in Close(): a shutdown that
// never connected must not overwrite persisted YAML with an empty map.
var componentEverConnected atomic.Bool
var sessionLock sync.Mutex
// call infrastructure, built in NewComponent
var jingleManager *jingle.Manager
var tgManager *tgsig.Manager
var orchestrator *calls.Orchestrator
var callDeps telegram.CallDeps
// latest XEP-0215 result; nil means no extdisco yet, pion uses host candidates only
var iceServers atomic.Pointer[[]webrtc.ICEServer]
const (
B uint64 = 1
KB = B << 10
MB = KB << 10
GB = MB << 10
TB = GB << 10
PB = TB << 10
EB = PB << 10
maxUint64 uint64 = (1 << 64) - 1
)
var sizeRegex = regexp.MustCompile("\\A([0-9]+) ?([KMGTPE]?B?)\\z")
// NewComponent starts a new component and wraps it in
// a stream manager that you should start yourself
func NewComponent(conf config.XMPPConfig, tc config.TelegramConfig, idsPath string, version string) (*xmpp.StreamManager, *xmpp.Component, error) {
var err error
gateway.Jid, err = stanza.NewJid(conf.Jid)
gateway.Version = version
if err != nil {
return nil, nil, err
}
if gateway.Jid.Resource == "" {
if tc.Tdlib.Client.DeviceModel != "" {
gateway.Jid.Resource = tc.Tdlib.Client.DeviceModel
} else {
gateway.Jid.Resource = "telegabber"
}
}
gateway.IdsDB = badger.IdsDBOpen(idsPath)
tgConf = tc
if tc.Content.Quota != "" {
gateway.StorageQuota, err = parseSize(tc.Content.Quota)
if err != nil {
log.Warnf("Error parsing the storage quota: %v; the cleaner is disabled", err)
}
}
gateway.MAMThreshold = tc.MAMThreshold
options := xmpp.ComponentOptions{
TransportConfiguration: xmpp.TransportConfiguration{
Address: conf.Host + ":" + conf.Port,
Domain: conf.Jid,
},
Domain: conf.Jid,
Secret: conf.Password,
Name: "telegabber",
}
router := xmpp.NewRouter()
router.HandleFunc("iq", HandleIq)
router.HandleFunc("presence", HandlePresence)
router.HandleFunc("message", HandleMessage)
component, err := xmpp.NewComponent(options, router, func(err error) {
log.Error(err)
})
if err != nil {
return nil, nil, err
}
// call routers + orchestrator before loadSessions so per-session
// telegram.Clients get wired with both at construction
jingleManager = &jingle.Manager{
LocalJID: gateway.Jid.Bare(),
Sender: component,
}
tgManager = tgsig.NewManager(nil) // incoming hook set after orchestrator exists
orchestrator = calls.New(calls.Config{
JingleManager: jingleManager,
TgManager: tgManager,
Lookup: sessionLookup,
})
// close the cycle; jingleManager.OnProposal was wired by calls.New
tgManager.SetIncomingHandler(orchestrator.NewFromTelegram)
callDeps = telegram.CallDeps{
TgManager: tgManager,
XmppConfig: xmppsig.AdapterConfig{
Sender: component,
LocalJID: gateway.Jid.Bare(),
Manager: jingleManager,
PCFactory: xmppsig.MakePCFactory(func() []webrtc.ICEServer {
if p := iceServers.Load(); p != nil {
return *p
}
return nil
}),
},
}
// probe all known sessions
err = loadSessions(conf.Db, component)
if err != nil {
return nil, nil, err
}
sm := xmpp.NewStreamManager(component, func(s xmpp.Sender) {
componentEverConnected.Store(true)
go heartbeat(component)
go refreshICEServers(s)
})
return sm, component, nil
}
// fetch XEP-0215 STUN/TURN list from the parent server (the c2s server, i.e.
// the component domain with its leftmost label stripped) and store it for
// PCFactory; runs on every (re)connect because TURN creds are short-lived
func refreshICEServers(s xmpp.Sender) {
parent := parentDomain(gateway.Jid.Domain)
if parent == "" {
log.Warnf("extdisco: cannot derive parent server from component JID %q; skipping", gateway.Jid.Domain)
return
}
services, err := xmppsig.QueryServices(s, gateway.Jid.Bare(), parent, 10*time.Second)
if err != nil {
log.Warnf("extdisco: query %s -> %s failed: %v", gateway.Jid.Bare(), parent, err)
return
}
servers := xmppsig.ToICEServers(services)
iceServers.Store(&servers)
log.Infof("extdisco: loaded %d ICE servers from %s (raw services: %d)", len(servers), parent, len(services))
for i, srv := range servers {
log.Debugf("extdisco[%d]: urls=%v hasUser=%t hasCred=%t", i, srv.URLs, srv.Username != "", srv.Credential != nil && srv.Credential != "")
}
}
// strips the leftmost label, e.g. tlgrm.example.com -> example.com;
// "" if there's no dot
func parentDomain(domain string) string {
if i := strings.IndexByte(domain, '.'); i > 0 && i < len(domain)-1 {
return domain[i+1:]
}
return ""
}
func heartbeat(component *xmpp.Component) {
var err error
probeType := gateway.SPType("probe")
sessionLock.Lock()
for jid := range sessions {
err = gateway.SendPresence(component, jid, probeType)
if err != nil {
log.Error(err)
}
}
sessionLock.Unlock()
quotaLowThreshold := gateway.StorageQuota / 10 * 9
log.Info("Starting heartbeat queue")
// status updater thread
for {
gateway.StorageLock.Lock()
if quotaLowThreshold > 0 && tgConf.Content.Path != "" {
gateway.MeasureStorageSize(tgConf.Content.Path)
if gateway.CachedStorageSize > quotaLowThreshold {
gateway.CleanOldFiles(tgConf.Content.Path, quotaLowThreshold)
}
}
gateway.StorageLock.Unlock()
time.Sleep(60e9)
now := time.Now().Unix()
sessionLock.Lock()
for _, session := range sessions {
session.DelayedStatusesLock.Lock()
for chatID, delayedStatus := range session.DelayedStatuses {
if delayedStatus.TimestampExpired <= now {
go session.ProcessStatusUpdate(
chatID,
session.LastSeenStatus(delayedStatus.TimestampOnline),
"away",
true,
)
delete(session.DelayedStatuses, chatID)
}
}
session.DelayedStatusesLock.Unlock()
// shrink message id maps
session.MessageIdChangesLock.Lock()
for _, idsMap := range session.MessageIdChanges {
for oldMessageId, newId := range idsMap {
if newId.Ts < now - 60 {
newId.Unlock()
delete(idsMap, oldMessageId)
}
}
}
session.MessageIdChangesLock.Unlock()
}
sessionLock.Unlock()
for key, presence := range gateway.Queue {
err = gateway.ResumableSend(component, presence)
if err != nil {
gateway.LogBadPresence(presence)
} else {
gateway.QueueLock.Lock()
delete(gateway.Queue, key)
gateway.QueueLock.Unlock()
}
}
if gateway.DirtySessions {
gateway.DirtySessions = false
// no problem if a dirty flag gets set again here,
// it would be resolved on the next iteration
SaveSessions()
}
gateway.IdsDB.Gc()
}
}
func loadSessions(dbPath string, component *xmpp.Component) error {
var err error
sessions = make(map[string]*telegram.Client)
db, err = persistence.LoadSessions(dbPath)
if err != nil {
return err
}
db.Transaction(func() bool {
for jid, session := range db.Data.Sessions {
// copy the session struct, otherwise all of them would reference
// the same temporary range variable
currentSession := session
getTelegramInstance(jid, &currentSession, component)
}
return false
}, persistence.SessionMarshaller)
return nil
}
func getTelegramInstance(jid string, savedSession *persistence.Session, component *xmpp.Component) (*telegram.Client, bool) {
var err error
session, ok := sessions[jid]
if !ok {
session, err = telegram.NewClient(tgConf, jid, component, savedSession)
if err != nil {
log.Error(errors.Wrap(err, "TDlib initialization failure"))
return session, false
}
session.SetCallDeps(callDeps)
if savedSession.KeepOnline {
if err = session.Connect("", false); err != nil {
log.Error(err)
return session, false
}
}
sessionLock.Lock()
sessions[jid] = session
sessionLock.Unlock()
}
return session, true
}
// orchestrator's SessionLookup callback
func sessionLookup(bareJID string) (*tgsig.Adapter, *xmppsig.Adapter, bool) {
sessionLock.Lock()
cl, ok := sessions[bareJID]
sessionLock.Unlock()
if !ok || cl == nil {
return nil, nil, false
}
return cl.CallAdapter(), cl.XmppCallAdapter(), true
}
// SaveSessions dumps current sessions to the file
func SaveSessions() {
sessionLock.Lock()
defer sessionLock.Unlock()
db.Transaction(func() bool {
for jid, session := range sessions {
db.Data.Sessions[jid] = *session.Session
}
return true
}, persistence.SessionMarshaller)
}
// gracefully terminate the component and save active sessions.
// SaveSessions is skipped when componentEverConnected is false to avoid
// overwriting persisted YAML with the empty-at-startup map (seen during
// a tunnel-down crash-loop)
func Close(component *xmpp.Component) {
log.Error("Disconnecting...")
sessionLock.Lock()
// close all sessions
for _, session := range sessions {
session.Disconnect("", true)
}
sessionLock.Unlock()
if componentEverConnected.Load() {
SaveSessions()
} else {
log.Warn("Close: component never connected, skipping SaveSessions to preserve on-disk state")
}
// flush the ids database
gateway.IdsDB.Close()
// close stream
component.Disconnect()
}
// based on https://github.com/c2h5oh/datasize/blob/master/datasize.go
func parseSize(sSize string) (uint64, error) {
sizeParts := sizeRegex.FindStringSubmatch(sSize)
if len(sizeParts) > 2 {
numPart, err := strconv.ParseInt(sizeParts[1], 10, 64)
if err != nil {
return 0, err
}
var divisor uint64
val := uint64(numPart)
if len(sizeParts[2]) > 0 {
switch sizeParts[2][0] {
case 'B':
divisor = 1
case 'K':
divisor = KB
case 'M':
divisor = MB
case 'G':
divisor = GB
case 'T':
divisor = TB
case 'P':
divisor = PB
case 'E':
divisor = EB
}
}
if divisor == 0 {
return 0, &strconv.NumError{"Wrong suffix", sSize, strconv.ErrSyntax}
}
if val > maxUint64/divisor {
return 0, &strconv.NumError{"Overflow", sSize, strconv.ErrRange}
}
return val * divisor, nil
}
return 0, &strconv.NumError{"Not enough parts", sSize, strconv.ErrSyntax}
}