feat: 实现登录、会话令牌、握手与顶号
EOF
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user