package broker import ( "context" "encoding/json" "log/slog" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) const handshakeTimeout = 30 * time.Second // PresenceSink 供身份线订阅上下线(与 presence.Service 的 SetOnline/SetOffline 对齐)。 type PresenceSink interface { SetOnline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error SetOffline(ctx context.Context, endpointID string, connID port.ConnID, atMs int64) error } // HelloLimits 握手响应里的服务器限制。 type HelloLimits struct { MaxBodyBytes int MaxMetaBytes int MaxFrameBytes int MaxTTLSeconds int64 MaxScheduleSeconds int64 AckTimeoutSeconds int64 ServerVersion string } // Session 处理握手、logout、上下线落库,并转发其余上行给 Inner。 type Session struct { b *Broker login *Login inner port.UplinkHandler presence PresenceSink limits HelloLimits log *slog.Logger now func() time.Time } // SessionOptions 装配 Session。 type SessionOptions struct { Login *Login Inner port.UplinkHandler Presence PresenceSink Limits HelloLimits Logger *slog.Logger Now func() time.Time } // NewSession 创建会话层;调用 Attach 绑定 Broker 后再接连接。 func NewSession(opts SessionOptions) *Session { inner := opts.Inner if inner == nil { inner = port.StubUplinkHandler{} } log := opts.Logger if log == nil { log = slog.Default() } now := opts.Now if now == nil { now = time.Now } lim := opts.Limits if lim.ServerVersion == "" { lim.ServerVersion = "0.1.0" } if lim.MaxBodyBytes == 0 { lim.MaxBodyBytes = protocol.DefaultMaxBodyBytes } if lim.MaxMetaBytes == 0 { lim.MaxMetaBytes = protocol.DefaultMaxMetaBytes } if lim.MaxFrameBytes == 0 { lim.MaxFrameBytes = protocol.DefaultMaxFrameBytes } if lim.MaxTTLSeconds == 0 { lim.MaxTTLSeconds = 2592000 } if lim.MaxScheduleSeconds == 0 { lim.MaxScheduleSeconds = 31536000 } if lim.AckTimeoutSeconds == 0 { lim.AckTimeoutSeconds = 300 } return &Session{ login: opts.Login, inner: inner, presence: opts.Presence, limits: lim, log: log, now: now, } } // Attach 绑定 Broker(PublishDown / Disconnect / 连接表)。 func (s *Session) Attach(b *Broker) { s.b = b } // SetPresence 接线时在创建最终 presence 实现后注入(可替换占位)。 func (s *Session) SetPresence(p PresenceSink) { s.presence = p } func (s *Session) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { if s.b != nil { s.b.startHandshakeDeadline(conn.EndpointID, conn.ConnID, handshakeTimeout) } return s.inner.OnSessionEstablished(ctx, conn) } func (s *Session) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error { return s.inner.OnHandshakeComplete(ctx, hs) } func (s *Session) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) { // 正常路径由 hooks 调 HandleDisconnect(带 isCurrent)。 // 此方法满足 UplinkHandler;直接调用时按非当前处理,避免误标离线。 s.HandleDisconnect(ctx, conn, reason, false) } // HandleDisconnect 由 hooks 在确知 isCurrent 后调用(含落库与 presence)。 func (s *Session) HandleDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason, isCurrent bool) { if s.b != nil { s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID) } if isCurrent && s.login != nil { atMs := s.now().UnixMilli() if err := s.login.SetOfflineSince(ctx, conn.EndpointID, atMs); err != nil { s.log.Error("set offline_since", "endpoint", conn.EndpointID, "err", err) } if s.presence != nil { if err := s.presence.SetOffline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil { s.log.Error("presence offline", "endpoint", conn.EndpointID, "err", err) } } } s.inner.OnDisconnect(ctx, conn, reason) } func (s *Session) HandleUplink(ctx context.Context, conn port.ConnInfo, payload []byte) error { if s.b == nil { return nil } st := s.b.connStateOf(conn.EndpointID, conn.ConnID) if st == nil { return nil } frame, err := protocol.Decode(payload) if err != nil { s.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error()) return nil } st.mu.Lock() ready := st.handshook st.mu.Unlock() switch f := frame.(type) { case *protocol.Hello: return s.handleHello(ctx, conn, st, f) case *protocol.SelfLogout: if !ready { s.replyErr(ctx, conn, f.RID, protocol.CodeNotReady, "handshake required") return nil } return s.handleLogout(ctx, conn, f) default: if !ready { rid := peekRID(payload) s.replyErr(ctx, conn, rid, protocol.CodeNotReady, "handshake required") return nil } return s.inner.HandleUplink(ctx, conn, payload) } } func (s *Session) handleHello(ctx context.Context, conn port.ConnInfo, st *connState, hello *protocol.Hello) error { if err := hello.Validate(); err != nil { code := protocol.CodeBadRequest if pe, ok := err.(*protocol.Error); ok { code = pe.Code } s.replyErr(ctx, conn, hello.RID, code, err.Error()) return nil } st.mu.Lock() if st.handshook { st.mu.Unlock() s.replyErr(ctx, conn, hello.RID, protocol.CodeBadRequest, "already handshook") return nil } st.mu.Unlock() if !s.b.hasDownSub(st) { go func() { _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectIdle) }() return nil } maxRecv := 0 if hello.MaxReceiveBytes != nil { maxRecv = *hello.MaxReceiveBytes } s.b.SetMaxReceiveBytes(conn.EndpointID, conn.ConnID, maxRecv) if conn.SessionToken != "" && s.login != nil { if keep, chkErr := s.login.TokenMatchesDB(ctx, conn.EndpointID, conn.SessionToken); chkErr != nil { s.log.Error("re-read session token", "endpoint", conn.EndpointID, "err", chkErr) } else if !keep { conn.SessionToken = "" } } data := protocol.HelloData{ ServerTimeMs: s.now().UnixMilli(), ServerVersion: s.limits.ServerVersion, MaxBodyBytes: s.limits.MaxBodyBytes, MaxMetaBytes: s.limits.MaxMetaBytes, MaxFrameBytes: s.limits.MaxFrameBytes, MaxTTLSeconds: s.limits.MaxTTLSeconds, MaxScheduleSeconds: s.limits.MaxScheduleSeconds, AckTimeoutSeconds: s.limits.AckTimeoutSeconds, } if conn.SessionToken != "" { data.SessionToken = conn.SessionToken } raw, err := protocol.Marshal(data) if err != nil { return err } resp := protocol.Resp{ V: protocol.Version, Type: protocol.TypeResp, RID: hello.RID, OK: true, Data: raw, } if err := s.publishJSON(ctx, conn, resp, 1); err != nil { return err } atMs := s.now().UnixMilli() if s.login != nil { if err := s.login.SetOnlineSince(ctx, conn.EndpointID, atMs); err != nil { s.log.Error("set online_since", "endpoint", conn.EndpointID, "err", err) } } if s.presence != nil { if err := s.presence.SetOnline(ctx, conn.EndpointID, conn.ConnID, atMs); err != nil { s.log.Error("presence online", "endpoint", conn.EndpointID, "err", err) } } st.mu.Lock() st.handshook = true st.mu.Unlock() s.b.cancelHandshakeDeadline(conn.EndpointID, conn.ConnID) hs := port.HandshakeInfo{ ConnInfo: conn, MaxReceiveBytes: maxRecv, Client: hello.Client, } return s.inner.OnHandshakeComplete(ctx, hs) } func (s *Session) handleLogout(ctx context.Context, conn port.ConnInfo, req *protocol.SelfLogout) error { if err := req.Validate(); err != nil { code := protocol.CodeBadRequest if pe, ok := err.(*protocol.Error); ok { code = pe.Code } s.replyErr(ctx, conn, req.RID, code, err.Error()) return nil } if s.login != nil { if err := s.login.ClearSession(ctx, conn.EndpointID); err != nil { s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "clear session failed") return nil } } resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: req.RID, OK: true} raw, err := protocol.Marshal(resp) if err != nil { s.replyErr(ctx, conn, req.RID, protocol.CodeBusy, "marshal logout resp") return nil } if pubErr := s.b.PublishThenDisconnect(ctx, conn.EndpointID, conn.ConnID, raw, 1, port.DisconnectNormal); pubErr != nil { s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", pubErr) go func() { _ = s.b.Disconnect(context.Background(), conn.EndpointID, conn.ConnID, port.DisconnectNormal) }() } return nil } // Kick 只断开当前连接,令牌不变。 func (s *Session) Kick(ctx context.Context, endpointID string) error { if s.b == nil { return ErrNoConnection } return s.b.Disconnect(ctx, endpointID, "", port.DisconnectKicked) } // Disable 清空令牌,发 fatal(disabled) 后断开。 func (s *Session) Disable(ctx context.Context, endpointID string) error { return s.fatalKick(ctx, endpointID, "disabled") } // Deleted 清空令牌,发 fatal(deleted) 后断开。 func (s *Session) Deleted(ctx context.Context, endpointID string) error { return s.fatalKick(ctx, endpointID, "deleted") } // ResetPassword 清空令牌,发 fatal(password_reset) 后断开。 func (s *Session) ResetPassword(ctx context.Context, endpointID string) error { return s.fatalKick(ctx, endpointID, "password_reset") } func (s *Session) fatalKick(ctx context.Context, endpointID, reason string) error { if s.login != nil { if err := s.login.ClearSession(ctx, endpointID); err != nil { return err } } if s.b == nil { return nil } info, ok := s.b.ConnInfoOf(endpointID) if !ok { return nil } fatal := protocol.Fatal{V: protocol.Version, Type: protocol.TypeFatal, Reason: reason} raw, err := protocol.Marshal(fatal) if err != nil { return err } if pubErr := s.b.PublishThenDisconnect(ctx, info.EndpointID, info.ConnID, raw, 1, port.DisconnectFatal); pubErr != nil { go func() { _ = s.b.Disconnect(context.Background(), endpointID, info.ConnID, port.DisconnectFatal) }() } return nil } func (s *Session) replyErr(ctx context.Context, conn port.ConnInfo, rid, code, message string) { if rid == "" { rid = "0" } resp := protocol.Resp{ V: protocol.Version, Type: protocol.TypeResp, RID: rid, OK: false, Error: &protocol.ErrorBody{Code: code, Message: message}, } _ = s.publishJSON(ctx, conn, resp, 1) } func (s *Session) publishJSON(ctx context.Context, conn port.ConnInfo, v any, qos byte) error { b, err := protocol.Marshal(v) if err != nil { return err } return s.b.PublishDown(ctx, conn.EndpointID, conn.ConnID, b, port.PublishOpts{QoS: qos}) } func peekRID(payload []byte) string { var peek struct { RID string `json:"rid"` } _ = json.Unmarshal(payload, &peek) return peek.RID } var _ port.UplinkHandler = (*Session)(nil)