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/e2ee" "dev.narayana.im/narayana/telegabber/e2ee/omemo" "dev.narayana.im/narayana/telegabber/e2ee/omemo/libsignal" "dev.narayana.im/narayana/telegabber/e2ee/store/badgerstore" "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, e2eeDbPath string) (*xmpp.StreamManager, *xmpp.Component, error) { var err error if err := setupE2EE(e2eeDbPath, conf.E2EE); err != nil { return nil, nil, err } 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 } // setupE2EE wires up gateway.E2EE (see its doc comment) from config. If // E2EE is disabled (the default), gateway.E2EE still ends up non-nil, but // its Backend() always reports ok=false - every call site is expected to // treat that as "behave exactly as if this feature didn't exist" rather // than nil-checking gateway.E2EE itself. func setupE2EE(dbPath string, ec config.E2EEConfig) error { if !ec.Enabled { gateway.E2EE = e2ee.NewManager(nil) return nil } if ec.Backend != omemo.Name { return errors.Errorf("e2ee: unknown backend %q (only %q is supported)", ec.Backend, omemo.Name) } libctx, err := libsignal.NewContext() if err != nil { return errors.Wrap(err, "e2ee: libsignal.NewContext") } e2eeDB, err := badgerstore.Open(dbPath, []byte(ec.EncryptionPassphrase)) if err != nil { return errors.Wrap(err, "e2ee: badgerstore.Open") } gateway.E2EE = e2ee.NewManager(omemo.New(libctx, e2eeDB)) return 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, ¤tSession, 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...") // cancel any in-flight background operations (e.g. e2ee bundle // fetches) derived from gateway.ShutdownCtx, rather than letting them // run out their own timeout regardless of teardown gateway.CancelShutdown() 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() // flush the e2ee database, if enabled if backend, ok := gateway.E2EE.Backend(); ok { if err := backend.Close(); err != nil { log.Error(err) } } // 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} }