Files

280 lines
6.3 KiB
Go

package nixmsg
import (
"context"
"encoding/json"
"time"
)
// Connect 连接服务器。credential 为密码或会话令牌。
func (c *Client) Connect(ctx context.Context, rawURL, endpointID string, credential Credential, opts Options) error {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return apiErr(CodeClosed, "已关闭")
}
if c.transport != nil {
c.mu.Unlock()
return apiErr(CodeBadRequest, "已在连接中")
}
o := opts.withDefaults()
c.opts = o
c.endpointID = endpointID
c.url = rawURL
c.stopReconnect = false
c.handshook = false
pass := credential.Password
if credential.SessionToken != "" {
pass = credential.SessionToken
c.credKind = "token"
} else {
c.credKind = "password"
}
var tr transport
if o.transport != nil {
tr = o.transport
} else {
tr = newMQTTTransport()
}
c.transport = tr
c.backoff = newReconnectBackoff()
inner, cancel := context.WithCancel(context.Background())
c.ctx = inner
c.cancel = cancel
c.setStateLocked(StateConnecting, "")
c.mu.Unlock()
go c.downLoop(inner)
tr.SetCredential(pass)
cfg := transportConfig{
URL: rawURL,
EndpointID: endpointID,
ConnectTimeout: o.ConnectTimeout,
AllowTCP: o.AllowTCP,
Backoff: c.backoff,
OnDown: c.handleDown,
OnOffline: func() {
c.mu.Lock()
c.handshook = false
if !c.stopReconnect && !c.closed {
c.setStateLocked(StateReconnecting, "")
}
c.mu.Unlock()
},
OnAuthFailed: func(reason AuthReason) { c.failAuth(reason) },
OnKicked: func() { c.failKicked() },
MQTTReady: func(readyCtx context.Context) error { return c.doHello(readyCtx) },
}
if err := tr.Start(inner, cfg); err != nil {
return err
}
deadline := time.Now().Add(o.ConnectTimeout)
for time.Now().Before(deadline) {
c.mu.Lock()
ok := c.handshook
failed := c.stopReconnect
st := c.state
c.mu.Unlock()
if ok {
return nil
}
if failed || st == StateAuthFailed || st == StateKicked {
return apiErr(CodeAuthFailed, string(st))
}
select {
case <-ctx.Done():
_ = c.Close()
return ctx.Err()
case <-time.After(20 * time.Millisecond):
}
}
_ = c.Close()
return apiErr(CodeNotConnected, "连接超时")
}
func (c *Client) failAuth(reason AuthReason) {
c.mu.Lock()
c.stopReconnect = true
c.handshook = false
c.setStateLocked(StateAuthFailed, string(reason))
c.failQueuedLocked(apiErr(string(reason), "认证失败,停止重连"))
tr := c.transport
cancel := c.cancel
c.mu.Unlock()
if cancel != nil {
cancel()
}
if tr != nil {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
_ = tr.Stop(ctx)
}
}
func (c *Client) failKicked() {
c.mu.Lock()
c.stopReconnect = true
c.handshook = false
c.setStateLocked(StateKicked, "0x8E")
c.failQueuedLocked(apiErr(CodeKicked, "被顶号,停止重连"))
tr := c.transport
cancel := c.cancel
c.mu.Unlock()
if cancel != nil {
cancel()
}
if tr != nil {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
_ = tr.Stop(ctx)
}
}
func (c *Client) setStateLocked(st ConnectionState, reason string) {
c.state = st
h := c.onConnection
ev := ConnectionEvent{State: st, Reason: reason}
go func() {
c.cbMu.Lock()
defer c.cbMu.Unlock()
if h != nil {
h(ev)
}
}()
}
func (c *Client) doHello(ctx context.Context) error {
c.mu.Lock()
c.helloSentAt = time.Now()
sentAt := c.helloSentAt
label := c.opts.ClientLabel
maxRecv := c.opts.MaxReceiveBytes
c.mu.Unlock()
req := map[string]any{
"v": 1,
"type": "hello",
"rid": c.nextRID(),
"client": label,
}
if maxRecv > 0 {
req["max_receive_bytes"] = maxRecv
}
data, err := c.request(ctx, req, true)
if err != nil {
return err
}
recvAt := time.Now()
var hd struct {
ServerTimeMs int64 `json:"server_time_ms"`
ServerVersion string `json:"server_version"`
MaxBodyBytes int `json:"max_body_bytes"`
MaxMetaBytes int `json:"max_meta_bytes"`
MaxFrameBytes int `json:"max_frame_bytes"`
MaxTTLSeconds int64 `json:"max_ttl_seconds"`
MaxScheduleSeconds int64 `json:"max_schedule_seconds"`
AckTimeoutSeconds int64 `json:"ack_timeout_seconds"`
SessionToken string `json:"session_token"`
}
if err := json.Unmarshal(data, &hd); err != nil {
return err
}
mid := (sentAt.UnixMilli() + recvAt.UnixMilli()) / 2
skew := hd.ServerTimeMs - mid
c.mu.Lock()
c.limits = HandshakeLimits{
ServerTimeMs: hd.ServerTimeMs,
ServerVersion: hd.ServerVersion,
MaxBodyBytes: hd.MaxBodyBytes,
MaxMetaBytes: hd.MaxMetaBytes,
MaxFrameBytes: hd.MaxFrameBytes,
MaxTTLSeconds: hd.MaxTTLSeconds,
MaxScheduleSeconds: hd.MaxScheduleSeconds,
AckTimeoutSeconds: hd.AckTimeoutSeconds,
}
c.clockSkew = skew
c.handshook = true
c.setStateLocked(StateOnline, "")
token := hd.SessionToken
if token != "" {
c.session = token
c.transport.SetCredential(token)
c.credKind = "token"
}
c.mu.Unlock()
if token != "" && c.onSession != nil {
c.cbMu.Lock()
c.onSession(token)
c.cbMu.Unlock()
}
c.drainSendQueue()
return nil
}
// ClockSkewMs 当前服务器时间偏差(毫秒)。
func (c *Client) ClockSkewMs() int64 {
c.mu.Lock()
defer c.mu.Unlock()
return c.clockSkew
}
// Limits 握手上限。
func (c *Client) Limits() HandshakeLimits {
c.mu.Lock()
defer c.mu.Unlock()
return c.limits
}
// Logout 作废会话并停止重连。
func (c *Client) Logout(ctx context.Context) error {
req := map[string]any{"v": 1, "type": "self.logout", "rid": c.nextRID()}
_, err := c.request(ctx, req, false)
c.mu.Lock()
c.stopReconnect = true
c.session = ""
cancel := c.cancel
c.mu.Unlock()
if cancel != nil {
cancel()
}
return err
}
// Close 关闭连接并停止重连。
func (c *Client) Close() error {
c.mu.Lock()
c.closed = true
c.stopReconnect = true
c.failQueuedLocked(apiErr(CodeClosed, "已关闭"))
tr := c.transport
cancel := c.cancel
c.setStateLocked(StateOffline, "")
c.mu.Unlock()
if cancel != nil {
cancel()
}
if tr != nil {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
return tr.Stop(ctx)
}
return nil
}
func (c *Client) failQueuedLocked(err error) {
for _, it := range c.sendQ {
if it.result != nil {
select {
case it.result <- sendOutcome{err: err}:
default:
}
}
}
c.sendQ = nil
c.inflight = 0
}