feat: 实现登录、会话令牌、握手与顶号

EOF
This commit is contained in:
Nixevol
2026-09-30 07:25:35 +08:00
parent d357082f3d
commit 34e7c2827f
6 changed files with 1590 additions and 14 deletions
+363
View File
@@ -0,0 +1,363 @@
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
}
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)