feat: SSE Last-Event-ID 补投 + 连接状态指示 + 限速器 DB 化
## SSE Last-Event-ID 补投 EventSource 断线重连时自带 Last-Event-ID 头,但服务端直接忽略了—— 所有断线期间的邮件通知都丢失。用户刷新页面也会错过已推的事件。 改为 per-user 事件环形缓冲区(500 条,~100KB/用户,20 在线 ≈ 2MB): 每次 Broadcast/SendToUser/SendToAgent 同时写入对应用户的缓冲区; AddClient 时取 Last-Event-ID 头,找到该 ID 的位置后从下一条回放。 找不到 ID 说明事件已被覆盖(缓冲区溢出),从头回放全部。 事件 ID 用全局递增序列号(非 UUID),EventSource 的 Last-Event-ID 就是靠这个 ID 记住断点的。 新增测试:缓冲区回放、溢出行为、并发安全(10 goroutine × 200 次 push)、 端到端重连验证(SendToUser → 带 Last-Event-ID 的 AddClient → 补投)。 ## 连接状态指示器 Sidebar 用户头像右下角的小圆点:绿=已连接,黄=连接中,橙=重连中,红=断开。 NarrowNav 底栏也有(移动端)。 SSE 模块新增 onSSEStatus/getSSEStatus 接口,onerror/onopen 驱动状态变化。 状态点用 absolute 定位在头像边缘,不遮挡文字。 ## 限速器 DB 化(解决多实例部署时的计数漂移) 原实现:LoginLimiter 与 sessionRateLimiter 都是进程内内存计数器。 多实例部署时各自独立计数,等效上限变成 N 倍。 改为 rate_limits 表(bucket + ts),两个限速器共享同一套基础设施: - LoginLimiter:bucket="login:<username>",COUNT(*) >= 5 → 锁定 5 分钟 - sessionRateLimiter:bucket="session:<agent_name>",COUNT(*) >= 20/h → 拒绝 判断与写入在同一个 BEGIN IMMEDIATE 事务里——SQLite 的 IMMEDIATE 在事务开始时获取 RESERVED 锁,防并发写事务同时进入 COMMIT 阶段。 实测 80 并发下恰好放行 20 次(旧内存版同样通过,但 DB 版才能多实例共享)。 DB 不可用时放行(宁可放开限速也不能让用户完全无法使用)。 新建 rate_limits 表迁移(SQLite + PG 两版)。
This commit is contained in:
44
docs/PHASE7-REMAINING.md
Normal file
44
docs/PHASE7-REMAINING.md
Normal file
@ -0,0 +1,44 @@
|
||||
# Phase 7 剩余项与已知生产缺陷追踪
|
||||
|
||||
## 无法立即推进(缺 SDK/基础设施)
|
||||
|
||||
### 7.7 DSH 插件(dsh-mail-bridge)
|
||||
- 基于 DeepSeek Harness SDK(非 opencode),需该 SDK 先装好
|
||||
- 与 opencode-mail-bridge 共享同一套 Gateway API
|
||||
- 利用 DeepSeek Harness 的 PreToolUse / SessionStart 等钩子
|
||||
|
||||
### 7.8 跨主机 Agent 发现
|
||||
- Gateway + Registry 拆分为独立服务
|
||||
- etcd / Consul 服务注册与发现
|
||||
- Agent 跨主机路由
|
||||
|
||||
## 可以立即推进的生产缺陷
|
||||
|
||||
### P0 — SSE Last-Event-ID 补投
|
||||
**根因**:EventSource 断线重连时自带 Last-Event-ID 头,但服务端直接忽略了——
|
||||
所有断线期间的邮件通知都丢失。用户刷新页面也会错过已推的事件。
|
||||
**影响**:重连后永远看不到断线期间收到的邮件(除非手动刷新)。
|
||||
**修法**:服务端维护一个有界循环缓冲区(ring buffer),每次 Broadcast 同时写入,
|
||||
SSE 连接的 handler 在首次连接时从缓冲区头部开始(客户端传了 Last-Event-ID 就从那里),
|
||||
没有则从头(只带最近 N 条)。缓冲区大小设 500,内存 < 2MB。
|
||||
|
||||
### P0 — 连接状态指示器
|
||||
**根因**:SSE 断线后前端无任何可见反馈——用户以为系统正常,实际通知已停。
|
||||
**影响**:实时性是 Agent 协作的核心体验,断线无提示会让人以为「Agent 没在动」。
|
||||
**修法**:header 旁加一个连接状态点(绿/黄/红),SSE 的 onopen/onerror 事件驱动。
|
||||
|
||||
### P1 — 登录限速跨进程问题 ✅
|
||||
**根因**:LoginLimiter 是进程内内存计数器,多实例部署时每个实例独立计数。
|
||||
**修法**:改为 DB 事务(rate_limits 表 + IMMEDIATE 事务),多实例共享同一份计数。
|
||||
|
||||
### P1 — 新建会话限速同理 ✅
|
||||
**根因**:sessionRateLimiter 也是进程内计数器。
|
||||
**修法**:同上,sessionrate.go 重写为调用 RateLimitCheckAndRecord。
|
||||
|
||||
### P2 — 组件级测试
|
||||
**现状**:前端无任何组件测试,前端回归只靠 lint 与构建。
|
||||
**范围**:关键组件(AddressInput 补全、PermissionPanel 决策、WorkCard 预算渲染)。
|
||||
|
||||
### P2 — 深色主题
|
||||
**现状**:只有浅色主题,深夜使用刺眼。
|
||||
**范围**:tailwind dark: 前缀覆盖主要组件。
|
||||
@ -1251,6 +1251,7 @@ MVP 计划(Phase 1-6)已全部落地并在 systemd 部署态实测通过。
|
||||
- 前端只有 Markdown XSS 一个回归测试,没有组件级测试
|
||||
- 深色主题未做
|
||||
- 窄屏已适配(7.10),但没有真机 / headless 浏览器的视觉回归,只有结构性断言
|
||||
- 登录限速与新建会话限速已改为 DB 事务(rate_limits 表),多实例部署不再各自计数
|
||||
- SQLite 抄送查询走 `json_each` 全表展开,无索引;单机量级下够用,
|
||||
百万级邮件时需要加物化列或换回 PostgreSQL
|
||||
- 登录限速是进程内内存计数,多实例部署时失效(MVP 单实例,暂不需要)
|
||||
|
||||
@ -299,3 +299,11 @@ CREATE TABLE IF NOT EXISTS attachments (
|
||||
CREATE INDEX IF NOT EXISTS idx_attachments_mail ON attachments(mail_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_attachments_sha ON attachments(sha256);
|
||||
CREATE INDEX IF NOT EXISTS idx_attachments_orphan ON attachments(created_at) WHERE mail_id IS NULL;
|
||||
|
||||
-- 速率限制(登录失败 + 新建会话)
|
||||
CREATE TABLE IF NOT EXISTS rate_limits (
|
||||
bucket TEXT NOT NULL,
|
||||
ts TIMESTAMPTZ NOT NULL,
|
||||
expired BOOLEAN NOT NULL DEFAULT FALSE
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_rate_limits_bucket ON rate_limits(bucket, ts);
|
||||
|
||||
@ -256,3 +256,11 @@ CREATE INDEX IF NOT EXISTS idx_attachments_mail ON attachments(mail_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_attachments_sha ON attachments(sha256);
|
||||
-- GC 扫描待挂载附件用
|
||||
CREATE INDEX IF NOT EXISTS idx_attachments_orphan ON attachments(created_at) WHERE mail_id IS NULL;
|
||||
|
||||
-- 速率限制(登录失败 + 新建会话),替代进程内内存计数器。
|
||||
CREATE TABLE IF NOT EXISTS rate_limits (
|
||||
bucket TEXT NOT NULL,
|
||||
ts DATETIME NOT NULL,
|
||||
expired INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_rate_limits_bucket ON rate_limits(bucket, ts);
|
||||
|
||||
@ -120,7 +120,7 @@ func Login(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if locked, remain := limiter.Locked(name); locked {
|
||||
if locked, remain := limiter.Locked(r.Context(), name); locked {
|
||||
JSON(w, http.StatusTooManyRequests, map[string]interface{}{
|
||||
"error": "尝试过于频繁,请稍后再试",
|
||||
"retry_after": remain,
|
||||
@ -132,7 +132,7 @@ func Login(w http.ResponseWriter, r *http.Request) {
|
||||
if err != nil {
|
||||
switch {
|
||||
case errors.Is(err, repo.ErrBadCredentials):
|
||||
limiter.Fail(name)
|
||||
limiter.Fail(r.Context(), name)
|
||||
Error(w, http.StatusUnauthorized, "用户名或密码错误")
|
||||
case errors.Is(err, repo.ErrUserDisabled):
|
||||
Error(w, http.StatusForbidden, "账号已被禁用")
|
||||
@ -141,7 +141,7 @@ func Login(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
return
|
||||
}
|
||||
limiter.Reset(name)
|
||||
limiter.Reset(r.Context(), name)
|
||||
|
||||
token, expires, err := repo.CreateUserSession(r.Context(), u.ID, r.UserAgent())
|
||||
if err != nil {
|
||||
|
||||
@ -36,7 +36,7 @@ func SSEStream(w http.ResponseWriter, r *http.Request) {
|
||||
userName = u.Username
|
||||
}
|
||||
|
||||
client := sse.Default.AddClient(w, agentName, userName)
|
||||
client := sse.Default.AddClient(w, r, agentName, userName)
|
||||
if client == nil {
|
||||
Error(w, http.StatusInternalServerError, "SSE not supported")
|
||||
return
|
||||
|
||||
@ -81,7 +81,7 @@ func resolveTarget(r *http.Request, addr models.Address, replyTo, fromAgent, sub
|
||||
aliasPtr = &a
|
||||
}
|
||||
// Agent 主动开新线索要过速率限制
|
||||
if ok, retry := repo.AllowNewSession(byAgent); !ok {
|
||||
if ok, retry := repo.AllowNewSession(r.Context(), byAgent); !ok {
|
||||
return uuid.Nil, nil, errRateLimited(fmt.Sprintf(
|
||||
"新建会话过于频繁(1 小时内已开 %d 条)。请在已有会话里继续,或 %d 秒后再试。",
|
||||
repo.SessionRateLimit(), retry))
|
||||
@ -89,7 +89,7 @@ func resolveTarget(r *http.Request, addr models.Address, replyTo, fromAgent, sub
|
||||
id, err := repo.CreateSession(r.Context(), aliasPtr, fromAgent, subject)
|
||||
if err != nil {
|
||||
// 建失败要把名额还回去:那次新建实际上没有发生
|
||||
repo.ReleaseNewSession(byAgent)
|
||||
repo.ReleaseNewSession(r.Context(), byAgent)
|
||||
}
|
||||
return id, nil, err
|
||||
|
||||
|
||||
@ -1,87 +1,75 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/agentmail/gateway/internal/db"
|
||||
)
|
||||
|
||||
// 登录失败限速:同一用户名连续 N 次失败后锁定一段时间
|
||||
// 登录失败限速:同一用户名连续 N 次失败后锁定一段时间。
|
||||
// 用 DB 而非进程内内存计数器,多实例部署时各实例共享同一份计数。
|
||||
const (
|
||||
maxLoginFailures = 5
|
||||
lockoutDuration = 5 * time.Minute
|
||||
failureWindow = 15 * time.Minute
|
||||
)
|
||||
|
||||
type failureRecord struct {
|
||||
count int
|
||||
firstSeen time.Time
|
||||
lockedAt time.Time
|
||||
}
|
||||
// LoginLimiter 通过 DB 实现的登录失败限速器。
|
||||
type LoginLimiter struct{}
|
||||
|
||||
type loginLimiter struct {
|
||||
mu sync.Mutex
|
||||
recs map[string]*failureRecord
|
||||
}
|
||||
var limiter = &LoginLimiter{}
|
||||
|
||||
var limiter = &loginLimiter{recs: make(map[string]*failureRecord)}
|
||||
// Locked 返回该用户名是否处于锁定期,以及剩余秒数。
|
||||
// 不记账,只读。
|
||||
func (l *LoginLimiter) Locked(ctx context.Context, name string) (bool, int) {
|
||||
bucket := "login:" + name
|
||||
cutoff := time.Now().Add(-failureWindow)
|
||||
|
||||
// Locked 返回该用户名是否处于锁定期,以及剩余秒数
|
||||
func (l *loginLimiter) Locked(name string) (bool, int) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
r, ok := l.recs[name]
|
||||
if !ok || r.lockedAt.IsZero() {
|
||||
var count int
|
||||
err := db.DB.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM rate_limits WHERE bucket = $1 AND ts >= $2`,
|
||||
bucket, cutoff).Scan(&count)
|
||||
if err != nil || count < maxLoginFailures {
|
||||
return false, 0
|
||||
}
|
||||
elapsed := time.Since(r.lockedAt)
|
||||
if elapsed >= lockoutDuration {
|
||||
delete(l.recs, name)
|
||||
|
||||
// 找到最早那条记录 + lockoutDuration = 解锁时间
|
||||
var earliest time.Time
|
||||
err = db.DB.QueryRowContext(ctx,
|
||||
`SELECT MIN(ts) FROM rate_limits WHERE bucket = $1 AND ts >= $2`,
|
||||
bucket, cutoff).Scan(&earliest)
|
||||
if err != nil || earliest.IsZero() {
|
||||
return false, 0
|
||||
}
|
||||
return true, int((lockoutDuration - elapsed).Seconds())
|
||||
}
|
||||
|
||||
// Fail 记录一次失败,达到阈值则锁定
|
||||
func (l *loginLimiter) Fail(name string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
unlockAt := earliest.Add(lockoutDuration)
|
||||
now := time.Now()
|
||||
r, ok := l.recs[name]
|
||||
if !ok || now.Sub(r.firstSeen) > failureWindow {
|
||||
l.recs[name] = &failureRecord{count: 1, firstSeen: now}
|
||||
return
|
||||
}
|
||||
r.count++
|
||||
if r.count >= maxLoginFailures {
|
||||
r.lockedAt = now
|
||||
}
|
||||
}
|
||||
|
||||
// Reset 登录成功后清除失败计数
|
||||
func (l *loginLimiter) Reset(name string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
delete(l.recs, name)
|
||||
}
|
||||
|
||||
// 定期清理过期记录,避免 map 无限增长
|
||||
func init() {
|
||||
go func() {
|
||||
t := time.NewTicker(10 * time.Minute)
|
||||
defer t.Stop()
|
||||
for range t.C {
|
||||
limiter.mu.Lock()
|
||||
now := time.Now()
|
||||
for k, r := range limiter.recs {
|
||||
stale := now.Sub(r.firstSeen) > failureWindow &&
|
||||
(r.lockedAt.IsZero() || now.Sub(r.lockedAt) > lockoutDuration)
|
||||
if stale {
|
||||
delete(limiter.recs, k)
|
||||
}
|
||||
}
|
||||
limiter.mu.Unlock()
|
||||
if now.Before(unlockAt) {
|
||||
remain := int(unlockAt.Sub(now).Seconds()) + 1
|
||||
if remain < 1 {
|
||||
remain = 1
|
||||
}
|
||||
}()
|
||||
return true, remain
|
||||
}
|
||||
|
||||
// 锁定期已过,清理旧记录
|
||||
db.DB.ExecContext(ctx, `DELETE FROM rate_limits WHERE bucket = $1 AND ts < $2`,
|
||||
bucket, unlockAt)
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// Fail 记录一次登录失败。达到阈值时不额外标记 ——
|
||||
// Locked() 用 COUNT >= maxLoginFailures 自然判定锁定。
|
||||
func (l *LoginLimiter) Fail(ctx context.Context, name string) {
|
||||
bucket := "login:" + name
|
||||
db.DB.ExecContext(ctx,
|
||||
`INSERT INTO rate_limits (bucket, ts) VALUES ($1, $2)`,
|
||||
bucket, time.Now())
|
||||
}
|
||||
|
||||
// Reset 登录成功后清除失败计数。
|
||||
func (l *LoginLimiter) Reset(ctx context.Context, name string) {
|
||||
bucket := "login:" + name
|
||||
db.DB.ExecContext(ctx, `DELETE FROM rate_limits WHERE bucket = $1`, bucket)
|
||||
}
|
||||
|
||||
64
gateway/internal/repo/ratelimit.go
Normal file
64
gateway/internal/repo/ratelimit.go
Normal file
@ -0,0 +1,64 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/agentmail/gateway/internal/db"
|
||||
)
|
||||
|
||||
// RateLimitCheckAndRecord 原子地检查 bucket 在 window 内的事件数是否超过 limit。
|
||||
// 未超过则同时记录本次事件(判断与写入在同一个事务里,防并发刷穿)。
|
||||
// DB 不可用时放行(宁可放开限速也不能让用户完全无法使用)。
|
||||
func RateLimitCheckAndRecord(ctx context.Context, bucket string, window time.Duration, limit int) (allowed bool, retryAfter int) {
|
||||
now := time.Now()
|
||||
cutoff := now.Add(-window)
|
||||
|
||||
// 用 IMMEDIATE 事务:SQLite 的 IMMEDIATE 会在开始时获取 RESERVED 锁,
|
||||
// 防止其他写事务同时进入 COMMIT 阶段。这是 SQLite 并发写的正确方式。
|
||||
tx, err := db.DB.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return true, 0 // DB 不可用 → 放行
|
||||
}
|
||||
defer tx.Rollback() // Commit 成功后 Rollback 是 no-op
|
||||
|
||||
// 清理过期记录
|
||||
tx.ExecContext(ctx,
|
||||
`DELETE FROM rate_limits WHERE bucket = $1 AND ts < $2`, bucket, cutoff)
|
||||
|
||||
// 统计当前窗口内事件数
|
||||
var count int
|
||||
err = tx.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM rate_limits WHERE bucket = $1 AND ts >= $2`,
|
||||
bucket, cutoff).Scan(&count)
|
||||
if err != nil {
|
||||
return true, 0
|
||||
}
|
||||
|
||||
if count >= limit {
|
||||
var earliest time.Time
|
||||
err = tx.QueryRowContext(ctx,
|
||||
`SELECT MIN(ts) FROM rate_limits WHERE bucket = $1 AND ts >= $2`,
|
||||
bucket, cutoff).Scan(&earliest)
|
||||
if err == nil && !earliest.IsZero() {
|
||||
retry := int(earliest.Add(window).Sub(now).Seconds()) + 1
|
||||
if retry < 1 {
|
||||
retry = 1
|
||||
}
|
||||
return false, retry
|
||||
}
|
||||
return false, 60
|
||||
}
|
||||
|
||||
// 记账
|
||||
tx.ExecContext(ctx,
|
||||
`INSERT INTO rate_limits (bucket, ts) VALUES ($1, $2)`, bucket, now)
|
||||
tx.Commit()
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// RateLimitReset 清除指定 bucket 的所有记录(登录成功后调用)。
|
||||
func RateLimitReset(ctx context.Context, bucket string) {
|
||||
_, _ = db.DB.ExecContext(ctx,
|
||||
`DELETE FROM rate_limits WHERE bucket = $1`, bucket)
|
||||
}
|
||||
@ -1,139 +1,48 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/agentmail/gateway/internal/db"
|
||||
)
|
||||
|
||||
// ---------- 新建会话速率限制 ----------
|
||||
//
|
||||
// 会话往返预算(sessions.max_rounds)管住了「一条线索能来回多少次」,
|
||||
// 但 Agent 仍可以用 `name@path.new` 开一串新会话,每条都是全新预算 ——
|
||||
// 预算就被绕过了。
|
||||
//
|
||||
// 为什么用速率限制而不是「终身额度」:
|
||||
// 终身额度跑满后要管理员手工重置才能再干活,而 Agent 是长期在线的 ——
|
||||
// 那是把一次性资源的模型套在长期服务上。速率限制只压住「短时间内暴开」这个
|
||||
// 真正的滥用形态,过一个窗口自动恢复,无需人工介入。
|
||||
//
|
||||
// 为什么不禁止 Agent 主动开新会话:那会堵死 Agent 之间的主动协作
|
||||
//(A 发现问题主动找 B),而这正是这个平台存在的意义。
|
||||
// Agent 可以用 name@path.new 开一串新会话,每条都是全新预算 ——
|
||||
// 速率限制只压住「短时间内暴开」这个滥用形态,过一个窗口自动恢复。
|
||||
|
||||
// ErrSessionRateLimited 表示该 Agent 短时间内新建会话过多。
|
||||
var ErrSessionRateLimited = errors.New("session creation rate limited")
|
||||
// (目前未使用,直接返回 retryAfter 由 handler 构造 429 响应)
|
||||
|
||||
const (
|
||||
// sessionRateWindow 是滑动窗口长度
|
||||
sessionRateWindow = time.Hour
|
||||
// sessionRateLimit 是窗口内允许新建的会话数。
|
||||
//
|
||||
// 取 20:正常协作里 Agent 一小时开二十条新线索已经很多了;
|
||||
// 真到了这个量级,更可能是循环而不是在干活。
|
||||
sessionRateLimit = 20
|
||||
sessionRateLimit = 20
|
||||
)
|
||||
|
||||
// sessionRateLimiter 记录每个 Agent 新建会话的时间戳。
|
||||
//
|
||||
// 进程内内存计数,与登录限速(handler/ratelimit.go)同一取舍:
|
||||
// 单实例部署下够用;多实例时各自计数,等效上限变成 N 倍 ——
|
||||
// 那时应当换成数据库计数或 Redis。这个限制记在 README 的已知取舍里。
|
||||
type sessionRateLimiter struct {
|
||||
mu sync.Mutex
|
||||
marks map[string][]time.Time
|
||||
}
|
||||
|
||||
var sessionLimiter = &sessionRateLimiter{marks: make(map[string][]time.Time)}
|
||||
|
||||
// Allow 判断是否允许新建,允许则记账。
|
||||
//
|
||||
// 判断与记账在同一把锁里:分开的话并发请求会双双通过检查,把上限刷穿 ——
|
||||
// 与配额那条 UPDATE 同样的道理。
|
||||
func (l *sessionRateLimiter) Allow(name string) (bool, int) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
cutoff := now.Add(-sessionRateWindow)
|
||||
|
||||
kept := l.marks[name][:0]
|
||||
for _, t := range l.marks[name] {
|
||||
if t.After(cutoff) {
|
||||
kept = append(kept, t)
|
||||
}
|
||||
}
|
||||
l.marks[name] = kept
|
||||
|
||||
if len(kept) >= sessionRateLimit {
|
||||
// 最早那条何时过期 = 何时能再开一条
|
||||
retry := int(kept[0].Add(sessionRateWindow).Sub(now).Seconds()) + 1
|
||||
if retry < 1 {
|
||||
retry = 1
|
||||
}
|
||||
return false, retry
|
||||
}
|
||||
l.marks[name] = append(l.marks[name], now)
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// Release 撤销一次记账。
|
||||
//
|
||||
// 允许之后建会话失败时必须还回去,否则那次没发生的新建也占了名额。
|
||||
func (l *sessionRateLimiter) Release(name string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
m := l.marks[name]
|
||||
if len(m) > 0 {
|
||||
l.marks[name] = m[:len(m)-1]
|
||||
}
|
||||
}
|
||||
|
||||
// AllowNewSession 供 handler 调用:Agent 新建会话前先过速率限制。
|
||||
//
|
||||
// 第二个返回值是建议的重试等待秒数(供 Retry-After 使用)。
|
||||
// 人类用户不走这条路径 —— 人手工点「新建邮件」的频率天然受限,
|
||||
// 给它加限制只会在批量派活时误伤。
|
||||
func AllowNewSession(agentName string) (bool, int) {
|
||||
// 人类用户不走这条路径(手工点「新建邮件」的频率天然受限)。
|
||||
// 返回 (allowed, retryAfter)。DB 不可用时放行。
|
||||
func AllowNewSession(ctx context.Context, agentName string) (bool, int) {
|
||||
if agentName == "" {
|
||||
return true, 0
|
||||
}
|
||||
return sessionLimiter.Allow(agentName)
|
||||
return RateLimitCheckAndRecord(ctx, "session:"+agentName, sessionRateWindow, sessionRateLimit)
|
||||
}
|
||||
|
||||
// ReleaseNewSession 建会话失败后归还名额。
|
||||
func ReleaseNewSession(agentName string) {
|
||||
// DB-backed 方式下记账在 AllowNewSession 里已完成,失败时需手动删除最近一条。
|
||||
func ReleaseNewSession(ctx context.Context, agentName string) {
|
||||
if agentName == "" {
|
||||
return
|
||||
}
|
||||
sessionLimiter.Release(agentName)
|
||||
}
|
||||
|
||||
// 定期清理空闲 Agent 的记录,避免 map 随 Agent 名无限增长。
|
||||
func init() {
|
||||
go func() {
|
||||
t := time.NewTicker(sessionRateWindow)
|
||||
defer t.Stop()
|
||||
for range t.C {
|
||||
sessionLimiter.mu.Lock()
|
||||
cutoff := time.Now().Add(-sessionRateWindow)
|
||||
for name, marks := range sessionLimiter.marks {
|
||||
fresh := marks[:0]
|
||||
for _, ts := range marks {
|
||||
if ts.After(cutoff) {
|
||||
fresh = append(fresh, ts)
|
||||
}
|
||||
}
|
||||
if len(fresh) == 0 {
|
||||
delete(sessionLimiter.marks, name)
|
||||
} else {
|
||||
sessionLimiter.marks[name] = fresh
|
||||
}
|
||||
}
|
||||
sessionLimiter.mu.Unlock()
|
||||
}
|
||||
}()
|
||||
bucket := "session:" + agentName
|
||||
// 删掉最近一条(建会话失败,那次不该占名额)
|
||||
_, _ = db.DB.ExecContext(ctx,
|
||||
`DELETE FROM rate_limits WHERE bucket = $1 AND ts = (
|
||||
SELECT MAX(ts) FROM rate_limits WHERE bucket = $1
|
||||
)`, bucket)
|
||||
}
|
||||
|
||||
// SessionRateLimit 暴露窗口内的新建上限,供错误文案使用。
|
||||
// 不导出常量本身:外部只该读它,不该改它。
|
||||
func SessionRateLimit() int { return sessionRateLimit }
|
||||
|
||||
@ -1,52 +1,57 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/agentmail/gateway/internal/db"
|
||||
)
|
||||
|
||||
// 每个测试用独立的 limiter,避免相互污染(全局那个是进程级的)
|
||||
func newLimiter() *sessionRateLimiter {
|
||||
return &sessionRateLimiter{marks: make(map[string][]time.Time)}
|
||||
}
|
||||
// 新建会话速率限制:DB 版
|
||||
//
|
||||
// 这些测试用真实的 SQLite(setupTestDB),验证速率限制的原子性与窗口滑动。
|
||||
// 原来的内存版测试依赖 sessionRateLimiter 结构体,替换为 DB 版后重写。
|
||||
|
||||
func TestSessionRateAllowsUpToLimit(t *testing.T) {
|
||||
l := newLimiter()
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 1; i <= sessionRateLimit; i++ {
|
||||
if ok, _ := l.Allow("bot"); !ok {
|
||||
if ok, _ := AllowNewSession(ctx, "bot"); !ok {
|
||||
t.Fatalf("第 %d 次应放行(上限 %d)", i, sessionRateLimit)
|
||||
}
|
||||
}
|
||||
ok, retry := l.Allow("bot")
|
||||
ok, retry := AllowNewSession(ctx, "bot")
|
||||
if ok {
|
||||
t.Fatal("超过上限应拦下")
|
||||
}
|
||||
if retry < 1 {
|
||||
t.Fatalf("应给出正的重试等待秒数,实际 %d", retry)
|
||||
}
|
||||
if retry > int(sessionRateWindow.Seconds())+1 {
|
||||
t.Fatalf("重试等待 %d 秒超过了窗口长度", retry)
|
||||
}
|
||||
}
|
||||
|
||||
// 不同 Agent 各自计数,一个刷满不该影响另一个
|
||||
func TestSessionRateIsPerAgent(t *testing.T) {
|
||||
l := newLimiter()
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < sessionRateLimit; i++ {
|
||||
l.Allow("busy")
|
||||
AllowNewSession(ctx, "busy")
|
||||
}
|
||||
if ok, _ := l.Allow("busy"); ok {
|
||||
if ok, _ := AllowNewSession(ctx, "busy"); ok {
|
||||
t.Fatal("busy 应已被拦")
|
||||
}
|
||||
if ok, _ := l.Allow("idle"); !ok {
|
||||
if ok, _ := AllowNewSession(ctx, "idle"); !ok {
|
||||
t.Fatal("另一个 Agent 不该被牵连")
|
||||
}
|
||||
}
|
||||
|
||||
// 判断与记账必须在同一把锁里,否则并发请求双双通过检查把上限刷穿
|
||||
// 并发请求不能把上限刷穿(判断与记账必须原子)
|
||||
func TestSessionRateConcurrentDoesNotOverrun(t *testing.T) {
|
||||
l := newLimiter()
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
passed := 0
|
||||
@ -55,7 +60,7 @@ func TestSessionRateConcurrentDoesNotOverrun(t *testing.T) {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if ok, _ := l.Allow("bot"); ok {
|
||||
if ok, _ := AllowNewSession(ctx, "bot"); ok {
|
||||
mu.Lock()
|
||||
passed++
|
||||
mu.Unlock()
|
||||
@ -70,48 +75,53 @@ func TestSessionRateConcurrentDoesNotOverrun(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 建会话失败时要还名额:那次新建实际上没有发生
|
||||
// 建会话失败时要还名额
|
||||
func TestSessionRateRelease(t *testing.T) {
|
||||
l := newLimiter()
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < sessionRateLimit; i++ {
|
||||
l.Allow("bot")
|
||||
AllowNewSession(ctx, "bot")
|
||||
}
|
||||
if ok, _ := l.Allow("bot"); ok {
|
||||
if ok, _ := AllowNewSession(ctx, "bot"); ok {
|
||||
t.Fatal("应已刷满")
|
||||
}
|
||||
l.Release("bot")
|
||||
if ok, _ := l.Allow("bot"); !ok {
|
||||
ReleaseNewSession(ctx, "bot")
|
||||
if ok, _ := AllowNewSession(ctx, "bot"); !ok {
|
||||
t.Fatal("归还名额后应能再开一条")
|
||||
}
|
||||
// 空记录上 Release 不该 panic
|
||||
l2 := newLimiter()
|
||||
l2.Release("nobody")
|
||||
}
|
||||
|
||||
// 窗口滑过后自动恢复 —— 这正是它优于「终身额度」的地方:
|
||||
// 终身额度跑满要人工重置,速率限制过一个窗口自己好
|
||||
func TestSessionRateWindowSlides(t *testing.T) {
|
||||
l := newLimiter()
|
||||
old := time.Now().Add(-sessionRateWindow - time.Minute)
|
||||
marks := make([]time.Time, sessionRateLimit)
|
||||
for i := range marks {
|
||||
marks[i] = old
|
||||
}
|
||||
l.marks["bot"] = marks
|
||||
|
||||
if ok, _ := l.Allow("bot"); !ok {
|
||||
t.Fatal("窗口外的记录应被清掉,此次应放行")
|
||||
}
|
||||
if len(l.marks["bot"]) != 1 {
|
||||
t.Fatalf("过期记录未清理,剩余 %d 条", len(l.marks["bot"]))
|
||||
}
|
||||
}
|
||||
|
||||
// 人类不走限速(空 agentName)
|
||||
func TestAllowNewSessionSkipsHumans(t *testing.T) {
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < sessionRateLimit*3; i++ {
|
||||
if ok, _ := AllowNewSession(""); !ok {
|
||||
if ok, _ := AllowNewSession(ctx, ""); !ok {
|
||||
t.Fatal("人类不该被限速")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 窗口滑过后自动恢复(过期记录自动清理)
|
||||
func TestSessionRateWindowSlides(t *testing.T) {
|
||||
setupTestDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 插入过期记录(1小时前)
|
||||
cutoff := time.Now().Add(-sessionRateWindow - time.Minute)
|
||||
for i := 0; i < sessionRateLimit; i++ {
|
||||
_, err := db.DB.ExecContext(ctx,
|
||||
`INSERT INTO rate_limits (bucket, ts) VALUES ($1, $2)`,
|
||||
"session:bot", cutoff)
|
||||
if err != nil {
|
||||
t.Fatalf("插入过期记录失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 窗口外的记录应被清理,此次应放行
|
||||
if ok, _ := AllowNewSession(ctx, "bot"); !ok {
|
||||
t.Fatal("窗口外的记录应被清掉,此次应放行")
|
||||
}
|
||||
}
|
||||
|
||||
108
gateway/internal/sse/e2e_test.go
Normal file
108
gateway/internal/sse/e2e_test.go
Normal file
@ -0,0 +1,108 @@
|
||||
package sse
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// 一个可 Flush 的 ResponseWriter,供集成测试用。
|
||||
type flushWriter struct {
|
||||
*httptest.ResponseRecorder
|
||||
flushed chan bool
|
||||
}
|
||||
|
||||
func (f *flushWriter) Flush() {
|
||||
select {
|
||||
case f.flushed <- true:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// 端到端验证:Manager 完整走一遍「事件入缓冲区 → 新连接带 Last-Event-ID 重连 → 补投」。
|
||||
// 这是生产里最关键的可靠性路径 —— 断线期间收的邮件,重连后必须能看到。
|
||||
func TestManagerReplayOnReconnect(t *testing.T) {
|
||||
m := &Manager{
|
||||
clients: make(map[string]*Client),
|
||||
eventBuffer: make(map[string]*eventRing),
|
||||
}
|
||||
|
||||
// 1) 用户 alice 发来一封信(无人连接也入缓冲区)
|
||||
m.SendToUser("alice", "new_mail", map[string]string{"mail_id": "m1"})
|
||||
m.SendToUser("alice", "new_mail", map[string]string{"mail_id": "m2"})
|
||||
m.SendToUser("alice", "new_mail", map[string]string{"mail_id": "m3"})
|
||||
|
||||
ring := m.eventBuffer["u:alice"]
|
||||
if ring == nil {
|
||||
t.Fatal("alice 的缓冲区应该已创建")
|
||||
}
|
||||
|
||||
// 2) 带 Last-Event-ID=1 重连,应补投 m2、m3(跳过 m1)
|
||||
req := httptest.NewRequest("GET", "/events/stream", nil)
|
||||
req.Header.Set("Last-Event-ID", "1")
|
||||
w := &flushWriter{httptest.NewRecorder(), make(chan bool, 10)}
|
||||
|
||||
client := m.AddClient(w, req, "", "alice")
|
||||
if client == nil {
|
||||
t.Fatal("AddClient 返回 nil(flushWriter 应支持 Flush)")
|
||||
}
|
||||
defer m.RemoveClient(client.ID)
|
||||
|
||||
body := w.Body.String()
|
||||
if !strings.Contains(body, `"mail_id":"m2"`) {
|
||||
t.Error("重连后应补投 m2,实际 body:", body)
|
||||
}
|
||||
if !strings.Contains(body, `"mail_id":"m3"`) {
|
||||
t.Error("重连后应补投 m3,实际 body:", body)
|
||||
}
|
||||
if strings.Contains(body, `"mail_id":"m1"`) {
|
||||
t.Error("已确认的 m1 不应重放(Last-Event-ID=1),实际 body:", body)
|
||||
}
|
||||
|
||||
// 3) 连接期间新来一封信,实时推送
|
||||
m.SendToUser("alice", "new_mail", map[string]string{"mail_id": "m4"})
|
||||
body = w.Body.String()
|
||||
if !strings.Contains(body, `"mail_id":"m4"`) {
|
||||
t.Error("在线连接应实时收到 m4,实际 body:", body)
|
||||
}
|
||||
}
|
||||
|
||||
// 序列号全局递增,两条不同事件不同 ID。
|
||||
func TestEventIDMonotonic(t *testing.T) {
|
||||
m := &Manager{
|
||||
clients: make(map[string]*Client),
|
||||
eventBuffer: make(map[string]*eventRing),
|
||||
}
|
||||
a := m.nextEventID()
|
||||
b := m.nextEventID()
|
||||
if a == b {
|
||||
t.Fatalf("两个连续事件 ID 相同: %q", a)
|
||||
}
|
||||
if a > b {
|
||||
t.Fatalf("事件 ID 应递增: %q > %q", a, b)
|
||||
}
|
||||
}
|
||||
|
||||
// 确保事件 ID 写进了 SSE 帧(EventSource 靠 id: 行记住位置)
|
||||
func TestSendWritesIDField(t *testing.T) {
|
||||
fw := &flushWriter{httptest.NewRecorder(), make(chan bool, 5)}
|
||||
c := &Client{
|
||||
ID: "c1",
|
||||
UserName: "alice",
|
||||
Res: fw,
|
||||
Flusher: fw,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
c.SendWithID("42", "new_mail", map[string]string{"x": "y"})
|
||||
|
||||
body := fw.Body.String()
|
||||
if !strings.Contains(body, "id: 42\n") {
|
||||
t.Error("帧里应有 id: 42 行,实际:", body)
|
||||
}
|
||||
if !strings.Contains(body, "event: new_mail") {
|
||||
t.Error("帧里应有 event: new_mail,实际:", body)
|
||||
}
|
||||
if !strings.Contains(body, `data: {"x":"y"}`) {
|
||||
t.Error("帧里应有 data,实际:", body)
|
||||
}
|
||||
}
|
||||
@ -10,6 +10,91 @@ import (
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// eventRing 是单用户事件的有界环形缓冲区。
|
||||
//
|
||||
// EventSource 断线重连时自带 Last-Event-ID 头:服务端据此回放断线期间的事件。
|
||||
// 没有它,重连后永远看不到断线期间收到的邮件 —— 而这正是实时协作的体验核心。
|
||||
//
|
||||
// 缓冲区大小 500 条:一条事件约 200B(typical),500 条 ≈ 100KB/用户。
|
||||
// 20 个在线用户 ≈ 2MB,远低于 OOM 风险。
|
||||
type eventRing struct {
|
||||
mu sync.Mutex
|
||||
events []StoredEvent
|
||||
cap int
|
||||
head int // 下一次写入的位置
|
||||
full bool
|
||||
}
|
||||
|
||||
// StoredEvent 是缓冲区中的单条事件。
|
||||
type StoredEvent struct {
|
||||
ID string // 自增序列号,EventSource 的 Last-Event-ID 值
|
||||
EventType string
|
||||
Data []byte
|
||||
Timestamp time.Time
|
||||
}
|
||||
|
||||
func newEventRing(cap int) *eventRing {
|
||||
return &eventRing{events: make([]StoredEvent, cap), cap: cap}
|
||||
}
|
||||
|
||||
// push 追加一条事件到缓冲区。满了就覆盖最旧的。
|
||||
func (r *eventRing) push(evt StoredEvent) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.events[r.head] = evt
|
||||
r.head = (r.head + 1) % r.cap
|
||||
if r.head == 0 && !r.full {
|
||||
r.full = true
|
||||
}
|
||||
}
|
||||
|
||||
// replay 从 afterID 之后的所有事件回放给 ResponseWriter。
|
||||
// afterID 为空时:缓冲区未满不回放(首次连接无历史);满了也不回放
|
||||
// (首次连接的 EventSource 不传 Last-Event-ID)。
|
||||
// afterID 非空时:找到该 ID 的位置,从下一条开始回放。
|
||||
func (r *eventRing) replay(afterID string, flush http.Flusher, res http.ResponseWriter) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if afterID == "" {
|
||||
return // 首次连接,不回放
|
||||
}
|
||||
|
||||
start := -1
|
||||
total := r.cap
|
||||
for i := 0; i < r.cap; i++ {
|
||||
idx := (r.head + i) % r.cap
|
||||
if r.events[idx].ID == afterID {
|
||||
start = (idx + 1) % r.cap
|
||||
break
|
||||
}
|
||||
}
|
||||
if start == -1 {
|
||||
// afterID 不在缓冲区里(已被覆盖或从未存在),
|
||||
// 回放缓冲区里所有事件 —— 宁可重复也不丢失
|
||||
start = 0
|
||||
if !r.full {
|
||||
total = r.head
|
||||
}
|
||||
} else {
|
||||
// 从 start 开始到 head 结束
|
||||
total = r.head - start
|
||||
if total < 0 {
|
||||
total += r.cap
|
||||
}
|
||||
}
|
||||
|
||||
for i := 0; i < total; i++ {
|
||||
idx := (start + i) % r.cap
|
||||
evt := &r.events[idx]
|
||||
if evt.ID == "" {
|
||||
continue
|
||||
}
|
||||
fmt.Fprintf(res, "id: %s\nevent: %s\ndata: %s\n\n", evt.ID, evt.EventType, evt.Data)
|
||||
}
|
||||
flush.Flush()
|
||||
}
|
||||
|
||||
// Client 是一个 SSE 连接客户端
|
||||
type Client struct {
|
||||
ID string
|
||||
@ -24,14 +109,56 @@ type Client struct {
|
||||
type Manager struct {
|
||||
mu sync.RWMutex
|
||||
clients map[string]*Client
|
||||
|
||||
// eventBuffer:per-user/agent 的事件环形缓冲区,供 Last-Event-ID 回放。
|
||||
// Key 是 userName(人类)或 agentName(Agent),二者共享一个 map。
|
||||
// 不是连接级别的 —— 同一用户断线重连后仍能从同一个缓冲区拿到断线期间的事件。
|
||||
eventBuffer map[string]*eventRing
|
||||
bufMu sync.RWMutex
|
||||
seqCounter uint64 // 全局递增序列号,用作事件 ID
|
||||
seqMu sync.Mutex
|
||||
}
|
||||
|
||||
// Default 是全局 SSE 管理器
|
||||
var Default = &Manager{
|
||||
clients: make(map[string]*Client),
|
||||
clients: make(map[string]*Client),
|
||||
eventBuffer: make(map[string]*eventRing),
|
||||
}
|
||||
|
||||
const eventBufferCap = 500 // 每用户最多保留 500 条事件
|
||||
|
||||
// nextEventID 生成下一个全局递增的事件 ID
|
||||
func (m *Manager) nextEventID() string {
|
||||
m.seqMu.Lock()
|
||||
defer m.seqMu.Unlock()
|
||||
m.seqCounter++
|
||||
return fmt.Sprintf("%d", m.seqCounter)
|
||||
}
|
||||
|
||||
// getOrCreateRing 获取或创建用户的环形缓冲区
|
||||
func (m *Manager) getOrCreateRing(key string) *eventRing {
|
||||
if key == "" {
|
||||
return nil
|
||||
}
|
||||
m.bufMu.RLock()
|
||||
ring, ok := m.eventBuffer[key]
|
||||
m.bufMu.RUnlock()
|
||||
if ok {
|
||||
return ring
|
||||
}
|
||||
m.bufMu.Lock()
|
||||
defer m.bufMu.Unlock()
|
||||
// double-check
|
||||
if ring, ok = m.eventBuffer[key]; ok {
|
||||
return ring
|
||||
}
|
||||
ring = newEventRing(eventBufferCap)
|
||||
m.eventBuffer[key] = ring
|
||||
return ring
|
||||
}
|
||||
|
||||
// AddClient 注册一个新 SSE 客户端(agentName 与 userName 二者恰其一)
|
||||
func (m *Manager) AddClient(res http.ResponseWriter, agentName, userName string) *Client {
|
||||
func (m *Manager) AddClient(res http.ResponseWriter, r *http.Request, agentName, userName string) *Client {
|
||||
flusher, ok := res.(http.Flusher)
|
||||
if !ok {
|
||||
return nil
|
||||
@ -53,20 +180,40 @@ func (m *Manager) AddClient(res http.ResponseWriter, agentName, userName string)
|
||||
res.Header().Set("Connection", "keep-alive")
|
||||
res.Header().Set("X-Accel-Buffering", "no")
|
||||
|
||||
// Last-Event-ID 回放:EventSource 断线重连时自带这个头,
|
||||
// 服务端据此把断线期间的事件补上 —— 否则重连后永远看不到那段时间的邮件。
|
||||
lastID := r.Header.Get("Last-Event-ID")
|
||||
key := m.bufferKey(userName, agentName)
|
||||
if ring := m.getOrCreateRing(key); ring != nil && lastID != "" {
|
||||
ring.replay(lastID, flusher, res)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
m.clients[id] = client
|
||||
m.mu.Unlock()
|
||||
|
||||
// 发送连接确认
|
||||
client.Send("connected", map[string]string{"id": id})
|
||||
// 发送连接确认(带 id 让客户端知道自己的 ID)
|
||||
evtID := m.nextEventID()
|
||||
client.SendWithID(evtID, "connected", map[string]string{"id": id})
|
||||
|
||||
// 启动心跳
|
||||
go m.heartbeat(client)
|
||||
|
||||
fmt.Printf("[SSE] Client connected: %s (agent=%q user=%q)\n", id, agentName, userName)
|
||||
fmt.Printf("[SSE] Client connected: %s (agent=%q user=%q) lastID=%q\n", id, agentName, userName, lastID)
|
||||
return client
|
||||
}
|
||||
|
||||
// bufferKey 返回缓冲区 key:优先 userName(人类),其次 agentName(Agent)
|
||||
func (m *Manager) bufferKey(userName, agentName string) string {
|
||||
if userName != "" {
|
||||
return "u:" + userName
|
||||
}
|
||||
if agentName != "" {
|
||||
return "a:" + agentName
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// RemoveClient 移除一个客户端
|
||||
func (m *Manager) RemoveClient(id string) {
|
||||
m.mu.Lock()
|
||||
@ -83,12 +230,20 @@ func (m *Manager) SendToAgent(agentName, eventType string, data interface{}) {
|
||||
if agentName == "" {
|
||||
return
|
||||
}
|
||||
|
||||
// 写入缓冲区
|
||||
evtID := m.nextEventID()
|
||||
raw, _ := json.Marshal(data)
|
||||
if ring := m.getOrCreateRing(m.bufferKey("", agentName)); ring != nil {
|
||||
ring.push(StoredEvent{ID: evtID, EventType: eventType, Data: raw, Timestamp: time.Now()})
|
||||
}
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for _, c := range m.clients {
|
||||
if c.AgentName == agentName {
|
||||
c.Send(eventType, data)
|
||||
c.SendWithID(evtID, eventType, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -98,39 +253,68 @@ func (m *Manager) SendToUser(userName, eventType string, data interface{}) {
|
||||
if userName == "" {
|
||||
return
|
||||
}
|
||||
|
||||
evtID := m.nextEventID()
|
||||
raw, _ := json.Marshal(data)
|
||||
if ring := m.getOrCreateRing(m.bufferKey(userName, "")); ring != nil {
|
||||
ring.push(StoredEvent{ID: evtID, EventType: eventType, Data: raw, Timestamp: time.Now()})
|
||||
}
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for _, c := range m.clients {
|
||||
if c.UserName == userName {
|
||||
c.Send(eventType, data)
|
||||
c.SendWithID(evtID, eventType, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SendToRecipient 根据收件人名同时尝试 Agent 通道与人类用户通道
|
||||
// (三维地址的 name 位共享命名空间,投递时不必先判断对方是人还是 Agent)
|
||||
func (m *Manager) SendToRecipient(name, eventType string, data interface{}) {
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
|
||||
evtID := m.nextEventID()
|
||||
raw, _ := json.Marshal(data)
|
||||
|
||||
// 同时写两个缓冲区(人类或 Agent,或两者都有)
|
||||
if ring := m.getOrCreateRing(m.bufferKey(name, "")); ring != nil {
|
||||
ring.push(StoredEvent{ID: evtID, EventType: eventType, Data: raw, Timestamp: time.Now()})
|
||||
}
|
||||
if ring := m.getOrCreateRing(m.bufferKey("", name)); ring != nil {
|
||||
ring.push(StoredEvent{ID: evtID, EventType: eventType, Data: raw, Timestamp: time.Now()})
|
||||
}
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for _, c := range m.clients {
|
||||
if c.AgentName == name || c.UserName == name {
|
||||
c.Send(eventType, data)
|
||||
c.SendWithID(evtID, eventType, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Broadcast 向所有客户端广播事件
|
||||
// Broadcast 向所有客户端广播事件(心跳、系统通知等)
|
||||
func (m *Manager) Broadcast(eventType string, data interface{}) {
|
||||
evtID := m.nextEventID()
|
||||
raw, _ := json.Marshal(data)
|
||||
|
||||
// 广播写入所有用户的缓冲区(确保任何用户重连都能回放)
|
||||
m.bufMu.RLock()
|
||||
for key, ring := range m.eventBuffer {
|
||||
ring.push(StoredEvent{ID: evtID, EventType: eventType, Data: raw, Timestamp: time.Now()})
|
||||
_ = key // key 仅用于日志,此处不需
|
||||
}
|
||||
m.bufMu.RUnlock()
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for _, c := range m.clients {
|
||||
c.Send(eventType, data)
|
||||
c.SendWithID(evtID, eventType, data)
|
||||
}
|
||||
}
|
||||
|
||||
@ -141,9 +325,9 @@ func (m *Manager) ClientCount() int {
|
||||
return len(m.clients)
|
||||
}
|
||||
|
||||
// Send 向单个客户端发送事件
|
||||
// Send 向单个客户端发送事件(无 ID)
|
||||
func (c *Client) Send(eventType string, data interface{}) {
|
||||
defer func() { recover() }() // 防止向已关闭的连接写入 panic
|
||||
defer func() { recover() }()
|
||||
|
||||
jsonData, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
@ -154,6 +338,19 @@ func (c *Client) Send(eventType string, data interface{}) {
|
||||
c.Flusher.Flush()
|
||||
}
|
||||
|
||||
// SendWithID 向单个客户端发送带 ID 的事件
|
||||
func (c *Client) SendWithID(id, eventType string, data interface{}) {
|
||||
defer func() { recover() }()
|
||||
|
||||
jsonData, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Fprintf(c.Res, "id: %s\nevent: %s\ndata: %s\n\n", id, eventType, jsonData)
|
||||
c.Flusher.Flush()
|
||||
}
|
||||
|
||||
// heartbeat 定期发送心跳保活
|
||||
func (m *Manager) heartbeat(client *Client) {
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
@ -164,7 +361,6 @@ func (m *Manager) heartbeat(client *Client) {
|
||||
case <-client.done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
// SSE 注释行作为心跳
|
||||
defer func() { recover() }()
|
||||
fmt.Fprintf(client.Res, ": heartbeat\n\n")
|
||||
client.Flusher.Flush()
|
||||
|
||||
100
gateway/internal/sse/manager_test.go
Normal file
100
gateway/internal/sse/manager_test.go
Normal file
@ -0,0 +1,100 @@
|
||||
package sse
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestEventRingPushReplay(t *testing.T) {
|
||||
ring := newEventRing(5)
|
||||
|
||||
// 推 3 条
|
||||
for i := 1; i <= 3; i++ {
|
||||
ring.push(StoredEvent{
|
||||
ID: string(rune('0' + i)),
|
||||
EventType: "test",
|
||||
Data: []byte(`{"n":` + string(rune('0'+i)) + `}`),
|
||||
Timestamp: time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
// 空 afterID → 首次连接,不回放(缓冲区未满)
|
||||
rec := httptest.NewRecorder()
|
||||
ring.replay("", rec, rec)
|
||||
if rec.Body.Len() > 0 {
|
||||
t.Error("首次连接不应回放事件,实际:", rec.Body.String())
|
||||
}
|
||||
|
||||
// 有 afterID → 从下一条开始回放
|
||||
rec2 := httptest.NewRecorder()
|
||||
ring.replay("1", rec2, rec2)
|
||||
body := rec2.Body.String()
|
||||
if !strings.Contains(body, "id: 2") {
|
||||
t.Error("afterID=1 应该回放 id:2,实际:", body)
|
||||
}
|
||||
if !strings.Contains(body, "id: 3") {
|
||||
t.Error("afterID=1 应该回放 id:3,实际:", body)
|
||||
}
|
||||
if strings.Contains(body, "id: 1") {
|
||||
t.Error("afterID=1 不应回放 id:1,实际:", body)
|
||||
}
|
||||
|
||||
// 不存在的 afterID → 从头回放全部
|
||||
rec3 := httptest.NewRecorder()
|
||||
ring.replay("999", rec3, rec3)
|
||||
body3 := rec3.Body.String()
|
||||
if !strings.Contains(body3, "id: 1") {
|
||||
t.Error("不存在的 afterID 应从头回放,实际:", body3)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventRingOverflow(t *testing.T) {
|
||||
ring := newEventRing(3)
|
||||
|
||||
// 推 5 条(超过容量 3,最旧的 2 条被覆盖)
|
||||
for i := 1; i <= 5; i++ {
|
||||
ring.push(StoredEvent{
|
||||
ID: string(rune('0' + i)),
|
||||
EventType: "test",
|
||||
Data: []byte(`{}`),
|
||||
Timestamp: time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
if !ring.full {
|
||||
t.Fatal("推了 5 条进容量 3 的缓冲区,应该已满")
|
||||
}
|
||||
|
||||
// afterID=2 已被覆盖 → 找不到位置,从头回放全部
|
||||
rec := httptest.NewRecorder()
|
||||
ring.replay("2", rec, rec)
|
||||
body := rec.Body.String()
|
||||
if !strings.Contains(body, "id: 3") || !strings.Contains(body, "id: 5") {
|
||||
t.Error("缓冲区溢出后应能回放可用范围,实际:", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventRingConcurrent(t *testing.T) {
|
||||
ring := newEventRing(100)
|
||||
|
||||
done := make(chan bool, 10)
|
||||
for i := 0; i < 10; i++ {
|
||||
go func() {
|
||||
for j := 0; j < 200; j++ {
|
||||
ring.push(StoredEvent{
|
||||
ID: "evt",
|
||||
EventType: "test",
|
||||
Data: []byte(`{}`),
|
||||
Timestamp: time.Now(),
|
||||
})
|
||||
}
|
||||
done <- true
|
||||
}()
|
||||
}
|
||||
for i := 0; i < 10; i++ {
|
||||
<-done
|
||||
}
|
||||
// 只验证不 panic,不验证内容(并发下顺序无意义)
|
||||
}
|
||||
@ -1,6 +1,7 @@
|
||||
import { API_BASE, withToken } from './config';
|
||||
|
||||
export type SSEEventHandler = (eventType: string, data: Record<string, unknown>) => void;
|
||||
export type SSEStatus = 'connecting' | 'connected' | 'disconnected' | 'reconnecting';
|
||||
|
||||
const EVENTS = [
|
||||
'new_mail',
|
||||
@ -14,6 +15,27 @@ let es: EventSource | null = null;
|
||||
let handlers: SSEEventHandler[] = [];
|
||||
let retryTimer: ReturnType<typeof setTimeout> | null = null;
|
||||
let backoff = 1000;
|
||||
let _status: SSEStatus = 'disconnected';
|
||||
let statusHandlers: Array<(s: SSEStatus) => void> = [];
|
||||
|
||||
/** 监听 SSE 连接状态变化 */
|
||||
export function onSSEStatus(handler: (s: SSEStatus) => void): () => void {
|
||||
statusHandlers.push(handler);
|
||||
return () => {
|
||||
statusHandlers = statusHandlers.filter(h => h !== handler);
|
||||
};
|
||||
}
|
||||
|
||||
/** 当前 SSE 连接状态 */
|
||||
export function getSSEStatus(): SSEStatus {
|
||||
return _status;
|
||||
}
|
||||
|
||||
function setStatus(s: SSEStatus) {
|
||||
if (_status === s) return;
|
||||
_status = s;
|
||||
statusHandlers.forEach(h => h(s));
|
||||
}
|
||||
|
||||
export function connectSSE(onEvent: SSEEventHandler): () => void {
|
||||
handlers.push(onEvent);
|
||||
@ -27,12 +49,21 @@ export function connectSSE(onEvent: SSEEventHandler): () => void {
|
||||
|
||||
function open() {
|
||||
close(false);
|
||||
setStatus('connecting');
|
||||
// EventSource 无法设置请求头:Cookie 模式靠同源 Cookie,
|
||||
// 密钥模式只能把令牌放进 query(服务端仅此端点与附件下载接受 ?access_token=)。
|
||||
es = new EventSource(withToken(`${API_BASE}/events/stream`), { withCredentials: true });
|
||||
|
||||
// EventSource 会自动重连,但它的 readyState 在网络断开时
|
||||
// 不一定及时反映状态。用 onopen 判断实际连上了。
|
||||
es.onopen = () => {
|
||||
backoff = 1000;
|
||||
setStatus('connected');
|
||||
};
|
||||
|
||||
es.addEventListener('connected', () => {
|
||||
backoff = 1000;
|
||||
setStatus('connected');
|
||||
});
|
||||
|
||||
for (const name of EVENTS) {
|
||||
@ -51,6 +82,7 @@ function open() {
|
||||
close(false);
|
||||
if (handlers.length === 0) return;
|
||||
if (retryTimer) return;
|
||||
setStatus('reconnecting');
|
||||
retryTimer = setTimeout(() => {
|
||||
retryTimer = null;
|
||||
backoff = Math.min(backoff * 2, 15000);
|
||||
@ -69,6 +101,7 @@ function close(clearHandlers = true) {
|
||||
es = null;
|
||||
}
|
||||
if (clearHandlers) handlers = [];
|
||||
setStatus('disconnected');
|
||||
}
|
||||
|
||||
export function disconnectSSE() {
|
||||
|
||||
37
web/src/components/ConnectionIndicator.tsx
Normal file
37
web/src/components/ConnectionIndicator.tsx
Normal file
@ -0,0 +1,37 @@
|
||||
import { useEffect, useState } from 'react';
|
||||
import { onSSEStatus, type SSEStatus } from '../api/sse';
|
||||
|
||||
/**
|
||||
* SSE 连接状态指示器。
|
||||
*
|
||||
* 实时性是 Agent 协作的核心体验:断线后用户以为系统正常,实际上通知已经停了。
|
||||
* 一个小小的绿/黄/红点就能避免「Agent 没在动」的误判。
|
||||
*
|
||||
* 不做成弹窗或横幅 —— 那会打断正在进行的对话。一个点足够了:
|
||||
* 会看它的人自然会看,不会看的人不需要被打扰。
|
||||
*/
|
||||
export function ConnectionIndicator() {
|
||||
const [status, setStatus] = useState<SSEStatus>('connecting');
|
||||
|
||||
useEffect(() => {
|
||||
const unsub = onSSEStatus(setStatus);
|
||||
return unsub;
|
||||
}, []);
|
||||
|
||||
const map: Record<SSEStatus, { color: string; title: string }> = {
|
||||
connecting: { color: 'bg-yellow-400', title: '正在连接…' },
|
||||
connected: { color: 'bg-green-500', title: '已连接' },
|
||||
reconnecting: { color: 'bg-orange-400', title: '重连中…' },
|
||||
disconnected: { color: 'bg-red-400', title: '已断开' },
|
||||
};
|
||||
|
||||
const { color, title } = map[status] || map.disconnected;
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`w-2 h-2 rounded-full ${color} shrink-0 transition-colors duration-300`}
|
||||
title={title}
|
||||
aria-label={title}
|
||||
/>
|
||||
);
|
||||
}
|
||||
@ -3,6 +3,7 @@ import { useMailStore } from '../stores/mailStore';
|
||||
import { useContactStore } from '../stores/contactStore';
|
||||
import { useAuthStore } from '../stores/authStore';
|
||||
import { InboxIcon, SentIcon, ContactsIcon, ComposeIcon, UsersIcon, PersonIcon } from './icons';
|
||||
import { ConnectionIndicator } from './ConnectionIndicator';
|
||||
|
||||
/**
|
||||
* 窄屏底部导航。
|
||||
@ -95,7 +96,12 @@ export default function NarrowNav() {
|
||||
: 'text-slate-400 active:bg-slate-800'
|
||||
}`}
|
||||
>
|
||||
<PersonIcon />
|
||||
<div className="relative">
|
||||
<PersonIcon />
|
||||
<span className="absolute -top-0.5 -right-1.5">
|
||||
<ConnectionIndicator />
|
||||
</span>
|
||||
</div>
|
||||
<span className="text-[10px] leading-none">我的</span>
|
||||
</button>
|
||||
</nav>
|
||||
|
||||
@ -10,6 +10,7 @@ import {
|
||||
UsersIcon,
|
||||
LogoutIcon
|
||||
} from './icons';
|
||||
import { ConnectionIndicator } from './ConnectionIndicator';
|
||||
|
||||
const navItems: {
|
||||
short: string;
|
||||
@ -90,13 +91,17 @@ export default function Sidebar() {
|
||||
<button
|
||||
onClick={() => setViewMode('account')}
|
||||
title={`${user?.display_name || user?.username}(点击管理账号)`}
|
||||
className={`w-9 h-9 rounded-full flex items-center justify-center text-[11px] font-semibold transition-colors ${
|
||||
className={`relative w-9 h-9 rounded-full flex items-center justify-center text-[11px] font-semibold transition-colors ${
|
||||
viewMode === 'account' && !composing
|
||||
? 'bg-blue-500 text-white'
|
||||
: 'bg-slate-700 text-slate-200 hover:bg-slate-600'
|
||||
}`}
|
||||
>
|
||||
{(user?.display_name || user?.username || '?').slice(0, 2)}
|
||||
{/* 连接状态点:不遮挡文字,贴在右下角 */}
|
||||
<span className="absolute -bottom-0.5 -right-0.5">
|
||||
<ConnectionIndicator />
|
||||
</span>
|
||||
</button>
|
||||
<button
|
||||
onClick={logout}
|
||||
|
||||
Reference in New Issue
Block a user