diff --git a/go.mod b/go.mod index f3e63111..eec95fc3 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,7 @@ tool go.mau.fi/util/cmd/maubuild require ( github.com/beeper/poly1305 v0.0.0-20250815183548-d4eede7bbf3c + github.com/coder/websocket v1.8.14 github.com/gabriel-vasile/mimetype v1.4.13 github.com/google/go-querystring v1.2.0 github.com/google/uuid v1.6.0 @@ -31,7 +32,6 @@ require ( require ( filippo.io/edwards25519 v1.2.0 // indirect github.com/beeper/argo-go v1.1.2 // indirect - github.com/coder/websocket v1.8.14 // indirect github.com/coreos/go-systemd/v22 v22.7.0 // indirect github.com/elliotchance/orderedmap/v3 v3.1.0 // indirect github.com/lib/pq v1.12.3 // indirect diff --git a/pkg/connector/client.go b/pkg/connector/client.go index 9d09084a..e4ff6631 100644 --- a/pkg/connector/client.go +++ b/pkg/connector/client.go @@ -76,6 +76,7 @@ type MetaClient struct { func (m *MetaConnector) getMessagixConfig() *messagix.Config { return &messagix.Config{ MayConnectToDGW: m.Config.ReceiveInstagramTypingIndicators, + UseMessengerDGWRealtime: m.Config.UseMessengerDGWRealtime, ClientSettings: m.Bridge.GetHTTPClientSettings(), LogRedactedBloksPayloads: m.Config.LogRedactedBloksPayloads, } @@ -335,7 +336,7 @@ func (m *MetaClient) connectWithRetry(retryCtx, ctx context.Context, attempts in } func (m *MetaClient) connectWithTable(ctx context.Context, initialTable *table.LSTable, currentUser types.UserInfo) { - zerolog.Ctx(ctx).Debug().Msg("Loaded messages page, connecting to MQTT with initial table") + zerolog.Ctx(ctx).Debug().Msg("Loaded messages page, connecting to Meta realtime with initial table") go m.handleTableLoop(ctx) var err error diff --git a/pkg/connector/config.go b/pkg/connector/config.go index 274cfcd4..5b41be9e 100644 --- a/pkg/connector/config.go +++ b/pkg/connector/config.go @@ -43,6 +43,7 @@ type Config struct { // Only affects E2EE chats right now. SendPresenceOnTyping bool `yaml:"send_presence_on_typing"` ReceiveInstagramTypingIndicators bool `yaml:"receive_instagram_typing_indicators"` + UseMessengerDGWRealtime bool `yaml:"use_messenger_dgw_realtime"` DisableViewOnce bool `yaml:"disable_view_once"` MarketplaceSpace bool `yaml:"marketplace_space"` LogRedactedBloksPayloads bool `yaml:"log_redacted_bloks_payloads"` @@ -98,6 +99,7 @@ func upgradeConfig(helper up.Helper) { helper.Copy(up.Bool, "disable_xma_always") helper.Copy(up.Bool, "send_presence_on_typing") helper.Copy(up.Bool, "receive_instagram_typing_indicators") + helper.Copy(up.Bool, "use_messenger_dgw_realtime") helper.Copy(up.Bool, "disable_view_once") helper.Copy(up.Bool, "marketplace_space") helper.Copy(up.Bool, "log_redacted_bloks_payloads") diff --git a/pkg/connector/example-config.yaml b/pkg/connector/example-config.yaml index 68a643da..a03b10c3 100644 --- a/pkg/connector/example-config.yaml +++ b/pkg/connector/example-config.yaml @@ -48,6 +48,11 @@ send_presence_on_typing: false # typing indicators; the existing connection(s) are sufficient to send # and receive typing indicators in all other cases. receive_instagram_typing_indicators: true +# Experimental: for facebook.com/messenger.com logins, use the Messenger +# Web DGW Lightspeed websocket instead of MQTT edge-chat for realtime +# Lightspeed requests. Leave disabled unless testing browser-parity +# realtime behavior. +use_messenger_dgw_realtime: false # Should view-once messages be disabled entirely? disable_view_once: false # Should FB marketplace chats have a separate space inside the main Facebook space? diff --git a/pkg/messagix/client.go b/pkg/messagix/client.go index a7d8c363..dde922a7 100644 --- a/pkg/messagix/client.go +++ b/pkg/messagix/client.go @@ -36,6 +36,7 @@ type EventHandler func(ctx context.Context, evt any) type Config struct { MayConnectToDGW bool + UseMessengerDGWRealtime bool ClientSettings exhttp.ClientSettings LogRedactedBloksPayloads bool } @@ -47,20 +48,22 @@ type Client struct { Logger zerolog.Logger Platform types.Platform - http *http.Client - httpSettings exhttp.ClientSettings - proxyAddr string - socket *Socket - dgwSocket *dgw.Socket - eventHandler EventHandler - configs *Configs - syncManager *SyncManager - - cookies *cookies.Cookies - httpProxy func(*http.Request) (*url.URL, error) - socksProxy proxy.Dialer - GetNewProxy func(reason string) (string, error) - mayConnectToDGW bool + http *http.Client + httpSettings exhttp.ClientSettings + proxyAddr string + socket *Socket + dgwSocket *dgw.Socket + dgwLightSpeed *DGWLightSpeedSocket + eventHandler EventHandler + configs *Configs + syncManager *SyncManager + + cookies *cookies.Cookies + httpProxy func(*http.Request) (*url.URL, error) + socksProxy proxy.Dialer + GetNewProxy func(reason string) (string, error) + mayConnectToDGW bool + useMessengerDGWRealtime bool device *store.Device @@ -114,7 +117,9 @@ func NewClient(cookies *cookies.Cookies, logger zerolog.Logger, cfg *Config) *Cl CSRBitmap: crypto.NewBitmap(), } cli.socket = cli.newSocketClient() + cli.dgwLightSpeed = cli.newDGWLightSpeedSocket() cli.mayConnectToDGW = cfg.MayConnectToDGW + cli.useMessengerDGWRealtime = cfg.UseMessengerDGWRealtime return cli } @@ -124,15 +129,32 @@ func (c *Client) GetCookies() *cookies.Cookies { } type dumpedState struct { - Configs *Configs - SyncStore map[int64]*socket.QueryMetadata - PacketsSent uint16 - SessionID int64 - Timestamp time.Time + Configs *Configs + SyncStore map[int64]*socket.QueryMetadata + PacketsSent uint16 + DGWStreamsSent uint16 + SessionID int64 + Timestamp time.Time } func (c *Client) DumpState() (json.RawMessage, error) { - if c == nil || c.configs == nil || c.syncManager == nil || c.socket == nil || c.socket.packetsSent == 0 || !c.socket.previouslyConnected { + if c == nil || c.configs == nil || c.syncManager == nil { + return nil, nil + } + if c.useDGWLightSpeedRealtime() { + if c.dgwLightSpeed == nil || c.dgwLightSpeed.packetsSent == 0 || !c.dgwLightSpeed.previouslyConnected { + return nil, nil + } + return json.Marshal(&dumpedState{ + Configs: c.configs, + SyncStore: c.syncManager.store, + PacketsSent: c.dgwLightSpeed.packetsSent, + DGWStreamsSent: c.dgwLightSpeed.streamsSent, + SessionID: c.socket.sessionID, + Timestamp: time.Now(), + }) + } + if c.socket == nil || c.socket.packetsSent == 0 || !c.socket.previouslyConnected { return nil, nil } return json.Marshal(&dumpedState{ @@ -162,8 +184,17 @@ func (c *Client) LoadState(state json.RawMessage) error { c.configs = dumped.Configs c.syncManager = c.newSyncManager() c.syncManager.store = dumped.SyncStore + if c.useDGWLightSpeedRealtime() { + c.dgwLightSpeed.packetsSent = dumped.PacketsSent + c.dgwLightSpeed.streamsSent = dumped.DGWStreamsSent + c.dgwLightSpeed.previouslyConnected = true + c.Logger.Info().Str("transport", "dgw_lightspeed").Msg("Loaded state") + return nil + } c.socket.packetsSent = dumped.PacketsSent - c.socket.sessionID = dumped.SessionID + if dumped.SessionID != 0 { + c.socket.sessionID = dumped.SessionID + } if c.Platform == types.Instagram { c.socket.broker = "wss://edge-chat.instagram.com/chat?" } else { @@ -308,8 +339,15 @@ func (c *Client) UpdateProxy(reason string) bool { func (c *Client) Connect(ctx context.Context) error { if c == nil { return ErrClientIsNil - } else if err := c.socket.CanConnect(); err != nil { - return err + } + if c.useDGWLightSpeedRealtime() { + if err := c.dgwLightSpeed.CanConnect(); err != nil { + return err + } + } else { + if err := c.socket.CanConnect(); err != nil { + return err + } } ctx, cancel := context.WithCancel(ctx) oldCancel := c.stopCurrentConnections.Swap(&cancel) @@ -328,7 +366,7 @@ func (c *Client) Connect(ctx context.Context) error { for { c.canSendMessages.Clear() // In case we're reconnecting from a normal network error connectStart := time.Now() - err := c.socket.Connect(ctx) + err := c.connectMainRealtime(ctx) c.canSendMessages.Clear() if ctx.Err() != nil { zerolog.Ctx(ctx).Warn(). @@ -382,6 +420,17 @@ func (c *Client) Connect(ctx context.Context) error { return nil } +func (c *Client) connectMainRealtime(ctx context.Context) error { + if c.useDGWLightSpeedRealtime() { + return c.dgwLightSpeed.Connect(ctx) + } + return c.socket.Connect(ctx) +} + +func (c *Client) useDGWLightSpeedRealtime() bool { + return c != nil && c.useMessengerDGWRealtime && (c.Platform == types.Facebook || c.Platform == types.Messenger) +} + func (c *Client) connectDGW(ctx context.Context) error { reconnectIn := 2 * time.Second for { @@ -431,7 +480,12 @@ func (c *Client) Disconnect() { if fn := c.stopCurrentConnections.Load(); fn != nil { (*fn)() } - c.socket.Disconnect() + if c.socket != nil { + c.socket.Disconnect() + } + if c.dgwLightSpeed != nil { + c.dgwLightSpeed.Disconnect() + } if c.dgwSocket != nil { c.dgwSocket.Disconnect() } @@ -441,7 +495,13 @@ func (c *Client) Disconnect() { } func (c *Client) IsConnected() bool { - return c != nil && c.socket.conn != nil + if c == nil { + return false + } + if c.useDGWLightSpeedRealtime() { + return c.dgwLightSpeed != nil && c.dgwLightSpeed.conn != nil + } + return c.socket != nil && c.socket.conn != nil } func (c *Client) GetEndpoint(name string) string { @@ -512,8 +572,15 @@ func (c *Client) ForceReconnect() { if c == nil { return } - c.socket.Disconnect() - c.dgwSocket.Disconnect() + if c.socket != nil { + c.socket.Disconnect() + } + if c.dgwLightSpeed != nil { + c.dgwLightSpeed.Disconnect() + } + if c.dgwSocket != nil { + c.dgwSocket.Disconnect() + } } func (c *Client) FetchMoreThreads(ctx context.Context, syncGroup int64) (*socket.KeyStoreData, *table.LSTable, error) { @@ -542,13 +609,13 @@ func (c *Client) FetchMoreThreads(ctx context.Context, syncGroup int64) (*socket return nil, nil, err } - resp, err := c.socket.makeLSRequest(ctx, payload, 3) + resp, err := c.makeRealtimeLSRequest(ctx, payload, 3) if err != nil { return nil, nil, err } resp.Finish() - c.socket.postHandlePublishResponse(resp.Table) + c.PostHandlePublishResponse(resp.Table) return keyStore, resp.Table, nil } diff --git a/pkg/messagix/configs.go b/pkg/messagix/configs.go index 2886c08a..f65b943e 100644 --- a/pkg/messagix/configs.go +++ b/pkg/messagix/configs.go @@ -34,6 +34,9 @@ func (c *Configs) SetupConfigs(ctx context.Context, ls *table.LSTable) (*table.L if c.client.socket != nil { c.client.socket.previouslyConnected = false } + if c.client.dgwLightSpeed != nil { + c.client.dgwLightSpeed.previouslyConnected = false + } authenticated := c.client.IsAuthenticated() c.WebSessionID = methods.GenerateWebsessionID(authenticated) c.LSDToken = c.BrowserConfigTable.LSD.Token @@ -58,7 +61,7 @@ func (c *Configs) SetupConfigs(ctx context.Context, ls *table.LSTable) (*table.L 936619743392459, ) } else { - if c.BrowserConfigTable.MqttWebConfig.Endpoint == "" { + if c.BrowserConfigTable.MqttWebConfig.Endpoint == "" && !c.client.useDGWLightSpeedRealtime() { return ls, fmt.Errorf("MQTT broker endpoint not found in page response (MqttWebConfig.Endpoint is empty)") } c.client.socket.broker = c.BrowserConfigTable.MqttWebConfig.Endpoint diff --git a/pkg/messagix/dgw_lightspeed.go b/pkg/messagix/dgw_lightspeed.go new file mode 100644 index 00000000..3ea04915 --- /dev/null +++ b/pkg/messagix/dgw_lightspeed.go @@ -0,0 +1,560 @@ +package messagix + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/http" + "net/url" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/coder/websocket" + "github.com/google/uuid" + "go.mau.fi/util/exhttp" + + "go.mau.fi/mautrix-meta/pkg/messagix/dgw" + "go.mau.fi/mautrix-meta/pkg/messagix/packets" + "go.mau.fi/mautrix-meta/pkg/messagix/socket" + "go.mau.fi/mautrix-meta/pkg/messagix/table" + "go.mau.fi/mautrix-meta/pkg/messagix/useragent" +) + +const dgwLightSpeedEndpoint = "wss://gateway.facebook.com/ws/lightspeed" + +type DGWLightSpeedSocket struct { + client *Client + conn *websocket.Conn + responseHandler *ResponseHandler + mu sync.Mutex + packetsSent uint16 + streamsSent uint16 + + previouslyConnected bool + lastReceived atomic.Int64 +} + +func (c *Client) newDGWLightSpeedSocket() *DGWLightSpeedSocket { + return &DGWLightSpeedSocket{ + client: c, + responseHandler: &ResponseHandler{ + client: c, + requestChannels: make(map[uint16]chan any), + packetChannels: make(map[uint16]chan any), + }, + } +} + +func (s *DGWLightSpeedSocket) CanConnect() error { + if s.conn != nil { + return fmt.Errorf("DGW Lightspeed: %w", socket.ErrSocketAlreadyOpen) + } else if !s.client.IsAuthenticated() { + return fmt.Errorf("DGW Lightspeed: %w", socket.ErrNotAuthenticated) + } + return nil +} + +func (s *DGWLightSpeedSocket) Connect(ctx context.Context) error { + if err := s.CanConnect(); err != nil { + return err + } + + opts := &websocket.DialOptions{ + // c.http can't be used directly: coder/websocket rejects clients with a Timeout set. + HTTPClient: &http.Client{Transport: s.client.http.Transport}, + HTTPHeader: s.getConnHeaders(), + } + socketURL := s.getConnURL() + + s.client.Logger.Debug().Str("url", socketURL).Msg("Dialing DGW Lightspeed socket") + conn, resp, err := websocket.Dial(ctx, socketURL, opts) + if err != nil { + statusCode := 999 + if resp != nil { + statusCode = resp.StatusCode + } + return fmt.Errorf("DGW Lightspeed: %w: %w (status code %d)", socket.ErrDial, err, statusCode) + } + conn.SetReadLimit(-1) + s.conn = conn + s.lastReceived.Store(time.Now().UnixNano()) + + runCtx, cancel := context.WithCancel(ctx) + defer cancel() + + readDone := make(chan error, 1) + go func() { + readDone <- s.readLoop(runCtx, conn) + cancel() + }() + + if err = s.handleReady(runCtx); err != nil { + _ = conn.CloseNow() + cancel() + readErr := <-readDone + if readErr != nil && !errors.Is(readErr, context.Canceled) { + s.client.Logger.Debug().Err(readErr).Msg("DGW Lightspeed read loop stopped after ready failure") + } + return fmt.Errorf("DGW Lightspeed ready failed: %w", err) + } + + err = <-readDone + if err != nil { + return fmt.Errorf("DGW Lightspeed: %w: %w", socket.ErrInReadLoop, err) + } + return nil +} + +func (s *DGWLightSpeedSocket) Disconnect() { + if s != nil && s.conn != nil { + _ = s.conn.Close(websocket.StatusNormalClosure, "") + } +} + +func (s *DGWLightSpeedSocket) readLoop(ctx context.Context, conn *websocket.Conn) error { + defer func() { + s.conn = nil + s.responseHandler.CancelAllRequests() + }() + + done := make(chan struct{}) + defer close(done) + + go s.pingLoop(ctx, conn, done) + go s.pongTimeoutLoop(conn, done) + + s.client.Logger.Debug().Msg("DGW Lightspeed connection established, starting read loop") + for { + messageType, payload, err := conn.Read(ctx) + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + var ce websocket.CloseError + if errors.As(err, &ce) { + s.client.Logger.Info().Int("code", int(ce.Code)).Str("text", ce.Reason).Msg("DGW Lightspeed websocket closed by server") + return fmt.Errorf("closed by server: %d %s", ce.Code, ce.Reason) + } + s.client.Logger.Err(err).Msg("Error reading message from DGW Lightspeed socket") + return fmt.Errorf("failed to read message: %w", err) + } + + s.lastReceived.Store(time.Now().UnixNano()) + switch messageType { + case websocket.MessageText: + s.client.Logger.Warn().Bytes("bytes", payload).Msg("Unexpected text message in DGW Lightspeed websocket") + case websocket.MessageBinary: + s.handleBinaryMessage(ctx, payload) + } + } +} + +func (s *DGWLightSpeedSocket) pingLoop(ctx context.Context, conn *websocket.Conn, done <-chan struct{}) { + ticker := time.NewTicker(PingInterval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + if err := s.sendFrames(&dgw.PingFrame{}); err != nil { + s.client.Logger.Err(err).Msg("Error sending DGW Lightspeed ping") + _ = conn.CloseNow() + return + } + case <-done: + return + case <-ctx.Done(): + return + } + } +} + +func (s *DGWLightSpeedSocket) pongTimeoutLoop(conn *websocket.Conn, done <-chan struct{}) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + for { + select { + case <-ticker.C: + lastReceived := time.Unix(0, s.lastReceived.Load()) + if time.Since(lastReceived) > PongTimeout { + s.client.Logger.Error().Msg("DGW Lightspeed pong timeout") + _ = conn.CloseNow() + return + } + case <-done: + return + } + } +} + +func (s *DGWLightSpeedSocket) handleBinaryMessage(ctx context.Context, data []byte) { + for len(data) > 0 { + frame := dgw.CheckFrameType(data) + rest, err := frame.Unmarshal(data) + if err != nil { + s.client.Logger.Warn().Err(err).Bytes("bytes", data).Msg("Failed to unmarshal DGW Lightspeed frame, dropping websocket message") + return + } + s.handleFrame(ctx, frame) + data = rest + } +} + +func (s *DGWLightSpeedSocket) handleFrame(ctx context.Context, frame dgw.Frame) { + switch f := frame.(type) { + case *dgw.UnsupportedFrame: + if f.MayHaveLostData() { + s.client.Logger.Warn().Bytes("bytes", f.Raw).Msg("Encountered unknown DGW Lightspeed frame, dropping with potential data loss") + } + case *dgw.PingFrame: + if err := s.sendFrames(&dgw.PongFrame{}); err != nil { + s.client.Logger.Err(err).Msg("Failed to send DGW Lightspeed pong") + } + case *dgw.PongFrame: + s.client.Logger.Trace().Msg("Got DGW Lightspeed pong") + case *dgw.OpenFrame: + if f.Parameters.StatusCode != 0 && f.Parameters.StatusCode != http.StatusOK { + s.client.Logger.Warn().Any("frame", f).Msg("DGW Lightspeed stream open returned non-OK status") + } + case *dgw.AckFrame: + s.client.Logger.Trace().Uint16("stream_id", uint16(f.StreamID)).Uint16("ack_id", f.AckID).Msg("Got DGW Lightspeed ack") + case *dgw.DataFrame: + if f.RequiresAck { + if err := s.sendFrames(&dgw.AckFrame{StreamID: f.StreamID, AckID: f.AckID}); err != nil { + s.client.Logger.Err(err).Uint16("stream_id", uint16(f.StreamID)).Uint16("ack_id", f.AckID).Msg("Failed to ack DGW Lightspeed data frame") + } + } + s.handleDataFrame(ctx, f) + } +} + +func (s *DGWLightSpeedSocket) handleDataFrame(ctx context.Context, frame *dgw.DataFrame) { + jsonPayload, err := extractDGWJSON(frame.Payload) + if err != nil { + s.client.Logger.Trace().Err(err).Bytes("payload", frame.Payload).Msg("Ignoring DGW Lightspeed data frame without JSON payload") + return + } + + var data PublishResponseData + if err = json.Unmarshal(jsonPayload, &data); err != nil { + s.client.Logger.Warn().Err(err).Bytes("json", jsonPayload).Msg("Failed to parse DGW Lightspeed publish response JSON") + return + } + if data.Payload == "" { + s.client.Logger.Trace().Bytes("json", jsonPayload).Msg("Ignoring DGW Lightspeed data frame without Lightspeed payload") + return + } + resp := &Event_PublishResponse{ + Topic: string(LS_RESP), + Data: data, + MessageIdentifier: uint16(frame.StreamID), + QoS: packets.QOS_LEVEL_0, + } + s.handlePublishResponseEvent(ctx, resp) +} + +func extractDGWJSON(payload []byte) ([]byte, error) { + msgStart := bytes.IndexByte(payload, '{') + msgEnd := bytes.LastIndexByte(payload, '}') + if msgStart < 0 || msgEnd < msgStart { + return nil, fmt.Errorf("JSON object not found") + } + return payload[msgStart : msgEnd+1], nil +} + +func (s *DGWLightSpeedSocket) handlePublishResponseEvent(ctx context.Context, resp *Event_PublishResponse) { + requestID := uint16(resp.Data.RequestID) + hasRequest := s.responseHandler.hasPacket(requestID) + switch resp.Topic { + case string(LS_RESP): + resp.Finish() + if hasRequest { + if ok := s.responseHandler.updateRequestChannel(requestID, resp); !ok { + s.client.Logger.Warn().Int64("request_id", resp.Data.RequestID).Msg("Dropped DGW Lightspeed response to request") + } + } else if resp.Data.RequestID == 0 { + s.client.HandleEvent(ctx, resp) + } else { + s.client.Logger.Debug().Int64("request_id", resp.Data.RequestID).Msg("Got unexpected DGW Lightspeed publish response") + } + default: + s.client.Logger.Info().Any("topic", resp.Topic).Any("data", resp.Data).Msg("Got unknown DGW Lightspeed publish response topic") + } +} + +func (c *Client) makeRealtimeLSRequest(ctx context.Context, payload []byte, t int) (*Event_PublishResponse, error) { + if c.useDGWLightSpeedRealtime() { + return c.dgwLightSpeed.makeLSRequest(ctx, payload, t) + } + return c.socket.makeLSRequest(ctx, payload, t) +} + +func (s *DGWLightSpeedSocket) makeLSRequest(ctx context.Context, payload []byte, t int) (*Event_PublishResponse, error) { + packetID := s.SafePacketID() + streamID := s.SafeStreamID() + lsPayload := &SocketLSRequestPayload{ + AppId: s.client.messengerAppID(), + Payload: string(payload), + RequestId: int(packetID), + Type: t, + } + + jsonPayload, err := json.Marshal(lsPayload) + if err != nil { + return nil, err + } + + if t != 4 { + s.responseHandler.addRequestChannel(packetID) + } + err = s.sendFrames( + &dgw.OpenFrame{StreamID: streamID}, + &dgw.DataFrame{ + StreamID: streamID, + Payload: append([]byte{0x80}, jsonPayload...), + RequiresAck: true, + AckID: 0, + }, + ) + if err != nil { + if t != 4 { + s.responseHandler.deleteDetails(packetID, RequestChannel) + } + return nil, err + } + + if t == 4 { + return nil, nil + } + return s.responseHandler.waitForPubResponseDetails(ctx, packetID) +} + +func (s *DGWLightSpeedSocket) sendFrames(frames ...dgw.Frame) error { + data, err := marshalDGWFrames(frames...) + if err != nil { + return err + } + return s.sendData(data) +} + +func marshalDGWFrames(frames ...dgw.Frame) ([]byte, error) { + var data []byte + for _, frame := range frames { + frameBytes, err := frame.Marshal() + if err != nil { + return nil, err + } + data = append(data, frameBytes...) + } + return data, nil +} + +func (s *DGWLightSpeedSocket) sendData(data []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + conn := s.conn + if conn == nil { + return fmt.Errorf("not connected") + } + ctx, cancel := context.WithTimeout(context.Background(), WriteTimeout) + defer cancel() + err := conn.Write(ctx, websocket.MessageBinary, data) + if exhttp.IsNetworkError(err) { + closeErr := conn.CloseNow() + if closeErr != nil && !errors.Is(closeErr, net.ErrClosed) { + s.client.Logger.Debug().Err(closeErr).Msg("Error closing DGW Lightspeed connection after network error") + return errors.Join(err, closeErr) + } + return err + } else if err != nil { + return fmt.Errorf("failed to write to DGW Lightspeed websocket: %w", err) + } + return nil +} + +func (s *DGWLightSpeedSocket) SafePacketID() uint16 { + s.mu.Lock() + defer s.mu.Unlock() + + s.packetsSent++ + if s.packetsSent == 0 { + s.packetsSent = 1 + } + return s.packetsSent +} + +func (s *DGWLightSpeedSocket) SafeStreamID() dgw.StreamID { + s.mu.Lock() + defer s.mu.Unlock() + + streamID := s.streamsSent + s.streamsSent++ + return dgw.StreamID(streamID) +} + +func (s *DGWLightSpeedSocket) handleReady(ctx context.Context) error { + if s.previouslyConnected { + s.client.canSendMessages.Set() + err := s.client.syncManager.EnsureSyncedSocket(ctx, minimalFBReconnectSync) + if err != nil { + return fmt.Errorf("failed to sync after DGW Lightspeed reconnect: %w", err) + } + s.client.HandleEvent(ctx, &Event_Reconnected{}) + return nil + } + + s.client.canSendMessages.Set() + if err := s.sendInitialSyncTasks(ctx); err != nil { + return err + } + if _, err := s.client.ExecuteTasks(ctx, &socket.ReportAppStateTask{AppState: table.FOREGROUND, RequestID: uuid.NewString()}); err != nil { + return fmt.Errorf("failed to report app state: %w", err) + } + if err := s.client.syncManager.EnsureSyncedSocket(ctx, minimalFBInitialSync); err != nil { + return fmt.Errorf("failed to ensure initial databases are synced through DGW Lightspeed: %w", err) + } + + s.client.HandleEvent(ctx, (&Event_Ready{ + client: s.client, + ConnectionCode: CONNECTION_ACCEPTED, + }).Finish()) + s.previouslyConnected = true + + return nil +} + +func (s *DGWLightSpeedSocket) sendInitialSyncTasks(ctx context.Context) error { + tskm := s.client.newTaskManager() + ptks := s.client.configs.ParentThreadKeys + if len(ptks) == 0 { + s.client.Logger.Warn().Msg("Parent thread keys are not known") + ptks = []int64{-1} + } else if !slicesContainsInt64(ptks, -1) { + s.client.Logger.Warn().Ints64("ptks", ptks).Msg("Parent thread keys don't contain -1") + ptks = append(ptks, -1) + } + for _, tk := range ptks { + tskm.AddNewTask(&socket.FetchThreadsTask{ + IsAfter: 0, + ParentThreadKey: tk, + ReferenceThreadKey: 0, + ReferenceActivityTimestamp: 9999999999999, + AdditionalPagesToFetch: 0, + Cursor: s.client.syncManager.GetCursor(1), + SyncGroup: 1, + }) + tskm.AddNewTask(&socket.FetchThreadsTask{ + IsAfter: 0, + ParentThreadKey: tk, + ReferenceThreadKey: 0, + ReferenceActivityTimestamp: 9999999999999, + AdditionalPagesToFetch: 0, + SyncGroup: 95, + }) + } + + syncGroupKeyStore1 := s.client.syncManager.getSyncGroupKeyStore(1) + if syncGroupKeyStore1 != nil { + tskm.AddNewTask(&socket.FetchThreadsTask{ + IsAfter: 0, + ParentThreadKey: syncGroupKeyStore1.ParentThreadKey, + ReferenceThreadKey: syncGroupKeyStore1.MinThreadKey, + ReferenceActivityTimestamp: syncGroupKeyStore1.MinLastActivityTimestampMs, + AdditionalPagesToFetch: 0, + Cursor: s.client.syncManager.GetCursor(1), + SyncGroup: 1, + }) + tskm.AddNewTask(&socket.FetchThreadsTask{ + IsAfter: 0, + ParentThreadKey: syncGroupKeyStore1.ParentThreadKey, + ReferenceThreadKey: syncGroupKeyStore1.MinThreadKey, + ReferenceActivityTimestamp: syncGroupKeyStore1.MinLastActivityTimestampMs, + AdditionalPagesToFetch: 0, + SyncGroup: 95, + }) + } + + payload, err := tskm.FinalizePayload() + if err != nil { + return fmt.Errorf("failed to finalize DGW Lightspeed sync tasks: %w", err) + } + + s.client.Logger.Trace().Any("data", string(payload)).Msg("DGW Lightspeed sync groups tasks") + if _, err = s.makeLSRequest(ctx, payload, 3); err != nil { + return fmt.Errorf("failed to send DGW Lightspeed sync tasks: %w", err) + } + return nil +} + +func slicesContainsInt64(values []int64, needle int64) bool { + for _, value := range values { + if value == needle { + return true + } + } + return false +} + +func (s *DGWLightSpeedSocket) getConnHeaders() http.Header { + h := http.Header{} + h.Set("cookie", s.client.cookies.String()) + h.Set("user-agent", useragent.UserAgent) + h.Set("origin", s.client.GetEndpoint("base_url")) + h.Set("cache-control", "no-cache") + h.Set("pragma", "no-cache") + return h +} + +func (s *DGWLightSpeedSocket) getConnURL() string { + query := &url.Values{} + query.Add("x-dgw-appid", s.client.messengerAppID()) + query.Add("x-dgw-appversion", "0") + query.Add("x-dgw-authtype", "1:0") + query.Add("x-dgw-version", "5") + query.Add("x-dgw-uuid", s.client.messengerDGWUUID()) + query.Add("x-dgw-tier", "prod") + query.Add("x-dgw-loggingid", uuid.NewString()) + if region := s.client.configs.BrowserConfigTable.MessengerWebRegion.Region; region != "" { + query.Add("x-dgw-regionhint", strings.ToUpper(region)) + } + query.Add("x-dgw-deviceid", uuid.NewString()) + return dgwLightSpeedEndpoint + "?" + query.Encode() +} + +func (c *Client) messengerAppID() string { + if appID := nonZeroString(c.configs.BrowserConfigTable.CurrentUserInitialData.AppID); appID != "" { + return appID + } + if appID := c.configs.BrowserConfigTable.MessengerWebInitData.AppID; appID != 0 { + return strconv.FormatInt(appID, 10) + } + if appID := c.configs.BrowserConfigTable.MqttWebConfig.AppID; appID != 0 { + return strconv.FormatInt(appID, 10) + } + return "2220391788200892" +} + +func (c *Client) messengerDGWUUID() string { + if accountID := nonZeroString(c.configs.BrowserConfigTable.CurrentUserInitialData.AccountID); accountID != "" { + return accountID + } + if userID := nonZeroString(c.configs.BrowserConfigTable.CurrentUserInitialData.UserID); userID != "" { + return userID + } + if userID := c.configs.BrowserConfigTable.MessengerWebInitData.UserID.String(); userID != "" { + return userID + } + return strconv.FormatInt(c.cookies.GetUserID(), 10) +} + +func nonZeroString(value string) string { + if value == "" || value == "0" { + return "" + } + return value +} diff --git a/pkg/messagix/dgw_lightspeed_test.go b/pkg/messagix/dgw_lightspeed_test.go new file mode 100644 index 00000000..912ff2fc --- /dev/null +++ b/pkg/messagix/dgw_lightspeed_test.go @@ -0,0 +1,87 @@ +package messagix + +import ( + "bytes" + "encoding/json" + "testing" + + "go.mau.fi/mautrix-meta/pkg/messagix/dgw" +) + +func TestDGWLightSpeedRequestFrameShape(t *testing.T) { + lsPayload := &SocketLSRequestPayload{ + AppId: "2220391788200892", + Payload: `{"epoch_id":1}`, + RequestId: 4, + Type: 3, + } + jsonPayload, err := json.Marshal(lsPayload) + if err != nil { + t.Fatal(err) + } + data, err := marshalDGWFrames( + &dgw.OpenFrame{StreamID: 0}, + &dgw.DataFrame{ + StreamID: 0, + Payload: append([]byte{0x80}, jsonPayload...), + RequiresAck: true, + AckID: 0, + }, + ) + if err != nil { + t.Fatal(err) + } + + expectedPrefix := []byte{0x0f, 0x00, 0x00, 0x02, 0x00, 0x00, '{', '}', 0x0d, 0x00, 0x00} + if !bytes.HasPrefix(data, expectedPrefix) { + t.Fatalf("unexpected DGW request frame prefix: % x", data[:len(expectedPrefix)]) + } + + frame := dgw.CheckFrameType(data) + rest, err := frame.Unmarshal(data) + if err != nil { + t.Fatal(err) + } + openFrame, ok := frame.(*dgw.OpenFrame) + if !ok { + t.Fatalf("first frame has type %T, expected *dgw.OpenFrame", frame) + } + if openFrame.StreamID != 0 { + t.Fatalf("unexpected open stream ID: %d", openFrame.StreamID) + } + + frame = dgw.CheckFrameType(rest) + _, err = frame.Unmarshal(rest) + if err != nil { + t.Fatal(err) + } + dataFrame, ok := frame.(*dgw.DataFrame) + if !ok { + t.Fatalf("second frame has type %T, expected *dgw.DataFrame", frame) + } + if !dataFrame.RequiresAck { + t.Fatal("DGW request data frame should request an ack") + } + if dataFrame.Payload[0] != 0x80 { + t.Fatalf("unexpected DGW data payload prefix: %x", dataFrame.Payload[0]) + } + var decoded SocketLSRequestPayload + err = json.Unmarshal(dataFrame.Payload[1:], &decoded) + if err != nil { + t.Fatal(err) + } + if decoded != *lsPayload { + t.Fatalf("decoded payload mismatch: got %+v, want %+v", decoded, *lsPayload) + } +} + +func TestExtractDGWJSON(t *testing.T) { + want := []byte(`{"request_id":4,"payload":"{\"step\":true}","sp":[],"target":3}`) + got, err := extractDGWJSON(append([]byte{0x80}, want...)) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, want) { + t.Fatalf("unexpected JSON: got %s, want %s", got, want) + } +} diff --git a/pkg/messagix/events.go b/pkg/messagix/events.go index ce25d4f1..4fac570a 100644 --- a/pkg/messagix/events.go +++ b/pkg/messagix/events.go @@ -147,15 +147,22 @@ func (s *Socket) handleACKEvent(ackData AckEvent) { } func (s *Socket) postHandlePublishResponse(tbl *table.LSTable) { + s.client.postHandlePublishResponse(tbl) +} + +func (c *Client) postHandlePublishResponse(tbl *table.LSTable) { + if c == nil || c.syncManager == nil || tbl == nil { + return + } syncGroupsNeedUpdate := methods.NeedUpdateSyncGroups(tbl) if syncGroupsNeedUpdate { - s.client.Logger.Debug(). + c.Logger.Debug(). Any("LSExecuteFirstBlockForSyncTransaction", tbl.LSExecuteFirstBlockForSyncTransaction). Any("LSUpsertSyncGroupThreadsRange", tbl.LSUpsertSyncGroupThreadsRange). Msg("Updating sync groups") - err := s.client.syncManager.updateSyncGroupCursors(tbl) + err := c.syncManager.updateSyncGroupCursors(tbl) if err != nil { - s.client.Logger.Err(err).Msg("Failed to sync transactions from publish response event") + c.Logger.Err(err).Msg("Failed to sync transactions from publish response event") } } } @@ -164,9 +171,7 @@ func (c *Client) PostHandlePublishResponse(tbl *table.LSTable) { if c == nil { return } - if s := c.socket; s != nil { - s.postHandlePublishResponse(tbl) - } + c.postHandlePublishResponse(tbl) } func (s *Socket) handlePublishResponseEvent(ctx context.Context, resp *Event_PublishResponse, isQueue bool) (addToQueue bool) { diff --git a/pkg/messagix/syncManager.go b/pkg/messagix/syncManager.go index e48ec4ef..6c273c5b 100644 --- a/pkg/messagix/syncManager.go +++ b/pkg/messagix/syncManager.go @@ -117,7 +117,7 @@ func (sm *SyncManager) SyncSocketData(ctx context.Context, databaseID int64, db RawJSON("payload", jsonPayload). Int64("database_id", databaseID). Msg("Syncing database via socket") - resp, err := sm.client.socket.makeLSRequest(ctx, jsonPayload, t) + resp, err := sm.client.makeRealtimeLSRequest(ctx, jsonPayload, t) if err != nil { return fmt.Errorf("failed to make lightspeed socket request with DatabaseQuery byte payload (databaseID=%d): %w", databaseID, err) } diff --git a/pkg/messagix/threads.go b/pkg/messagix/threads.go index b8bad56e..977bdae8 100644 --- a/pkg/messagix/threads.go +++ b/pkg/messagix/threads.go @@ -27,7 +27,7 @@ func (c *Client) ExecuteTasks(ctx context.Context, tasks ...socket.Task) (*table return nil, fmt.Errorf("failed to finalize payload: %w", err) } - resp, err := c.socket.makeLSRequest(ctx, payload, 3) + resp, err := c.makeRealtimeLSRequest(ctx, payload, 3) if err != nil { return nil, err } @@ -61,6 +61,6 @@ func (c *Client) ExecuteStatelessTask(ctx context.Context, task socket.Task) err if err != nil { return fmt.Errorf("failed to marshal outer task %s payload: %w", label, err) } - _, err = c.socket.makeLSRequest(ctx, outerPayloadMarshalled, 4) + _, err = c.makeRealtimeLSRequest(ctx, outerPayloadMarshalled, 4) return err }