Files
NixMsg/internal/broker/session.go
T

369 lines
10 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)