Files

355 lines
7.9 KiB
Go

package nixmsg
import (
"context"
"encoding/base64"
"encoding/json"
"time"
)
func (c *Client) handleDown(payload []byte) {
var head struct {
Type string `json:"type"`
RID string `json:"rid"`
}
if err := unmarshalJSON(payload, &head); err != nil {
return
}
// resp 必须在收包路径同步处理,否则 request/ack 在 downLoop 里等待时会死锁。
if head.Type == "resp" {
var rf respFrame
if err := unmarshalJSON(payload, &rf); err != nil {
return
}
c.mu.Lock()
p := c.pending[head.RID]
if p != nil {
delete(c.pending, head.RID)
}
c.mu.Unlock()
if p != nil {
select {
case p.ch <- rf:
default:
}
}
return
}
cp := append([]byte(nil), payload...)
select {
case c.downCh <- cp:
case <-c.ctx.Done():
}
}
func (c *Client) downLoop(ctx context.Context) {
for {
select {
case <-ctx.Done():
return
case payload := <-c.downCh:
c.handleDownApp(payload)
}
}
}
func (c *Client) handleDownApp(payload []byte) {
var head struct {
Type string `json:"type"`
}
if err := unmarshalJSON(payload, &head); err != nil {
return
}
switch head.Type {
case "msg":
c.handleMsg(payload)
case "receipt":
c.handleReceipt(payload)
case "revoked":
c.handleRevoked(payload)
case "presence":
c.handlePresence(payload)
case "group_event":
c.handleGroupEvent(payload)
case "fatal":
var f struct {
Reason string `json:"reason"`
}
_ = unmarshalJSON(payload, &f)
c.mu.Lock()
c.stopReconnect = true
c.setStateLocked(StateAuthFailed, f.Reason)
c.failQueuedLocked(apiErr("fatal", f.Reason))
cancel := c.cancel
c.mu.Unlock()
if cancel != nil {
cancel()
}
}
}
func (c *Client) handleMsg(payload []byte) {
var m struct {
ID string `json:"id"`
From string `json:"from"`
To Target `json:"to"`
Body Body `json:"body"`
Meta map[string]any `json:"meta"`
SendAtMs int64 `json:"send_at_ms"`
}
if err := unmarshalJSON(payload, &m); err != nil {
return
}
key := m.From + "\x00" + m.ID
c.mu.Lock()
ent := c.dedup[key]
manual := c.opts.ManualAck
if ent != nil && ent.state == dedupAcked {
c.mu.Unlock()
_ = c.sendAckFrame(m.From, m.ID)
return
}
if ent != nil && ent.state == dedupDelivered {
c.mu.Unlock()
return
}
c.rememberDedupLocked(key, m.From, m.ID, dedupDelivered)
c.mu.Unlock()
msg := Message{ID: m.ID, From: m.From, To: m.To, Body: m.Body, Meta: m.Meta, SendAtMs: m.SendAtMs}
var cbErr error
c.cbMu.Lock()
if c.onMessage != nil {
cbErr = c.onMessage(msg)
}
c.cbMu.Unlock()
if manual {
return
}
if cbErr != nil {
c.mu.Lock()
delete(c.dedup, key)
c.mu.Unlock()
return
}
_ = c.sendAckFrame(m.From, m.ID)
c.mu.Lock()
if e := c.dedup[key]; e != nil {
e.state = dedupAcked
}
c.mu.Unlock()
}
func (c *Client) rememberDedupLocked(key, from, id string, st dedupState) {
if _, ok := c.dedup[key]; !ok {
c.dedupOrd = append(c.dedupOrd, key)
for len(c.dedupOrd) > c.opts.DedupCapacity {
old := c.dedupOrd[0]
c.dedupOrd = c.dedupOrd[1:]
delete(c.dedup, old)
}
}
c.dedup[key] = &dedupEntry{state: st, from: from, id: id}
}
// Ack 手动确认。
func (c *Client) Ack(msg Message) error {
if err := c.sendAckFrame(msg.From, msg.ID); err != nil {
return err
}
key := msg.From + "\x00" + msg.ID
c.mu.Lock()
c.rememberDedupLocked(key, msg.From, msg.ID, dedupAcked)
c.mu.Unlock()
return nil
}
func (c *Client) sendAckFrame(from, id string) error {
req := map[string]any{
"v": 1, "type": "ack", "rid": c.nextRID(), "from": from, "id": id,
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
data, err := c.request(ctx, req, true)
if err != nil {
return err
}
if len(data) > 0 {
var d struct {
Result string `json:"result"`
}
if json.Unmarshal(data, &d) == nil && d.Result != "" && d.Result != "accepted" {
c.emitRevoked(RevokedEvent{ID: id, From: from, Reason: d.Result})
}
}
return nil
}
func (c *Client) handleReceipt(payload []byte) {
var r struct {
ReceiptID string `json:"receipt_id"`
ID string `json:"id"`
EndpointID string `json:"endpoint_id"`
State string `json:"state"`
Reason string `json:"reason"`
AtMs int64 `json:"at_ms"`
}
if err := unmarshalJSON(payload, &r); err != nil {
return
}
c.mu.Lock()
if _, ok := c.receiptSeen[r.ReceiptID]; ok {
c.mu.Unlock()
_ = c.sendReceiptAck(r.ReceiptID)
return
}
c.receiptSeen[r.ReceiptID] = struct{}{}
c.mu.Unlock()
ev := Receipt{ReceiptID: r.ReceiptID, ID: r.ID, EndpointID: r.EndpointID, State: r.State, Reason: r.Reason, AtMs: r.AtMs}
c.cbMu.Lock()
if c.onReceipt != nil {
c.onReceipt(ev)
}
c.cbMu.Unlock()
_ = c.sendReceiptAck(r.ReceiptID)
}
func (c *Client) sendReceiptAck(receiptID string) error {
req := map[string]any{"v": 1, "type": "receipt_ack", "rid": c.nextRID(), "receipt_id": receiptID}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
_, err := c.request(ctx, req, true)
return err
}
func (c *Client) handleRevoked(payload []byte) {
var r struct {
ID string `json:"id"`
From string `json:"from"`
Reason string `json:"reason"`
}
if err := unmarshalJSON(payload, &r); err != nil {
return
}
key := r.From + "\x00" + r.ID
c.mu.Lock()
ent := c.dedup[key]
if ent == nil || ent.state == dedupAcked {
c.mu.Unlock()
return
}
delete(c.dedup, key)
c.mu.Unlock()
c.emitRevoked(RevokedEvent{ID: r.ID, From: r.From, Reason: r.Reason})
}
func (c *Client) emitRevoked(e RevokedEvent) {
c.cbMu.Lock()
defer c.cbMu.Unlock()
if c.onRevoked != nil {
c.onRevoked(e)
}
}
func (c *Client) handlePresence(payload []byte) {
var p struct {
ID string `json:"id"`
Online bool `json:"online"`
AtMs int64 `json:"at_ms"`
}
if err := unmarshalJSON(payload, &p); err != nil {
return
}
c.cbMu.Lock()
defer c.cbMu.Unlock()
if c.onPresence != nil {
c.onPresence(PresenceEvent{ID: p.ID, Online: p.Online, AtMs: p.AtMs})
}
}
func (c *Client) handleGroupEvent(payload []byte) {
var g struct {
GroupID string `json:"group_id"`
Event string `json:"event"`
EndpointID string `json:"endpoint_id"`
AtMs int64 `json:"at_ms"`
}
if err := unmarshalJSON(payload, &g); err != nil {
return
}
c.cbMu.Lock()
defer c.cbMu.Unlock()
if c.onGroupEvent != nil {
c.onGroupEvent(GroupEvent{GroupID: g.GroupID, Event: g.Event, EndpointID: g.EndpointID, AtMs: g.AtMs})
}
}
func (c *Client) request(ctx context.Context, frame map[string]any, allowUnready bool) (json.RawMessage, error) {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return nil, apiErr(CodeClosed, "已关闭")
}
if !allowUnready && !c.handshook {
c.mu.Unlock()
return nil, apiErr(CodeNotConnected, "未握手")
}
tr := c.transport
c.mu.Unlock()
if tr == nil {
return nil, apiErr(CodeNotConnected, "未连接")
}
rid, _ := frame["rid"].(string)
if rid == "" {
rid = c.nextRID()
frame["rid"] = rid
}
payload, err := marshalJSON(frame)
if err != nil {
return nil, err
}
ch := make(chan respFrame, 1)
c.mu.Lock()
c.pending[rid] = &pendingReq{rid: rid, ch: ch}
c.mu.Unlock()
if err := tr.PublishUp(payload); err != nil {
c.mu.Lock()
delete(c.pending, rid)
c.mu.Unlock()
return nil, err
}
select {
case <-ctx.Done():
c.mu.Lock()
delete(c.pending, rid)
c.mu.Unlock()
return nil, ctx.Err()
case rf := <-ch:
if !rf.OK {
code, msg := CodeBadRequest, "请求失败"
if rf.Error != nil {
code, msg = rf.Error.Code, rf.Error.Message
}
return nil, apiErr(code, msg)
}
return rf.Data, nil
}
}
// bodyDecodedLen 按解码后字节计正文大小。
func bodyDecodedLen(b Body) (int, error) {
switch b.Enc {
case "utf8", "":
return len([]byte(b.Data)), nil
case "base64":
raw, err := base64.StdEncoding.DecodeString(b.Data)
if err != nil {
return 0, apiErr(CodeBadRequest, "base64 正文无效")
}
return len(raw), nil
default:
return 0, apiErr(CodeBadRequest, "未知 enc")
}
}