266 lines
6.0 KiB
Go
266 lines
6.0 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()
|
|
|
|
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), "认证失败,停止重连"))
|
|
cancel := c.cancel
|
|
c.mu.Unlock()
|
|
if cancel != nil {
|
|
cancel()
|
|
}
|
|
}
|
|
|
|
func (c *Client) failKicked() {
|
|
c.mu.Lock()
|
|
c.stopReconnect = true
|
|
c.handshook = false
|
|
c.setStateLocked(StateKicked, "0x8E")
|
|
c.failQueuedLocked(apiErr(CodeKicked, "被顶号,停止重连"))
|
|
cancel := c.cancel
|
|
c.mu.Unlock()
|
|
if cancel != nil {
|
|
cancel()
|
|
}
|
|
}
|
|
|
|
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
|
|
}
|