package main import ( "context" "encoding/json" "errors" "log/slog" "git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/app/message" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/app/presence" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) // appUplink 把 broker 生命周期接到消息连接表,并把已握手上行帧分发到各业务服务。 type appUplink struct { msg *message.App identity *identity.App presence *presence.App groups *group.App conns *message.MemoryConns down port.Downlink log *slog.Logger } func (u *appUplink) OnSessionEstablished(ctx context.Context, conn port.ConnInfo) error { u.conns.Set(conn.EndpointID, message.LiveConn{ ConnID: conn.ConnID, MaxPacketSize: conn.MaxPacketSize, }) return nil } func (u *appUplink) OnHandshakeComplete(ctx context.Context, hs port.HandshakeInfo) error { live := message.LiveConn{ ConnID: hs.ConnID, MaxReceiveBytes: hs.MaxReceiveBytes, MaxPacketSize: hs.MaxPacketSize, } u.conns.Set(hs.EndpointID, live) if err := u.msg.OnHandshakeComplete(ctx, hs.EndpointID, live); err != nil { u.log.Error("message handshake", "endpoint", hs.EndpointID, "err", err) return err } return nil } func (u *appUplink) OnDisconnect(ctx context.Context, conn port.ConnInfo, reason port.DisconnectReason) { if u.presence != nil { u.presence.ClearWatch(conn.ConnID) } live, ok := u.conns.Current(conn.EndpointID) isCurrent := ok && live.ConnID == conn.ConnID if err := u.msg.OnDisconnect(ctx, conn.EndpointID, conn.ConnID, isCurrent); err != nil { u.log.Error("message disconnect", "endpoint", conn.EndpointID, "err", err) } u.conns.Clear(conn.EndpointID, conn.ConnID) } func (u *appUplink) HandleUplink(ctx context.Context, conn port.ConnInfo, payload []byte) error { frame, err := protocol.Decode(payload) if err != nil { u.replyErr(ctx, conn, peekRID(payload), protocol.CodeBadRequest, err.Error()) return nil } rid, data, callErr := u.dispatch(ctx, conn, frame) if callErr != nil { u.replyFromErr(ctx, conn, rid, callErr) return nil } u.replyOK(ctx, conn, rid, data) return nil } func (u *appUplink) dispatch(ctx context.Context, conn port.ConnInfo, frame any) (rid string, data any, err error) { switch f := frame.(type) { case *protocol.Send: rid = f.RID res, e := u.msg.Submit(ctx, conn.EndpointID, conn, f) if e != nil { return rid, nil, e } return rid, protocol.SendData{ID: res.ID, SendAtMs: res.SendAtMs, State: res.State}, nil case *protocol.Ack: rid = f.RID res, e := u.msg.Ack(ctx, conn.EndpointID, f) if e != nil { return rid, nil, e } return rid, map[string]any{"result": res.Result}, nil case *protocol.ReceiptAck: rid = f.RID return rid, nil, u.msg.ReceiptAck(ctx, conn.EndpointID, f) case *protocol.Recall: rid = f.RID res, e := u.msg.Recall(ctx, conn.EndpointID, f) if e != nil { return rid, nil, e } return rid, res, nil case *protocol.Status: rid = f.RID res, e := u.msg.Status(ctx, conn.EndpointID, f) return rid, res, e case *protocol.Unlock: rid = f.RID if e := f.Validate(); e != nil { return rid, nil, e } return rid, nil, u.identity.UnlockTalk(ctx, conn.EndpointID, f.EndpointID, f.TalkPassword, conn.RemoteIP) case *protocol.SelfGet: rid = f.RID info, e := u.identity.SelfGet(ctx, conn.EndpointID) return rid, info, e case *protocol.SelfUpdate: rid = f.RID return rid, nil, u.identity.SelfUpdate(ctx, conn.EndpointID, f) case *protocol.SelfTalkPassword: rid = f.RID if e := f.Validate(); e != nil { return rid, nil, e } return rid, nil, u.identity.SelfSetTalkPassword(ctx, conn.EndpointID, f.TalkPassword) case *protocol.SelfLoginPassword: rid = f.RID if e := f.Validate(); e != nil { return rid, nil, e } tok, e := u.identity.SelfChangeLoginPassword(ctx, conn.EndpointID, f.OldPassword, f.NewPassword, conn.RemoteIP) if e != nil { return rid, nil, e } return rid, map[string]any{"session_token": tok}, nil case *protocol.PresenceGet: rid = f.RID if e := f.Validate(); e != nil { return rid, nil, e } items, e := u.presence.Get(ctx, f.IDs) if e != nil { return rid, nil, e } return rid, presence.EncodeGetData(items), nil case *protocol.DirectoryList: rid = f.RID items, next, e := u.presence.Directory(ctx, f) if e != nil { return rid, nil, e } return rid, map[string]any{"items": items, "next_cursor": next}, nil case *protocol.PresenceWatch: rid = f.RID return rid, nil, u.presence.Watch(ctx, conn.ConnID, conn.EndpointID, f) case *protocol.GroupCreate: rid = f.RID res, e := u.groups.Create(ctx, conn.EndpointID, f) return rid, res, e case *protocol.GroupAdd: rid = f.RID res, e := u.groups.Add(ctx, conn.EndpointID, f) return rid, res, e case *protocol.GroupRemove: rid = f.RID return rid, nil, u.groups.Remove(ctx, conn.EndpointID, f) case *protocol.GroupLeave: rid = f.RID return rid, nil, u.groups.Leave(ctx, conn.EndpointID, f) case *protocol.GroupTransfer: rid = f.RID return rid, nil, u.groups.Transfer(ctx, conn.EndpointID, f) case *protocol.GroupRename: rid = f.RID return rid, nil, u.groups.Rename(ctx, conn.EndpointID, f) case *protocol.GroupDissolve: rid = f.RID return rid, nil, u.groups.Dissolve(ctx, conn.EndpointID, f) case *protocol.GroupList: rid = f.RID items, next, e := u.groups.List(ctx, conn.EndpointID, f) if e != nil { return rid, nil, e } return rid, map[string]any{"items": items, "next_cursor": next}, nil case *protocol.GroupGet: rid = f.RID res, e := u.groups.Get(ctx, conn.EndpointID, f) return rid, res, e case *protocol.Hello, *protocol.SelfLogout: // Session 已处理;不应落到此处。 rid = frameRID(frame) return rid, nil, &protocol.Error{Code: protocol.CodeBadRequest, Message: "unexpected frame"} default: rid = frameRID(frame) return rid, nil, &protocol.Error{Code: protocol.CodeBadRequest, Message: "unsupported uplink type"} } } func (u *appUplink) replyOK(ctx context.Context, conn port.ConnInfo, rid string, data any) { if rid == "" { rid = "0" } var raw json.RawMessage if data != nil { b, err := protocol.Marshal(data) if err != nil { u.replyErr(ctx, conn, rid, protocol.CodeBusy, "marshal resp data") return } raw = b } resp := protocol.Resp{V: protocol.Version, Type: protocol.TypeResp, RID: rid, OK: true, Data: raw} u.publishResp(ctx, conn, resp) } func (u *appUplink) replyFromErr(ctx context.Context, conn port.ConnInfo, rid string, err error) { if rid == "" { rid = "0" } var pe *protocol.Error if errors.As(err, &pe) && pe != nil { u.replyErr(ctx, conn, rid, pe.Code, pe.Message) return } u.log.Error("uplink handler", "endpoint", conn.EndpointID, "err", err) u.replyErr(ctx, conn, rid, protocol.CodeBusy, "internal error") } func (u *appUplink) 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}, } u.publishResp(ctx, conn, resp) } func (u *appUplink) publishResp(ctx context.Context, conn port.ConnInfo, resp protocol.Resp) { if u.down == nil { return } b, err := protocol.Marshal(resp) if err != nil { u.log.Error("marshal resp", "err", err) return } if live, ok := u.conns.Current(conn.EndpointID); ok && live.ConnID == conn.ConnID { limit := respPayloadLimit(live.MaxPacketSize, live.MaxReceiveBytes) if limit > 0 && len(b) > limit { tooLarge := protocol.Resp{ V: protocol.Version, Type: protocol.TypeResp, RID: resp.RID, OK: false, Error: &protocol.ErrorBody{Code: protocol.CodeResponseTooLarge, Message: "response too large"}, } b, err = protocol.Marshal(tooLarge) if err != nil { return } } } if pubErr := u.down.PublishDown(ctx, conn.EndpointID, conn.ConnID, b, port.PublishOpts{QoS: 1}); pubErr != nil { u.log.Error("publish resp", "endpoint", conn.EndpointID, "err", pubErr) } } func respPayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { limit := 0 if maxRecvBytes > 0 { limit = maxRecvBytes } if maxPacketSize > 0 { n := int(maxPacketSize) if limit == 0 || n < limit { limit = n } } return limit } func peekRID(payload []byte) string { var peek struct { RID string `json:"rid"` } _ = json.Unmarshal(payload, &peek) return peek.RID } func frameRID(frame any) string { b, err := protocol.Marshal(frame) if err != nil { return "" } return peekRID(b) } // presenceConnTable 把消息连接表暴露给 presence.ConnTable。 type presenceConnTable struct { conns *message.MemoryConns } func (p *presenceConnTable) IsOnline(endpointID string) bool { _, ok := p.conns.Current(endpointID) return ok } func (p *presenceConnTable) CurrentConn(endpointID string) (port.ConnID, bool) { live, ok := p.conns.Current(endpointID) if !ok { return "", false } return live.ConnID, true }