369 lines
10 KiB
Go
369 lines
10 KiB
Go
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)
|
||
|
||
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}
|
||
if err := s.publishJSON(ctx, conn, resp, 1); err != nil {
|
||
s.log.Error("logout resp", "endpoint", conn.EndpointID, "err", err)
|
||
}
|
||
go func() {
|
||
// 稍等让 QoS1 resp 写入连接,再断开
|
||
time.Sleep(50 * time.Millisecond)
|
||
_ = 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}
|
||
_ = s.publishJSON(ctx, info, fatal, 1)
|
||
go func() {
|
||
time.Sleep(20 * time.Millisecond)
|
||
_ = 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)
|