Files
NixMsg/cmd/nixmsg/uplink.go
T

345 lines
9.0 KiB
Go

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
}