239 lines
5.0 KiB
Go
239 lines
5.0 KiB
Go
package nixmsg
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"time"
|
||
)
|
||
|
||
// Send 发送消息;可在未连接时入队,重连后按原内容重交。
|
||
func (c *Client) Send(ctx context.Context, to Target, body Body, opt SendOptions) (SendResult, error) {
|
||
c.mu.Lock()
|
||
limits := c.limits
|
||
skew := c.clockSkew
|
||
handshook := c.handshook
|
||
maxQ := c.opts.SendQueueSize
|
||
qLen := len(c.sendQ)
|
||
c.mu.Unlock()
|
||
|
||
if body.Enc == "" {
|
||
body.Enc = "utf8"
|
||
}
|
||
if body.ContentType == "" && opt.ContentType != "" {
|
||
body.ContentType = opt.ContentType
|
||
} else if body.ContentType == "" {
|
||
if body.Enc == "base64" {
|
||
body.ContentType = "application/octet-stream"
|
||
} else {
|
||
body.ContentType = "text/plain; charset=utf-8"
|
||
}
|
||
} else if opt.ContentType != "" {
|
||
body.ContentType = opt.ContentType
|
||
}
|
||
|
||
n, err := bodyDecodedLen(body)
|
||
if err != nil {
|
||
return SendResult{}, err
|
||
}
|
||
maxBody := limits.MaxBodyBytes
|
||
if maxBody <= 0 {
|
||
maxBody = 262144
|
||
}
|
||
if handshook && n > maxBody {
|
||
return SendResult{}, apiErr(CodeBodyTooLarge, "正文超限")
|
||
}
|
||
if !handshook && n > 262144 {
|
||
return SendResult{}, apiErr(CodeBodyTooLarge, "正文超限")
|
||
}
|
||
|
||
id := opt.ID
|
||
if id == "" {
|
||
id = newMessageID()
|
||
}
|
||
|
||
frame := map[string]any{
|
||
"v": 1,
|
||
"type": "send",
|
||
"rid": c.nextRID(),
|
||
"id": id,
|
||
"to": to,
|
||
"body": body,
|
||
}
|
||
if opt.Meta != nil {
|
||
frame["meta"] = opt.Meta
|
||
}
|
||
if opt.TalkPassword != "" {
|
||
frame["talk_password"] = opt.TalkPassword
|
||
}
|
||
if opt.Receipt != nil {
|
||
frame["receipt"] = *opt.Receipt
|
||
}
|
||
if opt.Keep {
|
||
off := map[string]any{"keep": true}
|
||
if opt.TTL != nil {
|
||
off["ttl_seconds"] = *opt.TTL
|
||
}
|
||
frame["offline"] = off
|
||
}
|
||
|
||
var sendAtMs *int64
|
||
if opt.SendAt != nil && opt.Delay != nil {
|
||
return SendResult{}, apiErr(CodeBadRequest, "sendAt 与 delay 互斥")
|
||
}
|
||
if opt.SendAt != nil {
|
||
// sendAt 使用本机时间 + 服务器偏差,换算后写入,重交不重算
|
||
ms := opt.SendAt.UnixMilli() + skew
|
||
sendAtMs = &ms
|
||
frame["send_at_ms"] = ms
|
||
} else if opt.Delay != nil {
|
||
frame["delay_ms"] = opt.Delay.Milliseconds()
|
||
}
|
||
|
||
payload, err := marshalJSON(frame)
|
||
if err != nil {
|
||
return SendResult{}, err
|
||
}
|
||
maxFrame := limits.MaxFrameBytes
|
||
if maxFrame <= 0 {
|
||
maxFrame = 786432
|
||
}
|
||
if handshook && len(payload) > maxFrame {
|
||
return SendResult{}, apiErr(CodeFrameTooLarge, "整帧超限")
|
||
}
|
||
|
||
item := &sendItem{
|
||
frame: frame,
|
||
payload: payload,
|
||
id: id,
|
||
sendAtMs: sendAtMs,
|
||
result: make(chan sendOutcome, 1),
|
||
}
|
||
|
||
c.mu.Lock()
|
||
if c.closed || c.stopReconnect {
|
||
c.mu.Unlock()
|
||
return SendResult{}, apiErr(CodeClosed, "已关闭")
|
||
}
|
||
if qLen >= maxQ {
|
||
c.mu.Unlock()
|
||
return SendResult{}, apiErr(CodeQueueFull, "发送队列已满")
|
||
}
|
||
c.sendQ = append(c.sendQ, item)
|
||
c.mu.Unlock()
|
||
|
||
c.drainSendQueue()
|
||
|
||
select {
|
||
case <-ctx.Done():
|
||
return SendResult{}, ctx.Err()
|
||
case out := <-item.result:
|
||
return out.res, out.err
|
||
}
|
||
}
|
||
|
||
func (c *Client) drainSendQueue() {
|
||
for {
|
||
c.mu.Lock()
|
||
if !c.handshook || c.transport == nil {
|
||
c.mu.Unlock()
|
||
return
|
||
}
|
||
var next *sendItem
|
||
for _, it := range c.sendQ {
|
||
if !it.inflight {
|
||
next = it
|
||
break
|
||
}
|
||
}
|
||
if next == nil || c.inflight >= c.opts.MaxInflight {
|
||
c.mu.Unlock()
|
||
return
|
||
}
|
||
next.inflight = true
|
||
c.inflight++
|
||
tr := c.transport
|
||
payload := next.payload
|
||
rid, _ := next.frame["rid"].(string)
|
||
item := next
|
||
c.mu.Unlock()
|
||
|
||
go c.dispatchSend(tr, item, rid, payload)
|
||
}
|
||
}
|
||
|
||
func (c *Client) dispatchSend(tr transport, item *sendItem, rid string, payload []byte) {
|
||
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)
|
||
item.inflight = false
|
||
c.inflight--
|
||
c.mu.Unlock()
|
||
// 网络错误:保留队列等重连
|
||
return
|
||
}
|
||
|
||
rf := <-ch
|
||
if !rf.OK {
|
||
code, msg := CodeBadRequest, "发送失败"
|
||
if rf.Error != nil {
|
||
code, msg = rf.Error.Code, rf.Error.Message
|
||
}
|
||
if code == CodeRateLimited {
|
||
// 自动重交:保持同一 payload(含 id / send_at_ms)
|
||
c.mu.Lock()
|
||
item.inflight = false
|
||
c.inflight--
|
||
delete(c.pending, rid)
|
||
c.mu.Unlock()
|
||
time.AfterFunc(time.Second, func() { c.drainSendQueue() })
|
||
return
|
||
}
|
||
c.finishSend(item, SendResult{}, apiErr(code, msg))
|
||
return
|
||
}
|
||
var sd SendResult
|
||
_ = json.Unmarshal(rf.Data, &sd)
|
||
if sd.ID == "" {
|
||
sd.ID = item.id
|
||
}
|
||
c.finishSend(item, sd, nil)
|
||
}
|
||
|
||
func (c *Client) finishSend(item *sendItem, res SendResult, err error) {
|
||
c.mu.Lock()
|
||
// 从队列移除
|
||
out := item.result
|
||
nq := c.sendQ[:0]
|
||
for _, it := range c.sendQ {
|
||
if it != item {
|
||
nq = append(nq, it)
|
||
}
|
||
}
|
||
c.sendQ = nq
|
||
if item.inflight {
|
||
c.inflight--
|
||
item.inflight = false
|
||
}
|
||
c.mu.Unlock()
|
||
select {
|
||
case out <- sendOutcome{res: res, err: err}:
|
||
default:
|
||
}
|
||
c.drainSendQueue()
|
||
}
|
||
|
||
// ResendPayloadForTest 返回队列中第一条发送帧的编码(测试用)。
|
||
func (c *Client) ResendPayloadForTest() []byte {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
if len(c.sendQ) == 0 {
|
||
return nil
|
||
}
|
||
return append([]byte(nil), c.sendQ[0].payload...)
|
||
}
|