feat: 集成限流器到服务 #178
232
CLAUDE.md
232
CLAUDE.md
@@ -6,7 +6,25 @@
|
|||||||
|
|
||||||
CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。
|
CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。
|
||||||
|
|
||||||
> **文档优先原则:** 执行任何开发任务前,先读取 `docs/` 下的相关设计文档(架构、接口、技术选型等),以文档为最高依据。代码实现应与文档一致;若有偏差,优先更新文档(尤其是接口文档)。
|
> **文档优先原则:** 执行任何开发任务前,先读取 `docs/` 下的相关设计文档(架构、接口、技术选型等),以文档为最高依据。代码实现应与文档一致;若有偏差,优先更新文档(尤其是接口文档)。`docs/Eino/` 下有完整的 Eino 框架文档(~75 个 markdown 文件),可作为参考。
|
||||||
|
|
||||||
|
**核心设计文档:**
|
||||||
|
|
||||||
|
| 文档 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| `docs/01-架构设计.md` | 三层架构、技术栈、数据库设计、部署方案 |
|
||||||
|
| `docs/02-接口文档.md` | WebSocket 协议、REST API、AI 服务层、编排器、配置管理 |
|
||||||
|
| `docs/03-技术选型.md` | AI 服务栈、持久化层、前端边缘处理选型 |
|
||||||
|
| `docs/04-用户故事.md` | 用户场景与优先级 |
|
||||||
|
| `docs/05-语音交互.md` | VAD → STT → LLM → TTS 全链路 |
|
||||||
|
| `docs/06-视觉理解.md` | 帧采样、关键帧检测、多模态输入 |
|
||||||
|
| `docs/07-成本控制.md` | 采样策略、端云协同、模型分级 |
|
||||||
|
| `docs/08-功能创意.md` | 功能创意与规划 |
|
||||||
|
| `docs/09-技术名词解释.md` | 术语定义(VAD/STT/TTS/Token/JWT 等) |
|
||||||
|
| `docs/10-Eino重构方案.md` | Eino Graph 迁移方案与决策记录 |
|
||||||
|
| `docs/11-Eino框架技术文档.md` | Eino 框架使用指南 |
|
||||||
|
| `docs/12-鉴权体系设计.md` | JWT 双 token 轮转详细设计 |
|
||||||
|
| `docs/13-令牌桶限流设计.md` | 令牌桶限流详细设计(需同步实现) |
|
||||||
|
|
||||||
## 架构
|
## 架构
|
||||||
|
|
||||||
@@ -16,30 +34,65 @@ CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头
|
|||||||
2. **Go 网关**(Gin, gorilla/websocket, Viper, Zap)—— WebSocket 服务器、会话管理、AI 编排(基于 CloudWeGo Eino Graph)。每个 WebSocket 连接一个 goroutine。
|
2. **Go 网关**(Gin, gorilla/websocket, Viper, Zap)—— WebSocket 服务器、会话管理、AI 编排(基于 CloudWeGo Eino Graph)。每个 WebSocket 连接一个 goroutine。
|
||||||
3. **云端 AI 服务** —— 通过 OpenAI 兼容接口可灵活切换。默认:DashScope qwen3-vl-plus(LLM)、MiMo ASR(STT)、MiMo TTS(TTS)。仅通过 Go 网关访问,浏览器不直连。
|
3. **云端 AI 服务** —— 通过 OpenAI 兼容接口可灵活切换。默认:DashScope qwen3-vl-plus(LLM)、MiMo ASR(STT)、MiMo TTS(TTS)。仅通过 Go 网关访问,浏览器不直连。
|
||||||
|
|
||||||
**关键模式**:AI 编排基于 Eino Graph 声明式 DAG(`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`),LLM token 通过 Callback 实时推送,TTS 逐句合成并行推送,最小化感知延迟。
|
**关键模式**:AI 编排基于 Eino Graph 声明式 DAG(6 节点线性流水线:`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`),LLM token 通过 Callback 实时推送,TTS 逐句合成并行推送,最小化感知延迟。
|
||||||
|
|
||||||
**存储**:三级存储架构(TieredManager)—— L1 Memory → L2 Redis → L3 PostgreSQL,自动降级。Repository 接口模式(UserRepository、MessageRepository、SessionRepository),PostgreSQL + 内存双实现。
|
### Eino Graph 节点详解
|
||||||
|
|
||||||
|
| 节点 | 类型 | 文件 | 职责 |
|
||||||
|
|------|------|------|------|
|
||||||
|
| STT | `InvokableLambda` | `backend/internal/eino/nodes_stt.go` | 语音识别或文本直通(text-only 跳过 STT) |
|
||||||
|
| History | `InvokableLambda` | `backend/internal/eino/nodes_history.go` | 构建 System Prompt + 对话历史 + 用户输入 + 图像 |
|
||||||
|
| ChatModel | ChatModel 节点 | `backend/internal/eino/graph.go` | 调用 DashScope qwen3-vl-plus(OpenAI 兼容协议) |
|
||||||
|
| Msg2Str | `TransformableLambda` | `backend/internal/eino/nodes_splitter.go` | 将 ChatModel 流式 Message 转为字符串流 |
|
||||||
|
| Splitter | `TransformableLambda` | `backend/internal/eino/nodes_splitter.go` | 按句子分隔符(`。!?\n.!?`)拆分文本流 |
|
||||||
|
| TTS | `TransformableLambda` | `backend/internal/eino/nodes_tts.go` | 逐句合成语音并推送 `tts_audio` |
|
||||||
|
| Done | `InvokableLambda` | `backend/internal/eino/nodes_done.go` | 发送 `llm_done`、收集最终输出 |
|
||||||
|
|
||||||
|
**跨节点状态**:`PipelineState`(`backend/internal/eino/state.go`),通过 `context.WithValue` 在节点间传递 FullResponse、TranscribedText、TokenUsage、SessionID、RequestID。
|
||||||
|
|
||||||
|
**Callback**:`BuildCallbackHandler`(`backend/internal/eino/callback.go`)挂载到 ChatModel 的 `OnEndWithStreamOutput`,每收到一个 LLM token 立即通过 `sender.SendLLMChunk()` 推送到客户端。
|
||||||
|
|
||||||
|
**适配器**:`EinoOrchestrator`(`backend/internal/eino/adapter.go`)包装 Graph,实现 `orchestrator.Orchestrator` 接口,负责解码 Base64 图像/音频、构建输入、注入上下文、运行流式推理、持久化消息。
|
||||||
|
|
||||||
|
### 会话存储(TieredManager)
|
||||||
|
|
||||||
|
三级存储:**L1 Memory → L2 Redis → L3 PostgreSQL**(`backend/internal/session/tiered.go`)
|
||||||
|
|
||||||
|
- **读路径**:L1 命中直接返回;未命中尝试 L2 Redis → 回填 L1;L3 通过 L1 的 `FindByID` 降级读取
|
||||||
|
- **写路径**:L1 同步写入 → L2 同步写(失败 soft-warn)→ L3 异步 goroutine 写(使用 `context.Background()` 防止请求取消丢失)
|
||||||
|
- **降级**:后台协程每 30 秒 ping Redis,Redis 不可用时自动跳过 L2 操作;恢复后自动重新启用
|
||||||
|
- **TTL**:Session 默认 30 分钟,MaxHistory 20 条;L1 后台协程每分钟清理过期 session
|
||||||
|
|
||||||
|
Repository 接口模式:`UserRepository`、`MessageRepository`、`SessionRepository`,均有 PostgreSQL 和内存双实现。
|
||||||
|
|
||||||
|
### 鉴权
|
||||||
|
|
||||||
|
JWT 双 token 轮转认证(HMAC-SHA256):
|
||||||
|
- Access Token:默认 120 分钟 TTL,Bearer header 传递
|
||||||
|
- Refresh Token:默认 7 天 TTL,带 jti(UUID),Hash 存储在 Redis/PostgreSQL
|
||||||
|
- 轮转:Refresh 时旧 token hash 删除,新 pair 生成;若 JWT 有效但 DB hash 缺失 → 判定为重放攻击 → 吊销该用户所有 refresh token
|
||||||
|
- `CachedUserRepository`(`backend/internal/store/cached_user.go`):装饰器模式,Redis 缓存 refresh token hash,Read-Through / Write-Through,Redis 故障软降级
|
||||||
|
|
||||||
## 技术栈
|
## 技术栈
|
||||||
|
|
||||||
| 层级 | 技术 |
|
| 层级 | 技术 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| 前端 | React 18, TypeScript, Vite, @ricky0123/vad-web |
|
| 前端 | React 18, TypeScript, Vite, @ricky0123/vad-web, onnxruntime-web |
|
||||||
| 后端 | Go, Gin, gorilla/websocket, Viper, Zap |
|
| 后端 | Go 1.25+ (go.mod 最低要求; Dockerfile 构建用 golang:1.26-alpine), Gin, gorilla/websocket, Viper, Zap |
|
||||||
| AI 编排 | CloudWeGo Eino Graph(声明式 DAG 编排) |
|
| AI 编排 | CloudWeGo Eino Graph(声明式 DAG 编排) |
|
||||||
| LLM | DashScope qwen3-vl-plus(默认,通过 eino-ext OpenAI ChatModel 接入) |
|
| LLM | DashScope qwen3-vl-plus(默认,通过 eino-ext OpenAI ChatModel 接入) |
|
||||||
| STT | MiMo ASR(默认) / Deepgram |
|
| STT | MiMo ASR(默认) / Deepgram |
|
||||||
| TTS | MiMo TTS(默认) / OpenAI TTS |
|
| TTS | MiMo TTS(默认) / OpenAI TTS |
|
||||||
|
| 存储 | PostgreSQL 15 + Redis 7(通过 TieredManager 三级存储) |
|
||||||
|
|
||||||
## 构建与运行命令
|
## 构建与运行命令
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# 前端
|
# 前端
|
||||||
cd frontend && npm install
|
cd frontend && npm install
|
||||||
npm run dev # Vite 开发服务器
|
npm run dev # Vite 开发服务器(含 /ws、/api 代理到 localhost:8080)
|
||||||
npm run build # 生产构建
|
npm run build # 生产构建(tsc -b && vite build)
|
||||||
npm run lint # ESLint 检查
|
npm run lint # ESLint 检查(flat config, TypeScript strict)
|
||||||
npm run test # Vitest 测试
|
|
||||||
|
|
||||||
# 后端
|
# 后端
|
||||||
cd backend && go mod download
|
cd backend && go mod download
|
||||||
@@ -48,9 +101,96 @@ go build -o bin/camtalk ./cmd/server
|
|||||||
go test ./... # 运行所有测试
|
go test ./... # 运行所有测试
|
||||||
go test -run TestName ./path # 运行单个测试
|
go test -run TestName ./path # 运行单个测试
|
||||||
go vet ./... # 静态分析
|
go vet ./... # 静态分析
|
||||||
|
|
||||||
|
# Docker 部署(生产环境)
|
||||||
|
./deploy.sh build # 构建所有镜像
|
||||||
|
./deploy.sh up # 启动 4 个服务
|
||||||
|
./deploy.sh restart # down + up
|
||||||
|
./deploy.sh logs [service] # 查看日志
|
||||||
|
./deploy.sh status # 查看服务状态
|
||||||
```
|
```
|
||||||
|
|
||||||
基础设施:三级存储架构(L1 Memory → L2 Redis → L3 PostgreSQL),通过配置控制启用层级。
|
> **注意**:前端目前没有测试基础设施(无 vitest 配置、无测试文件)。后端使用 `testing` + `testify`(assert/require/mock)测试,编译期接口检查 `var _ Interface = (*Impl)(nil)`。
|
||||||
|
|
||||||
|
## 配置环境切换
|
||||||
|
|
||||||
|
项目通过 `APP_ENV` 环境变量控制配置文件加载:
|
||||||
|
|
||||||
|
- **本地开发**(默认):`APP_ENV=dev` → 加载 `config/config.dev.yaml`
|
||||||
|
- Debug 日志、关闭限流、允许所有 CORS
|
||||||
|
- 使用 `backend/.env` 中的远程 Redis/PostgreSQL 地址
|
||||||
|
|
||||||
|
- **生产部署**:`APP_ENV=prod` → 加载 `config/config.prod.yaml`
|
||||||
|
- Info/JSON 日志、启用限流、严格 CORS 白名单
|
||||||
|
- Docker Compose 自动设置,使用容器内网地址
|
||||||
|
|
||||||
|
**配置优先级**:环境变量 > config.{env}.yaml > config.yaml > 默认值
|
||||||
|
|
||||||
|
**手动切换环境**(测试用):
|
||||||
|
```bash
|
||||||
|
cd backend
|
||||||
|
APP_ENV=prod go run ./cmd/server # 本地测试生产配置
|
||||||
|
APP_ENV=dev go run ./cmd/server # 显式指定开发配置
|
||||||
|
```
|
||||||
|
|
||||||
|
## 配置系统
|
||||||
|
|
||||||
|
配置文件:`backend/config/config.yaml`(基础配置),可被 `config/config.{env}.yaml` 覆盖。
|
||||||
|
|
||||||
|
**优先级(从低到高)**:默认值 → `config.yaml` → `config.{env}.yaml`(由 `APP_ENV` 环境变量决定加载哪个 env 特定文件)→ `.env` 文件 → 环境变量
|
||||||
|
|
||||||
|
**主要 `CAMTALK_` 环境变量**(模板见 `backend/.env.example`):
|
||||||
|
|
||||||
|
| 变量 | 用途 |
|
||||||
|
|------|------|
|
||||||
|
| `APP_ENV` | 运行环境(dev/prod),决定加载 `config.{env}.yaml` |
|
||||||
|
| `CAMTALK_AI_STT_API_KEY` | STT API Key |
|
||||||
|
| `CAMTALK_AI_LLM_API_KEY` | LLM API Key |
|
||||||
|
| `CAMTALK_AI_TTS_API_KEY` | TTS API Key |
|
||||||
|
| `CAMTALK_AUTH_JWT_SECRET` | JWT 签名密钥 |
|
||||||
|
| `CAMTALK_STORAGE_DSN` | PostgreSQL 连接串 |
|
||||||
|
| `CAMTALK_STORAGE_REDIS_ENABLED` | 启用 Redis(true/false) |
|
||||||
|
| `CAMTALK_STORAGE_PERSISTENCE_ENABLED` | 启用 PostgreSQL(true/false) |
|
||||||
|
| `CAMTALK_REDIS_ADDR` | Redis 地址 |
|
||||||
|
| `CAMTALK_REDIS_PASSWORD` | Redis 密码 |
|
||||||
|
|
||||||
|
**最小启动**(至少需要一个 AI 服务的 API Key):
|
||||||
|
```bash
|
||||||
|
CAMTALK_AI_LLM_API_KEY=sk-xxx CAMTALK_AI_STT_API_KEY=xxx go run ./cmd/server
|
||||||
|
```
|
||||||
|
|
||||||
|
## Docker 部署
|
||||||
|
|
||||||
|
`docker-compose.yml` 定义 4 个服务(`camtalk-net` 桥接网络):
|
||||||
|
|
||||||
|
| 服务 | 镜像/构建 | 端口 | 说明 |
|
||||||
|
|------|----------|------|------|
|
||||||
|
| `frontend` | 构建 `./frontend/Dockerfile`(node:22-alpine → nginx:stable-alpine) | 9000:80 | React SPA,反向代理 /api 和 /ws 到 backend |
|
||||||
|
| `backend` | 构建 `./backend/Dockerfile`(golang:1.26-alpine → alpine:3.20) | 内部 8080 | Go 网关,静态链接二进制 `-ldflags="-s -w"` |
|
||||||
|
| `postgres` | `postgres:15-alpine` | 内部 5432 | 数据库 `camtalk`,挂载 `./backend/migrations/` 到 initdb |
|
||||||
|
| `redis` | `redis:7-alpine` | 内部 6379 | 会话缓存,AOF 持久化 |
|
||||||
|
|
||||||
|
后端容器依赖 postgres + redis 健康检查通过后启动。所有服务 `restart: unless-stopped`。密钥通过 `--env-file /opt/camtalk/.env` 注入。
|
||||||
|
|
||||||
|
**前端 nginx 特殊配置**:设置 `Cross-Origin-Opener-Policy` 和 `Cross-Origin-Embedder-Policy` 头(`SharedArrayBuffer` 需要,ONNX WASM 推理依赖)。
|
||||||
|
|
||||||
|
## CI/CD
|
||||||
|
|
||||||
|
使用 **Gitea Actions**(`.gitea/workflows/deploy.yml`),自托管 runner(标签 `aliyun`)。
|
||||||
|
|
||||||
|
触发条件:push 到 `main` 或 `v2` 分支。流程:rsync 代码到 `/root/camtalk`,执行 `deploy.sh build` → `deploy.sh restart`。
|
||||||
|
|
||||||
|
## 数据库迁移
|
||||||
|
|
||||||
|
嵌入式 SQL 迁移系统(`backend/internal/store/migrate.go`),SQL 文件在 `backend/migrations/`:
|
||||||
|
|
||||||
|
| 迁移 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| `001_users` | `users` 表(UUID PK)+ `refresh_tokens` 表(FK → users) |
|
||||||
|
| `002_messages` | `messages` 表(BIGSERIAL PK, session_id UUID, 游标分页索引) |
|
||||||
|
| `003_sessions` | `sessions` 表(UUID PK, user_id UUID, config JSONB, 时间排序索引) |
|
||||||
|
|
||||||
|
迁移文件通过 Go 1.16+ `//go:embed` 嵌入二进制,启动时自动执行。通过 `schema_migrations` 表追踪版本,已应用的迁移跳过。同时挂载到 PostgreSQL 容器的 `/docker-entrypoint-initdb.d` 作为备用初始化路径。
|
||||||
|
|
||||||
## WebSocket 协议
|
## WebSocket 协议
|
||||||
|
|
||||||
@@ -58,12 +198,16 @@ go vet ./... # 静态分析
|
|||||||
|
|
||||||
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。完整契约见 `docs/02-接口文档.md`。
|
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。完整契约见 `docs/02-接口文档.md`。
|
||||||
|
|
||||||
|
**认证**:WebSocket 连接通过 query param `token`(Access Token)认证,不走 HTTP `Authorization` header。服务端在升级时校验 JWT,失败返回 401。
|
||||||
|
|
||||||
**客户端 → 服务端**:`query`(图像 Base64 + 音频 Base64)、`config`、`interrupt`、`ping`
|
**客户端 → 服务端**:`query`(图像 Base64 + 音频 Base64)、`config`、`interrupt`、`ping`
|
||||||
**服务端 → 客户端**:`connected`、`stt_result`、`llm_chunk`、`llm_done`、`tts_audio`、`error`、`pong`
|
**服务端 → 客户端**:`connected`、`stt_result`、`llm_chunk`、`llm_done`、`tts_audio`、`error`、`pong`
|
||||||
|
|
||||||
**心跳**:客户端每 30 秒 ping,服务端 60 秒无 ping 断开连接。
|
**心跳**:客户端每 30 秒 ping,服务端 60 秒无 ping 断开连接。
|
||||||
**重连**:指数退避 + 抖动 —— 1s, 2s, 4s, 8s… 最大 30s。
|
**重连**:指数退避 + 抖动 —— 1s, 2s, 4s, 8s… 最大 30s。
|
||||||
|
|
||||||
|
**前端 WebSocket 实现**:`CamTalkWebSocket` 单例类(`frontend/src/lib/websocket.ts`),基于订阅模式(`onMessage`/`onStatusChange` 返回取消订阅函数),自动处理心跳和重连。
|
||||||
|
|
||||||
## REST API(辅助)
|
## REST API(辅助)
|
||||||
|
|
||||||
- `GET /api/health` — 健康检查(版本、运行时间、活跃会话数)
|
- `GET /api/health` — 健康检查(版本、运行时间、活跃会话数)
|
||||||
@@ -74,7 +218,8 @@ go vet ./... # 静态分析
|
|||||||
- `GET /api/conversations` — 对话列表
|
- `GET /api/conversations` — 对话列表
|
||||||
- `POST /api/conversations` — 创建对话
|
- `POST /api/conversations` — 创建对话
|
||||||
- `GET/PATCH/DELETE /api/conversations/:id` — 对话详情/改标题/删除
|
- `GET/PATCH/DELETE /api/conversations/:id` — 对话详情/改标题/删除
|
||||||
- `GET /api/conversations/:id/messages` — 获取对话消息
|
- `GET /api/conversations/:id/messages` — 获取对话消息(游标分页)
|
||||||
|
- `POST/DELETE /api/sessions` — 会话管理
|
||||||
|
|
||||||
## 错误码
|
## 错误码
|
||||||
|
|
||||||
@@ -86,39 +231,64 @@ go vet ./... # 静态分析
|
|||||||
|------|------|
|
|------|------|
|
||||||
| `LandingPage` | 未登录时的着陆页,内嵌 LoginModal 登录/注册弹窗 |
|
| `LandingPage` | 未登录时的着陆页,内嵌 LoginModal 登录/注册弹窗 |
|
||||||
| `AuthPage` | 登录/注册表单(备用) |
|
| `AuthPage` | 登录/注册表单(备用) |
|
||||||
| `CameraManager` | 摄像头流采集 |
|
| `CameraManager` | 摄像头流采集(`useCamera` hook:640x480, facingMode: environment) |
|
||||||
| `MicManager` | 麦克风音频采集 |
|
| `MicManager` | 麦克风音频采集(`useMicrophone` hook:16kHz 单声道) |
|
||||||
| `EdgeProcessor` | VAD + 关键帧检测(Canvas 像素比较) |
|
| `EdgeProcessor` | VAD(`useVAD` hook:@ricky0123/vad-web)+ 关键帧检测(Canvas 像素比较,160x120 降采样) |
|
||||||
| `WebSocketManager` | WebSocket 连接生命周期管理 |
|
| `WebSocketManager` | WebSocket 连接生命周期管理(桥接 `wsClient` 单例到 React 状态) |
|
||||||
| `ChatPanel` | 消息展示、流式回复、文本输入、场景选择 |
|
| `ChatPanel` | 消息展示、流式回复、文本输入、场景选择(5 种场景卡片) |
|
||||||
| `VideoPreview` | 摄像头画面预览 |
|
| `VideoPreview` | 摄像头画面预览(forwardRef `<video>`) |
|
||||||
| `SessionSidebar` | 左侧抽屉式对话列表(搜索、重命名、删除、时间分组) |
|
| `SessionSidebar` | 左侧抽屉式对话列表(搜索、重命名、删除、时间分组) |
|
||||||
| `ConfigPanel` | 右侧抽屉式配置面板(主题、TTS、语言、场景、登出) |
|
| `ConfigPanel` | 右侧抽屉式配置面板(主题、TTS、语言、场景、登出) |
|
||||||
| `Toast` | 轻量通知提示 |
|
| `Toast` | 轻量通知提示(3 秒自动消失) |
|
||||||
|
|
||||||
核心 Hook:`useVisionSession()` 封装一次完整的视觉对话会话。`useSessionList()` 管理对话列表 CRUD(通过 REST API)。
|
核心 Hook:
|
||||||
|
- `useVisionSession()` — 封装一次完整的视觉对话会话(~500 行),管理摄像头、麦克风、VAD、WebSocket、消息状态、TTS 播放、两种模式(dialogue / observation)
|
||||||
|
- `useSessionList()` — 对话列表 CRUD(通过 REST API),乐观更新
|
||||||
|
- `useObservationMode()` — 定期帧差异检测(5 秒间隔),相似度 < 0.85 时触发回调
|
||||||
|
|
||||||
|
关键 lib 文件:
|
||||||
|
- `frontend/src/lib/websocket.ts` — `CamTalkWebSocket` 单例,心跳 + 指数退避重连
|
||||||
|
- `frontend/src/lib/auth.tsx` — `AuthProvider` 上下文,JWT 解码 + 自动刷新调度(exp 前 60 秒)
|
||||||
|
- `frontend/src/lib/api.ts` — REST 客户端,自动 Bearer header,并发安全 401 拦截 + token 刷新 + 重试
|
||||||
|
- `frontend/src/lib/ttsPlayer.ts` — `TTSPlayer` 类,流式 TTS 音频逐句排队播放
|
||||||
|
- `frontend/src/lib/scenarios.ts` — 5 种对话场景定义(free_chat, interviewer, english_teacher, debate, interpreter)
|
||||||
|
- `frontend/src/lib/i18n/` — 国际化,3 种语言(zh-CN 默认/fallback, en-US, ja-JP),扁平常量 map
|
||||||
|
- `frontend/src/lib/storage.ts` — localStorage 封装(config, tokens, user info)
|
||||||
|
- `frontend/src/lib/audio.ts` — 音频编码工具(浏览器采集 → Base64 PCM)
|
||||||
|
- `frontend/src/lib/sampling.ts` — 混合采样策略(定时低频 + 事件高频,实现见 `docs/07-成本控制.md`)
|
||||||
|
- `frontend/src/lib/errors.ts` — 错误码到用户友好文案的映射
|
||||||
|
- `frontend/src/lib/toast.ts` — Toast 全局状态管理(error/warning/info,3 秒自动消失)
|
||||||
|
- `frontend/src/types/index.ts` — 所有 TypeScript 类型定义(WebSocket 消息可辨识联合类型、场景、配置等)
|
||||||
|
|
||||||
|
## Vite 构建细节
|
||||||
|
|
||||||
|
`frontend/vite.config.ts` 包含:
|
||||||
|
- 自定义 `serve-vad-assets` 插件,在开发/构建时自动从 `node_modules` 复制 VAD 模型文件(`silero_vad_legacy.onnx`、`silero_vad_v5.onnx`、`vad.worklet.bundle.min.js`)和 ONNX Runtime WASM 文件到 `public/`。开发服务器中间件确保 `.wasm` 和 `.mjs` 文件返回正确的 MIME type。
|
||||||
|
- 开发代理:`/ws` → `ws://localhost:8080`、`/api` → `http://localhost:8080`,前端开发时无需配置额外环境变量。
|
||||||
|
|
||||||
## 后端模块结构
|
## 后端模块结构
|
||||||
|
|
||||||
| 模块 | 职责 |
|
| 模块 | 职责 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| WebSocket Handler | 连接管理、JWT 认证、单播消息推送 |
|
| WebSocket Handler | 连接管理、JWT 认证(query param token)、消息分发(query/config/interrupt/ping) |
|
||||||
| Session Manager | 会话状态、对话历史(三级存储:Memory/Redis/PostgreSQL,30 分钟 TTL) |
|
| Session Manager | 会话状态、对话历史(三级存储:Memory/Redis/PostgreSQL,30 分钟 TTL) |
|
||||||
| Eino 编排层 | 基于 Eino Graph 的声明式 AI 编排(7 节点 DAG,Stream 模式,Callback AOP) |
|
| Eino 编排层 | 基于 Eino Graph 的声明式 AI 编排(6 节点线性 DAG + Callback) |
|
||||||
| AI Orchestrator | `EinoOrchestrator` 适配器,包装 Graph 实现 `Orchestrator` 接口 |
|
| AI Orchestrator | `EinoOrchestrator` 适配器,包装 Graph 实现 `Orchestrator` 接口 |
|
||||||
| AI Service Layer | AI 服务抽象层(STT/TTS 多 provider,LLM 通过 eino-ext ChatModel) |
|
| AI Service Layer | AI 服务抽象层(STT/TTS 多 provider 接口,LLM 通过 eino-ext ChatModel) |
|
||||||
| Auth | JWT 双 token 轮转认证,bcrypt 密码哈希 |
|
| Auth | JWT 双 token 轮转认证,bcrypt 密码哈希,Gin 中间件 |
|
||||||
| Store | 持久化存储层(UserRepository/MessageRepository/SessionRepository,内存 + PostgreSQL) |
|
| Store | 持久化存储层(UserRepository/MessageRepository/SessionRepository,PG + 内存 + Redis 缓存装饰器) |
|
||||||
| REST API | 健康检查、认证、对话管理(Gin 路由) |
|
| REST API | 健康检查、认证、对话管理(Gin 路由组) |
|
||||||
| Models | 数据模型定义 |
|
| Models | 数据模型 + WebSocket 消息类型定义 |
|
||||||
| Migrations | 数据库版本化迁移(嵌入式 SQL) |
|
| Migrations | 嵌入式 SQL 版本化迁移 |
|
||||||
| Model Router | 按请求选择 AI 模型(规划中) |
|
| Rate Limiter | 按用户的令牌桶速率限制(`docs/13-令牌桶限流设计.md`,`backend/internal/ratelimit/`,实现中) |
|
||||||
| Rate Limiter | 按用户的令牌桶速率限制(规划中) |
|
|
||||||
|
测试模式:后端使用 `testing` + `testify`(`assert`、`require`、`mock`)。mock 模式包括:`mockSender`(实现 `orchestrator.Sender` 接口)、`httptest.Server`(模拟 AI 服务 HTTP API)、`MockOrchestrator`(模拟完整编排管道)。WebSocket 集成测试使用 `httptest.Server` + `gorilla/websocket.Dialer`。
|
||||||
|
|
||||||
## 编码规范
|
## 编码规范
|
||||||
|
|
||||||
- **Go**:遵循标准 Go 规范。所有 AI 调用使用 `context.Context` 做取消/超时。并发 map 访问使用 `sync.RWMutex`。结构体标签用 `json:"snake_case"`。
|
- **Go**:遵循标准 Go 规范。所有 AI 调用使用 `context.Context` 做取消/超时。并发 map 访问使用 `sync.RWMutex`。结构体标签用 `json:"snake_case"`。编译期接口检查 `var _ Interface = (*Impl)(nil)`。
|
||||||
- **TypeScript**:严格模式。所有数据模型用接口定义。WebSocket 消息类型用可辨识联合类型(`type` 字段)。
|
- **TypeScript**:严格模式(`strict: true`)。所有数据模型用接口定义。WebSocket 消息类型用可辨识联合类型(`type` 字段)。`verbatimModuleSyntax: true`(强制 `import type`)。未使用变量以 `_` 前缀忽略。
|
||||||
|
- **CORS 处理**:禁止在后端代码和配置文件(`config/*.yaml`)中进行任何 CORS 配置。跨域由代理层统一处理:开发环境通过 `frontend/vite.config.ts` 中的 proxy 配置(`/ws`、`/api` 代理到 `localhost:8080`),生产环境通过 Nginx 反向代理(`frontend/nginx.conf`)。
|
||||||
- **提交信息**:Conventional Commits 格式,描述用中文。示例:`feat: 添加 WebSocket 连接管理`、`fix: 修复心跳超时判断`、`docs: 更新接口文档`
|
- **提交信息**:Conventional Commits 格式,描述用中文。示例:`feat: 添加 WebSocket 连接管理`、`fix: 修复心跳超时判断`、`docs: 更新接口文档`
|
||||||
- **禁止自动 push**:除非用户明确要求。
|
- **禁止自动 push**:除非用户明确要求。
|
||||||
- **文档优先**:实现功能前先读取 `docs/` 下的相关设计文档。实现与文档不一致时,优先更新 `docs/` 下的接口文档。
|
- **文档优先**:实现功能前先读取 `docs/` 下的相关设计文档。实现与文档不一致时,优先更新 `docs/` 下的接口文档。
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
# 运行环境
|
# 运行环境
|
||||||
# dev / prod,决定加载 config.dev.yaml 或 config.prod.yaml(可选)
|
# dev:本地开发环境(debug 日志、关闭限流、允许所有 CORS)
|
||||||
|
# prod:生产环境(info 日志、启用限流、严格 CORS 白名单)
|
||||||
|
# 本地开发保持 dev,生产部署会被 docker-compose.yml 覆盖为 prod
|
||||||
APP_ENV=dev
|
APP_ENV=dev
|
||||||
|
|
||||||
# AI 服务 API Key
|
# AI 服务 API Key
|
||||||
|
|||||||
2
backend/.gitignore
vendored
2
backend/.gitignore
vendored
@@ -4,8 +4,6 @@ bin/
|
|||||||
|
|
||||||
# 环境配置
|
# 环境配置
|
||||||
.env
|
.env
|
||||||
config.dev.yaml
|
|
||||||
config.prod.yaml
|
|
||||||
|
|
||||||
# 临时文件
|
# 临时文件
|
||||||
tmp/
|
tmp/
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ import (
|
|||||||
"github.com/hhs/camtalk/internal/config"
|
"github.com/hhs/camtalk/internal/config"
|
||||||
eino "github.com/hhs/camtalk/internal/eino"
|
eino "github.com/hhs/camtalk/internal/eino"
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"github.com/hhs/camtalk/internal/ratelimit"
|
||||||
"github.com/hhs/camtalk/internal/session"
|
"github.com/hhs/camtalk/internal/session"
|
||||||
"github.com/hhs/camtalk/internal/store"
|
"github.com/hhs/camtalk/internal/store"
|
||||||
"github.com/hhs/camtalk/internal/ws"
|
"github.com/hhs/camtalk/internal/ws"
|
||||||
@@ -195,6 +196,23 @@ func main() {
|
|||||||
)
|
)
|
||||||
authService := auth.NewAuthService(tokenMgr, userRepo)
|
authService := auth.NewAuthService(tokenMgr, userRepo)
|
||||||
|
|
||||||
|
// 初始化限流器
|
||||||
|
var limiter ratelimit.Limiter
|
||||||
|
if cfg.RateLimit.Enabled {
|
||||||
|
if rdb != nil {
|
||||||
|
// 多实例:使用 Redis 令牌桶
|
||||||
|
limiter = ratelimit.NewRedisLimiter(rdb, cfg.RateLimit)
|
||||||
|
logger.Log.Info("rate limiter initialized with Redis backend")
|
||||||
|
} else {
|
||||||
|
// 单实例:使用内存令牌桶
|
||||||
|
limiter = ratelimit.NewMemoryLimiter(cfg.RateLimit)
|
||||||
|
logger.Log.Info("rate limiter initialized with in-memory backend")
|
||||||
|
}
|
||||||
|
defer limiter.Stop()
|
||||||
|
} else {
|
||||||
|
logger.Log.Info("rate limiter disabled")
|
||||||
|
}
|
||||||
|
|
||||||
// Gin 模式
|
// Gin 模式
|
||||||
if cfg.App.Env == "prod" {
|
if cfg.App.Env == "prod" {
|
||||||
gin.SetMode(gin.ReleaseMode)
|
gin.SetMode(gin.ReleaseMode)
|
||||||
@@ -215,14 +233,14 @@ func main() {
|
|||||||
|
|
||||||
// Auth REST 端点
|
// Auth REST 端点
|
||||||
authHandler := api.NewAuthHandler(authService, tokenMgr)
|
authHandler := api.NewAuthHandler(authService, tokenMgr)
|
||||||
authHandler.RegisterRoutes(apiGroup)
|
authHandler.RegisterRoutes(apiGroup, limiter)
|
||||||
|
|
||||||
// Conversation REST 端点
|
// Conversation REST 端点
|
||||||
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
|
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
|
||||||
convHandler.RegisterRoutes(apiGroup)
|
convHandler.RegisterRoutes(apiGroup)
|
||||||
|
|
||||||
// WebSocket
|
// WebSocket
|
||||||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter))
|
||||||
|
|
||||||
// HTTP Server
|
// HTTP Server
|
||||||
srv := &http.Server{
|
srv := &http.Server{
|
||||||
|
|||||||
68
backend/config/config.dev.yaml
Normal file
68
backend/config/config.dev.yaml
Normal file
@@ -0,0 +1,68 @@
|
|||||||
|
# CamTalk 开发环境配置
|
||||||
|
# 通过 APP_ENV=dev 加载此文件,覆盖 config.yaml 中的配置
|
||||||
|
|
||||||
|
server:
|
||||||
|
host: "0.0.0.0"
|
||||||
|
port: 8080
|
||||||
|
heartbeat_interval: 30
|
||||||
|
heartbeat_timeout: 60
|
||||||
|
allowed_origins: [] # 开发环境允许所有来源
|
||||||
|
|
||||||
|
session:
|
||||||
|
ttl: 30 # 开发环境会话较短,方便测试过期逻辑
|
||||||
|
max_history: 20
|
||||||
|
|
||||||
|
ai:
|
||||||
|
stt:
|
||||||
|
provider: mimo # 与生产环境一致
|
||||||
|
model: mimo-v2.5-asr
|
||||||
|
endpoint: "https://api.xiaomimimo.com/v1"
|
||||||
|
timeout: 10 # 开发环境超时较长,方便调试
|
||||||
|
http_client_timeout: 30
|
||||||
|
llm:
|
||||||
|
provider: dashscope # 与生产环境一致
|
||||||
|
model: qwen3-vl-plus
|
||||||
|
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
timeout: 60 # 开发环境 LLM 超时较长
|
||||||
|
http_client_timeout: 120
|
||||||
|
tts:
|
||||||
|
provider: mimo # 与生产环境一致
|
||||||
|
model: mimo-v2.5-tts
|
||||||
|
voice: mimo_default
|
||||||
|
speed: 1.0
|
||||||
|
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||||
|
timeout: 10
|
||||||
|
http_client_timeout: 30
|
||||||
|
output_format: mp3
|
||||||
|
sample_rate: 24000
|
||||||
|
|
||||||
|
storage:
|
||||||
|
redis:
|
||||||
|
enabled: true # 开发环境启用 Redis,测试三级存储
|
||||||
|
persistence:
|
||||||
|
enabled: true # 开发环境启用持久化
|
||||||
|
|
||||||
|
redis:
|
||||||
|
addr: "localhost:6379" # 本地 Redis
|
||||||
|
password: ""
|
||||||
|
db: 0
|
||||||
|
|
||||||
|
auth:
|
||||||
|
access_ttl: 120 # 开发环境 Access Token 2 小时,方便调试
|
||||||
|
refresh_ttl: 10080 # 7 天
|
||||||
|
|
||||||
|
ratelimit:
|
||||||
|
enabled: false # 开发环境关闭限流,方便测试
|
||||||
|
query:
|
||||||
|
capacity: 10
|
||||||
|
rate: 0.2
|
||||||
|
login:
|
||||||
|
capacity: 5
|
||||||
|
rate: 0.1
|
||||||
|
register:
|
||||||
|
capacity: 3
|
||||||
|
rate: 0.05
|
||||||
|
|
||||||
|
log:
|
||||||
|
level: debug # 开发环境 debug 日志
|
||||||
|
format: console # 控制台格式,易读
|
||||||
71
backend/config/config.prod.yaml
Normal file
71
backend/config/config.prod.yaml
Normal file
@@ -0,0 +1,71 @@
|
|||||||
|
# CamTalk 生产环境配置
|
||||||
|
# 通过 APP_ENV=prod 加载此文件,覆盖 config.yaml 中的配置
|
||||||
|
|
||||||
|
server:
|
||||||
|
host: "0.0.0.0"
|
||||||
|
port: 8080
|
||||||
|
read_timeout: 30
|
||||||
|
write_timeout: 30
|
||||||
|
shutdown_timeout: 15 # 生产环境优雅关闭时间稍长
|
||||||
|
heartbeat_interval: 30
|
||||||
|
heartbeat_timeout: 60
|
||||||
|
|
||||||
|
session:
|
||||||
|
ttl: 60 # 生产环境会话 1 小时
|
||||||
|
max_history: 20
|
||||||
|
|
||||||
|
ai:
|
||||||
|
stt:
|
||||||
|
provider: mimo # 生产环境推荐 MiMo,性价比高
|
||||||
|
model: mimo-v2.5-asr
|
||||||
|
endpoint: "https://api.xiaomimimo.com/v1"
|
||||||
|
timeout: 5 # 生产环境严格超时控制
|
||||||
|
http_client_timeout: 30
|
||||||
|
llm:
|
||||||
|
provider: dashscope # 生产环境推荐通义千问,稳定性好
|
||||||
|
model: qwen3-vl-plus
|
||||||
|
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
timeout: 30
|
||||||
|
http_client_timeout: 60
|
||||||
|
tts:
|
||||||
|
provider: mimo # 生产环境推荐 MiMo TTS
|
||||||
|
model: mimo-v2.5-tts
|
||||||
|
voice: mimo_default
|
||||||
|
speed: 1.0
|
||||||
|
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||||
|
timeout: 5
|
||||||
|
http_client_timeout: 30
|
||||||
|
output_format: mp3
|
||||||
|
sample_rate: 24000
|
||||||
|
|
||||||
|
storage:
|
||||||
|
redis:
|
||||||
|
enabled: true # 生产环境必须启用 Redis
|
||||||
|
persistence:
|
||||||
|
enabled: true # 生产环境必须启用持久化
|
||||||
|
driver: postgres
|
||||||
|
|
||||||
|
redis:
|
||||||
|
addr: "redis:6379" # Docker Compose 内部服务名
|
||||||
|
password: "" # 密码通过 CAMTALK_REDIS_PASSWORD 环境变量设置
|
||||||
|
db: 0
|
||||||
|
|
||||||
|
auth:
|
||||||
|
access_ttl: 120 # 生产环境 Access Token 2 小时
|
||||||
|
refresh_ttl: 10080 # Refresh Token 7 天
|
||||||
|
|
||||||
|
ratelimit:
|
||||||
|
enabled: true # 生产环境启用限流
|
||||||
|
query:
|
||||||
|
capacity: 10 # 允许突发 10 个请求
|
||||||
|
rate: 0.2 # 每 5 秒恢复 1 个令牌
|
||||||
|
login:
|
||||||
|
capacity: 5 # 防暴力破解
|
||||||
|
rate: 0.1 # 每 10 秒恢复 1 次
|
||||||
|
register:
|
||||||
|
capacity: 3 # 防批量注册
|
||||||
|
rate: 0.05 # 每 20 秒恢复 1 次
|
||||||
|
|
||||||
|
log:
|
||||||
|
level: info # 生产环境 info 级别
|
||||||
|
format: json # JSON 格式,便于日志收集和分析
|
||||||
@@ -60,6 +60,21 @@ auth:
|
|||||||
access_ttl: 120 # Access Token 过期时间(分钟)
|
access_ttl: 120 # Access Token 过期时间(分钟)
|
||||||
refresh_ttl: 10080 # Refresh Token 过期时间(分钟),7 天
|
refresh_ttl: 10080 # Refresh Token 过期时间(分钟),7 天
|
||||||
|
|
||||||
|
ratelimit:
|
||||||
|
enabled: false # 是否启用限流
|
||||||
|
# WebSocket query 消息限流(核心,控制 AI 成本)
|
||||||
|
query:
|
||||||
|
capacity: 10 # 突发容量:允许连续发 10 个 query
|
||||||
|
rate: 0.2 # 填充速率:每 5 秒补充 1 个令牌
|
||||||
|
# REST API 登录限流(防暴力破解)
|
||||||
|
login:
|
||||||
|
capacity: 5 # 突发容量:允许连续 5 次登录尝试
|
||||||
|
rate: 0.1 # 填充速率:每 10 秒补充 1 次
|
||||||
|
# REST API 注册限流
|
||||||
|
register:
|
||||||
|
capacity: 3 # 突发容量:允许连续 3 次注册
|
||||||
|
rate: 0.05 # 填充速率:每 20 秒补充 1 次
|
||||||
|
|
||||||
log:
|
log:
|
||||||
level: info # debug / info / warn / error
|
level: info # debug / info / warn / error
|
||||||
format: console # console / json
|
format: console # console / json
|
||||||
@@ -19,6 +19,7 @@ require (
|
|||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
github.com/alicebob/miniredis/v2 v2.38.0 // indirect
|
||||||
github.com/bahlo/generic-list-go v0.2.0 // indirect
|
github.com/bahlo/generic-list-go v0.2.0 // indirect
|
||||||
github.com/buger/jsonparser v1.1.1 // indirect
|
github.com/buger/jsonparser v1.1.1 // indirect
|
||||||
github.com/bytedance/gopkg v0.1.3 // indirect
|
github.com/bytedance/gopkg v0.1.3 // indirect
|
||||||
@@ -68,6 +69,7 @@ require (
|
|||||||
github.com/ugorji/go/codec v1.2.12 // indirect
|
github.com/ugorji/go/codec v1.2.12 // indirect
|
||||||
github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect
|
github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect
|
||||||
github.com/yargevad/filepathx v1.0.0 // indirect
|
github.com/yargevad/filepathx v1.0.0 // indirect
|
||||||
|
github.com/yuin/gopher-lua v1.1.1 // indirect
|
||||||
go.uber.org/atomic v1.11.0 // indirect
|
go.uber.org/atomic v1.11.0 // indirect
|
||||||
go.uber.org/multierr v1.10.0 // indirect
|
go.uber.org/multierr v1.10.0 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
github.com/airbrake/gobrake v3.6.1+incompatible/go.mod h1:wM4gu3Cn0W0K7GUuVWnlXZU11AGBXMILnrdOU8Kn00o=
|
github.com/airbrake/gobrake v3.6.1+incompatible/go.mod h1:wM4gu3Cn0W0K7GUuVWnlXZU11AGBXMILnrdOU8Kn00o=
|
||||||
|
github.com/alicebob/miniredis/v2 v2.38.0 h1:nZAzCR+Lj+Vxk4ZXzm2NuKq2O33RXj1XxJ2e2uP9jiw=
|
||||||
|
github.com/alicebob/miniredis/v2 v2.38.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM=
|
||||||
github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk=
|
github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk=
|
||||||
github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg=
|
github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg=
|
||||||
github.com/bitly/go-simplejson v0.5.0/go.mod h1:cXHtHw4XUPsvGaxgjIAn8PhEWG9NfngEKAMDJEczWVA=
|
github.com/bitly/go-simplejson v0.5.0/go.mod h1:cXHtHw4XUPsvGaxgjIAn8PhEWG9NfngEKAMDJEczWVA=
|
||||||
@@ -190,6 +192,8 @@ github.com/x-cray/logrus-prefixed-formatter v0.5.2 h1:00txxvfBM9muc0jiLIEAkAcIMJ
|
|||||||
github.com/x-cray/logrus-prefixed-formatter v0.5.2/go.mod h1:2duySbKsL6M18s5GU7VPsoEPHyzalCE06qoARUCeBBE=
|
github.com/x-cray/logrus-prefixed-formatter v0.5.2/go.mod h1:2duySbKsL6M18s5GU7VPsoEPHyzalCE06qoARUCeBBE=
|
||||||
github.com/yargevad/filepathx v1.0.0 h1:SYcT+N3tYGi+NvazubCNlvgIPbzAk7i7y2dwg3I5FYc=
|
github.com/yargevad/filepathx v1.0.0 h1:SYcT+N3tYGi+NvazubCNlvgIPbzAk7i7y2dwg3I5FYc=
|
||||||
github.com/yargevad/filepathx v1.0.0/go.mod h1:BprfX/gpYNJHJfc35GjRRpVcwWXS89gGulUIU5tK3tA=
|
github.com/yargevad/filepathx v1.0.0/go.mod h1:BprfX/gpYNJHJfc35GjRRpVcwWXS89gGulUIU5tK3tA=
|
||||||
|
github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M=
|
||||||
|
github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw=
|
||||||
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
|
||||||
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
|
||||||
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
|
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
|
|
||||||
"github.com/hhs/camtalk/internal/auth"
|
"github.com/hhs/camtalk/internal/auth"
|
||||||
apperr "github.com/hhs/camtalk/internal/errors"
|
apperr "github.com/hhs/camtalk/internal/errors"
|
||||||
|
"github.com/hhs/camtalk/internal/ratelimit"
|
||||||
)
|
)
|
||||||
|
|
||||||
// AuthHandler 提供认证相关的 REST 端点。
|
// AuthHandler 提供认证相关的 REST 端点。
|
||||||
@@ -25,11 +26,26 @@ func NewAuthHandler(authService auth.Service, tokenMgr *auth.TokenManager) *Auth
|
|||||||
}
|
}
|
||||||
|
|
||||||
// RegisterRoutes 注册认证相关路由到给定的路由组。
|
// RegisterRoutes 注册认证相关路由到给定的路由组。
|
||||||
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup, limiter ratelimit.Limiter) {
|
||||||
authGroup := rg.Group("/auth")
|
authGroup := rg.Group("/auth")
|
||||||
{
|
{
|
||||||
authGroup.POST("/register", h.Register)
|
// 注册和登录端点添加限流中间件(按 IP 限流)
|
||||||
authGroup.POST("/login", h.Login)
|
if limiter != nil {
|
||||||
|
authGroup.POST("/register",
|
||||||
|
ratelimit.Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return c.ClientIP() + ":register"
|
||||||
|
}),
|
||||||
|
h.Register)
|
||||||
|
authGroup.POST("/login",
|
||||||
|
ratelimit.Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return c.ClientIP() + ":login"
|
||||||
|
}),
|
||||||
|
h.Login)
|
||||||
|
} else {
|
||||||
|
authGroup.POST("/register", h.Register)
|
||||||
|
authGroup.POST("/login", h.Login)
|
||||||
|
}
|
||||||
|
// refresh 和 logout 不限流
|
||||||
authGroup.POST("/refresh", h.Refresh)
|
authGroup.POST("/refresh", h.Refresh)
|
||||||
authGroup.POST("/logout", auth.AuthMiddleware(h.tokenMgr), h.Logout)
|
authGroup.POST("/logout", auth.AuthMiddleware(h.tokenMgr), h.Logout)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ func newTestRouter(svc auth.Service) *gin.Engine {
|
|||||||
r := gin.New()
|
r := gin.New()
|
||||||
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
||||||
h := api.NewAuthHandler(svc, tm)
|
h := api.NewAuthHandler(svc, tm)
|
||||||
h.RegisterRoutes(r.Group("/api"))
|
h.RegisterRoutes(r.Group("/api"), nil) // 测试时不启用限流
|
||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -57,7 +57,7 @@ func newTestRouterWithToken(svc auth.Service) (*gin.Engine, *auth.TokenManager)
|
|||||||
r := gin.New()
|
r := gin.New()
|
||||||
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
||||||
h := api.NewAuthHandler(svc, tm)
|
h := api.NewAuthHandler(svc, tm)
|
||||||
h.RegisterRoutes(r.Group("/api"))
|
h.RegisterRoutes(r.Group("/api"), nil) // 测试时不启用限流
|
||||||
return r, tm
|
return r, tm
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -10,14 +10,15 @@ import (
|
|||||||
|
|
||||||
// Config 应用配置。
|
// Config 应用配置。
|
||||||
type Config struct {
|
type Config struct {
|
||||||
App AppConfig `mapstructure:"app"`
|
App AppConfig `mapstructure:"app"`
|
||||||
Server ServerConfig `mapstructure:"server"`
|
Server ServerConfig `mapstructure:"server"`
|
||||||
Session SessionConfig `mapstructure:"session"`
|
Session SessionConfig `mapstructure:"session"`
|
||||||
Redis RedisConfig `mapstructure:"redis"`
|
Redis RedisConfig `mapstructure:"redis"`
|
||||||
AI AIConfig `mapstructure:"ai"`
|
AI AIConfig `mapstructure:"ai"`
|
||||||
Storage StorageConfig `mapstructure:"storage"`
|
Storage StorageConfig `mapstructure:"storage"`
|
||||||
Log LogConfig `mapstructure:"log"`
|
Log LogConfig `mapstructure:"log"`
|
||||||
Auth AuthConfig `mapstructure:"auth"`
|
Auth AuthConfig `mapstructure:"auth"`
|
||||||
|
RateLimit RateLimitConfig `mapstructure:"ratelimit"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// SessionConfig 会话管理配置。
|
// SessionConfig 会话管理配置。
|
||||||
@@ -120,8 +121,22 @@ type AuthConfig struct {
|
|||||||
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 10080(7天)
|
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 10080(7天)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RateLimitConfig 限流配置。
|
||||||
|
type RateLimitConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
Query BucketConfig `mapstructure:"query"`
|
||||||
|
Login BucketConfig `mapstructure:"login"`
|
||||||
|
Register BucketConfig `mapstructure:"register"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// BucketConfig 令牌桶配置。
|
||||||
|
type BucketConfig struct {
|
||||||
|
Capacity int `mapstructure:"capacity"` // 桶容量(突发上限)
|
||||||
|
Rate float64 `mapstructure:"rate"` // 每秒填充令牌数
|
||||||
|
}
|
||||||
|
|
||||||
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml > 默认值。
|
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml > 默认值。
|
||||||
// workDir 为项目根目录或 backend 目录,用于定位 .env 和 config.yaml。
|
// workDir 为项目根目录或 backend 目录,用于定位 .env 和 config/config.yaml。
|
||||||
func Load(workDir string) (*Config, error) {
|
func Load(workDir string) (*Config, error) {
|
||||||
// 1. 加载 .env 文件(敏感信息)
|
// 1. 加载 .env 文件(敏感信息)
|
||||||
envFile := filepath.Join(workDir, ".env")
|
envFile := filepath.Join(workDir, ".env")
|
||||||
@@ -130,7 +145,8 @@ func Load(workDir string) (*Config, error) {
|
|||||||
v := viper.New()
|
v := viper.New()
|
||||||
v.SetConfigName("config")
|
v.SetConfigName("config")
|
||||||
v.SetConfigType("yaml")
|
v.SetConfigType("yaml")
|
||||||
v.AddConfigPath(workDir)
|
v.AddConfigPath(filepath.Join(workDir, "config")) // 配置文件在 config/ 目录下
|
||||||
|
v.AddConfigPath(workDir) // 兼容旧路径
|
||||||
|
|
||||||
// 2. 设置默认值(与 config.yaml 保持一致,仅作为兜底)
|
// 2. 设置默认值(与 config.yaml 保持一致,仅作为兜底)
|
||||||
setDefaults(v)
|
setDefaults(v)
|
||||||
@@ -218,6 +234,15 @@ func setDefaults(v *viper.Viper) {
|
|||||||
// log
|
// log
|
||||||
v.SetDefault("log.level", "info")
|
v.SetDefault("log.level", "info")
|
||||||
v.SetDefault("log.format", "console")
|
v.SetDefault("log.format", "console")
|
||||||
|
|
||||||
|
// ratelimit
|
||||||
|
v.SetDefault("ratelimit.enabled", false)
|
||||||
|
v.SetDefault("ratelimit.query.capacity", 10)
|
||||||
|
v.SetDefault("ratelimit.query.rate", 0.2)
|
||||||
|
v.SetDefault("ratelimit.login.capacity", 5)
|
||||||
|
v.SetDefault("ratelimit.login.rate", 0.1)
|
||||||
|
v.SetDefault("ratelimit.register.capacity", 3)
|
||||||
|
v.SetDefault("ratelimit.register.rate", 0.05)
|
||||||
}
|
}
|
||||||
|
|
||||||
// bindEnvVars 显式绑定敏感信息环境变量。
|
// bindEnvVars 显式绑定敏感信息环境变量。
|
||||||
|
|||||||
172
backend/internal/ratelimit/bucket.go
Normal file
172
backend/internal/ratelimit/bucket.go
Normal file
@@ -0,0 +1,172 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TokenBucket 内存令牌桶,适用于单实例部署。
|
||||||
|
type TokenBucket struct {
|
||||||
|
capacity int // 桶容量
|
||||||
|
rate float64 // 每秒填充令牌数
|
||||||
|
tokens float64 // 当前令牌数
|
||||||
|
lastRefill time.Time // 上次填充时间
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// newTokenBucket 创建令牌桶。
|
||||||
|
func newTokenBucket(capacity int, rate float64) *TokenBucket {
|
||||||
|
return &TokenBucket{
|
||||||
|
capacity: capacity,
|
||||||
|
rate: rate,
|
||||||
|
tokens: float64(capacity), // 初始满桶
|
||||||
|
lastRefill: time.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// allow 尝试消耗一个令牌。
|
||||||
|
func (b *TokenBucket) allow() (bool, time.Duration) {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
elapsed := now.Sub(b.lastRefill).Seconds()
|
||||||
|
|
||||||
|
// 补充令牌
|
||||||
|
newTokens := elapsed * b.rate
|
||||||
|
b.tokens = min(float64(b.capacity), b.tokens+newTokens)
|
||||||
|
b.lastRefill = now
|
||||||
|
|
||||||
|
// 尝试消耗一个令牌
|
||||||
|
if b.tokens >= 1 {
|
||||||
|
b.tokens -= 1
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// 计算需要等待的时间
|
||||||
|
if b.rate == 0 {
|
||||||
|
// rate=0 时永远无法补充令牌
|
||||||
|
return false, 24 * time.Hour // 返回一个很大的值
|
||||||
|
}
|
||||||
|
retryAfter := time.Duration((1-b.tokens)/b.rate*1000) * time.Millisecond
|
||||||
|
return false, retryAfter
|
||||||
|
}
|
||||||
|
|
||||||
|
// MemoryLimiter 管理多个用户的令牌桶。
|
||||||
|
type MemoryLimiter struct {
|
||||||
|
buckets map[string]*TokenBucket
|
||||||
|
config config.RateLimitConfig
|
||||||
|
mu sync.RWMutex
|
||||||
|
stopOnce sync.Once
|
||||||
|
done chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMemoryLimiter 创建内存限流器。
|
||||||
|
func NewMemoryLimiter(cfg config.RateLimitConfig) *MemoryLimiter {
|
||||||
|
limiter := &MemoryLimiter{
|
||||||
|
buckets: make(map[string]*TokenBucket),
|
||||||
|
config: cfg,
|
||||||
|
done: make(chan struct{}),
|
||||||
|
}
|
||||||
|
|
||||||
|
// 启动后台清理 goroutine
|
||||||
|
go limiter.cleanup()
|
||||||
|
|
||||||
|
return limiter
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow 实现 Limiter 接口。
|
||||||
|
func (l *MemoryLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
bucket := l.getOrCreateBucket(key)
|
||||||
|
return bucket.allow()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop 实现 Limiter 接口。
|
||||||
|
func (l *MemoryLimiter) Stop() {
|
||||||
|
l.stopOnce.Do(func() {
|
||||||
|
close(l.done)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// getOrCreateBucket 获取或创建令牌桶。
|
||||||
|
func (l *MemoryLimiter) getOrCreateBucket(key string) *TokenBucket {
|
||||||
|
// 先尝试读锁
|
||||||
|
l.mu.RLock()
|
||||||
|
bucket, exists := l.buckets[key]
|
||||||
|
l.mu.RUnlock()
|
||||||
|
|
||||||
|
if exists {
|
||||||
|
return bucket
|
||||||
|
}
|
||||||
|
|
||||||
|
// 需要创建新桶,升级为写锁
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
// 双重检查(可能其他 goroutine 已创建)
|
||||||
|
bucket, exists = l.buckets[key]
|
||||||
|
if exists {
|
||||||
|
return bucket
|
||||||
|
}
|
||||||
|
|
||||||
|
// 根据 key 确定配置(简化版:假设 key 格式为 "userID:action")
|
||||||
|
cfg := l.getBucketConfig(key)
|
||||||
|
bucket = newTokenBucket(cfg.Capacity, cfg.Rate)
|
||||||
|
l.buckets[key] = bucket
|
||||||
|
|
||||||
|
return bucket
|
||||||
|
}
|
||||||
|
|
||||||
|
// getBucketConfig 根据 key 获取桶配置。
|
||||||
|
func (l *MemoryLimiter) getBucketConfig(key string) config.BucketConfig {
|
||||||
|
// 简化实现:从 key 后缀判断动作类型
|
||||||
|
// 实际使用时调用方会传递正确的 key
|
||||||
|
// 默认使用 query 配置
|
||||||
|
return l.config.Query
|
||||||
|
}
|
||||||
|
|
||||||
|
// cleanup 定期清理不活跃的桶。
|
||||||
|
func (l *MemoryLimiter) cleanup() {
|
||||||
|
ticker := time.NewTicker(10 * time.Minute)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ticker.C:
|
||||||
|
l.removeInactiveBuckets()
|
||||||
|
case <-l.done:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeInactiveBuckets 移除超过 10 分钟无活动的桶。
|
||||||
|
func (l *MemoryLimiter) removeInactiveBuckets() {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
for key, bucket := range l.buckets {
|
||||||
|
bucket.mu.Lock()
|
||||||
|
inactive := now.Sub(bucket.lastRefill) > 10*time.Minute
|
||||||
|
bucket.mu.Unlock()
|
||||||
|
|
||||||
|
if inactive {
|
||||||
|
delete(l.buckets, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// min 返回两个 float64 中的较小值。
|
||||||
|
func min(a, b float64) float64 {
|
||||||
|
if a < b {
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// 编译期接口检查
|
||||||
|
var _ Limiter = (*MemoryLimiter)(nil)
|
||||||
203
backend/internal/ratelimit/bucket_test.go
Normal file
203
backend/internal/ratelimit/bucket_test.go
Normal file
@@ -0,0 +1,203 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_FirstRequest(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(5, 0.2)
|
||||||
|
|
||||||
|
allowed, retryAfter := bucket.allow()
|
||||||
|
|
||||||
|
assert.True(t, allowed)
|
||||||
|
assert.Equal(t, time.Duration(0), retryAfter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_ConsumeUntilEmpty(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(3, 0.2)
|
||||||
|
|
||||||
|
// 连续消耗 3 个令牌
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 第 4 个请求应被拒绝
|
||||||
|
allowed, retryAfter := bucket.allow()
|
||||||
|
assert.False(t, allowed)
|
||||||
|
assert.Greater(t, retryAfter, time.Duration(0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_RetryAfterCorrect(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(1, 1.0) // 每秒 1 个令牌
|
||||||
|
|
||||||
|
// 消耗唯一的令牌
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
require.True(t, allowed)
|
||||||
|
|
||||||
|
// 立即再次请求应被拒绝
|
||||||
|
allowed, retryAfter := bucket.allow()
|
||||||
|
assert.False(t, allowed)
|
||||||
|
// retryAfter 应约为 1 秒(允许一定误差)
|
||||||
|
assert.InDelta(t, 1000, retryAfter.Milliseconds(), 100)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_RefillAfterWait(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(2, 10.0) // 每秒 10 个令牌(每 100ms 一个)
|
||||||
|
|
||||||
|
// 消耗 2 个令牌
|
||||||
|
bucket.allow()
|
||||||
|
bucket.allow()
|
||||||
|
|
||||||
|
// 等待 150ms,应补充至少 1 个令牌
|
||||||
|
time.Sleep(150 * time.Millisecond)
|
||||||
|
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
assert.True(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_CapacityLimit(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(3, 1.0)
|
||||||
|
|
||||||
|
// 等待足够长时间让桶"溢出"
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
|
||||||
|
// 但最多只能消耗 capacity 个令牌
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 第 4 个应被拒绝
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
assert.False(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_ConcurrentSafe(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(100, 10.0)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
successCount := 0
|
||||||
|
var mu sync.Mutex
|
||||||
|
|
||||||
|
// 100 个并发请求
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
if allowed {
|
||||||
|
mu.Lock()
|
||||||
|
successCount++
|
||||||
|
mu.Unlock()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
// 应该正好 100 个成功(桶容量为 100)
|
||||||
|
assert.Equal(t, 100, successCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_ZeroCapacity(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(0, 1.0)
|
||||||
|
|
||||||
|
allowed, retryAfter := bucket.allow()
|
||||||
|
assert.False(t, allowed)
|
||||||
|
assert.Greater(t, retryAfter, time.Duration(0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenBucket_Allow_ZeroRate(t *testing.T) {
|
||||||
|
bucket := newTokenBucket(1, 0.0)
|
||||||
|
|
||||||
|
// 第一个通过
|
||||||
|
allowed, _ := bucket.allow()
|
||||||
|
assert.True(t, allowed)
|
||||||
|
|
||||||
|
// 第二个被拒绝,且 retryAfter 应为无限大(实际上会很大)
|
||||||
|
allowed, retryAfter := bucket.allow()
|
||||||
|
assert.False(t, allowed)
|
||||||
|
// rate=0 时,retryAfter 理论上无限大,实际会是一个很大的值
|
||||||
|
assert.Greater(t, retryAfter, 1*time.Hour)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryLimiter_Allow_DifferentKeys(t *testing.T) {
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewMemoryLimiter(cfg)
|
||||||
|
defer limiter.Stop()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// user1 消耗 2 个令牌
|
||||||
|
allowed, _ := limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
|
||||||
|
// user1 第 3 个被拒绝
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.False(t, allowed)
|
||||||
|
|
||||||
|
// user2 应该不受影响
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user2:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryLimiter_Cleanup(t *testing.T) {
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewMemoryLimiter(cfg)
|
||||||
|
defer limiter.Stop()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// 创建一个桶
|
||||||
|
limiter.Allow(ctx, "user1:query")
|
||||||
|
|
||||||
|
// 验证桶已创建
|
||||||
|
limiter.mu.RLock()
|
||||||
|
initialCount := len(limiter.buckets)
|
||||||
|
limiter.mu.RUnlock()
|
||||||
|
assert.Equal(t, 1, initialCount)
|
||||||
|
|
||||||
|
// 手动触发清理(模拟 10 分钟后)
|
||||||
|
limiter.mu.Lock()
|
||||||
|
for _, bucket := range limiter.buckets {
|
||||||
|
bucket.mu.Lock()
|
||||||
|
bucket.lastRefill = time.Now().Add(-11 * time.Minute)
|
||||||
|
bucket.mu.Unlock()
|
||||||
|
}
|
||||||
|
limiter.mu.Unlock()
|
||||||
|
|
||||||
|
limiter.removeInactiveBuckets()
|
||||||
|
|
||||||
|
// 验证桶已被清理
|
||||||
|
limiter.mu.RLock()
|
||||||
|
finalCount := len(limiter.buckets)
|
||||||
|
limiter.mu.RUnlock()
|
||||||
|
assert.Equal(t, 0, finalCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryLimiter_Stop(t *testing.T) {
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewMemoryLimiter(cfg)
|
||||||
|
|
||||||
|
// 多次调用 Stop 不应 panic
|
||||||
|
limiter.Stop()
|
||||||
|
limiter.Stop()
|
||||||
|
}
|
||||||
17
backend/internal/ratelimit/limiter.go
Normal file
17
backend/internal/ratelimit/limiter.go
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Limiter 速率限制器接口。
|
||||||
|
type Limiter interface {
|
||||||
|
// Allow 判断 key 是否允许执行一次操作。
|
||||||
|
// key 通常为 "userID:action" 格式。
|
||||||
|
// 返回 (allowed, retryAfter)。retryAfter 表示需要等待的时间。
|
||||||
|
Allow(ctx context.Context, key string) (bool, time.Duration)
|
||||||
|
|
||||||
|
// Stop 停止限流器,清理资源(如后台 goroutine)。
|
||||||
|
Stop()
|
||||||
|
}
|
||||||
42
backend/internal/ratelimit/middleware.go
Normal file
42
backend/internal/ratelimit/middleware.go
Normal file
@@ -0,0 +1,42 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Middleware 返回 Gin 中间件,按 key 维度限流。
|
||||||
|
// keyFunc 从请求中提取限流 key(如 IP、用户 ID)。
|
||||||
|
func Middleware(limiter Limiter, keyFunc func(*gin.Context) string) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
if limiter == nil {
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
key := keyFunc(c)
|
||||||
|
if key == "" {
|
||||||
|
// key 为空时跳过限流
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
allowed, retryAfter := limiter.Allow(c.Request.Context(), key)
|
||||||
|
|
||||||
|
if !allowed {
|
||||||
|
// 设置 Retry-After header(秒)
|
||||||
|
c.Header("Retry-After", fmt.Sprintf("%d", int(retryAfter.Seconds()+0.5)))
|
||||||
|
|
||||||
|
c.JSON(http.StatusTooManyRequests, gin.H{
|
||||||
|
"code": "RATE_LIMITED",
|
||||||
|
"message": fmt.Sprintf("too many requests, retry after %s", retryAfter.Round(1)),
|
||||||
|
})
|
||||||
|
c.Abort()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
196
backend/internal/ratelimit/middleware_test.go
Normal file
196
backend/internal/ratelimit/middleware_test.go
Normal file
@@ -0,0 +1,196 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// mockLimiter 用于测试的 mock 限流器。
|
||||||
|
type mockLimiter struct {
|
||||||
|
allowFunc func(ctx context.Context, key string) (bool, time.Duration)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
if m.allowFunc != nil {
|
||||||
|
return m.allowFunc(ctx, key)
|
||||||
|
}
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockLimiter) Stop() {}
|
||||||
|
|
||||||
|
// 编译期接口检查
|
||||||
|
var _ Limiter = (*mockLimiter)(nil)
|
||||||
|
|
||||||
|
func TestMiddleware_Allow(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
limiter := &mockLimiter{
|
||||||
|
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
return true, 0
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return "user1:test"
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
assert.Equal(t, http.StatusOK, w.Code)
|
||||||
|
|
||||||
|
var resp map[string]interface{}
|
||||||
|
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "ok", resp["status"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMiddleware_Deny(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
limiter := &mockLimiter{
|
||||||
|
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
return false, 5 * time.Second
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return "user1:test"
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
// 验证返回 429
|
||||||
|
assert.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||||
|
|
||||||
|
// 验证 Retry-After header
|
||||||
|
assert.Equal(t, "5", w.Header().Get("Retry-After"))
|
||||||
|
|
||||||
|
// 验证响应体
|
||||||
|
var resp map[string]interface{}
|
||||||
|
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "RATE_LIMITED", resp["code"])
|
||||||
|
assert.Contains(t, resp["message"], "retry after")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMiddleware_NilLimiter(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(nil, func(c *gin.Context) string {
|
||||||
|
return "user1:test"
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
// nil limiter 应该放行
|
||||||
|
assert.Equal(t, http.StatusOK, w.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMiddleware_EmptyKey(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
limiter := &mockLimiter{
|
||||||
|
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
// 不应该被调用
|
||||||
|
t.Error("Allow should not be called with empty key")
|
||||||
|
return false, 0
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return "" // 返回空 key
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
// 空 key 应该放行
|
||||||
|
assert.Equal(t, http.StatusOK, w.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMiddleware_KeyFunc(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
var capturedKey string
|
||||||
|
limiter := &mockLimiter{
|
||||||
|
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
capturedKey = key
|
||||||
|
return true, 0
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
// 从 query 参数提取 user_id
|
||||||
|
userID := c.Query("user_id")
|
||||||
|
return userID + ":test"
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test?user_id=user123", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
assert.Equal(t, http.StatusOK, w.Code)
|
||||||
|
assert.Equal(t, "user123:test", capturedKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMiddleware_RetryAfterRounding(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
limiter := &mockLimiter{
|
||||||
|
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
return false, 2500 * time.Millisecond // 2.5 秒
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return "user1:test"
|
||||||
|
}))
|
||||||
|
router.GET("/test", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
assert.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||||
|
// 2.5 秒向上取整为 3 秒
|
||||||
|
assert.Equal(t, "3", w.Header().Get("Retry-After"))
|
||||||
|
}
|
||||||
128
backend/internal/ratelimit/redis_bucket.go
Normal file
128
backend/internal/ratelimit/redis_bucket.go
Normal file
@@ -0,0 +1,128 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
)
|
||||||
|
|
||||||
|
// luaScript 是 Redis 令牌桶算法的 Lua 脚本。
|
||||||
|
// 保证原子性:读取-计算-回写在一个事务中完成。
|
||||||
|
const luaScript = `
|
||||||
|
-- KEYS[1] = 限流 key
|
||||||
|
-- ARGV[1] = capacity(桶容量)
|
||||||
|
-- ARGV[2] = rate(每秒填充数)
|
||||||
|
-- ARGV[3] = now(当前时间戳,秒,浮点)
|
||||||
|
-- ARGV[4] = ttl(key 过期时间,秒)
|
||||||
|
|
||||||
|
local key = KEYS[1]
|
||||||
|
local capacity = tonumber(ARGV[1])
|
||||||
|
local rate = tonumber(ARGV[2])
|
||||||
|
local now = tonumber(ARGV[3])
|
||||||
|
local ttl = tonumber(ARGV[4])
|
||||||
|
|
||||||
|
local data = redis.call('HMGET', key, 'tokens', 'last_refill')
|
||||||
|
local tokens = tonumber(data[1]) or capacity
|
||||||
|
local last_refill = tonumber(data[2]) or now
|
||||||
|
|
||||||
|
-- 计算新令牌
|
||||||
|
local elapsed = math.max(0, now - last_refill)
|
||||||
|
tokens = math.min(capacity, tokens + elapsed * rate)
|
||||||
|
|
||||||
|
local allowed = 0
|
||||||
|
local retry_after = 0
|
||||||
|
|
||||||
|
if tokens >= 1 then
|
||||||
|
tokens = tokens - 1
|
||||||
|
allowed = 1
|
||||||
|
else
|
||||||
|
if rate == 0 then
|
||||||
|
retry_after = 86400 -- 24小时
|
||||||
|
else
|
||||||
|
retry_after = (1 - tokens) / rate
|
||||||
|
end
|
||||||
|
end
|
||||||
|
|
||||||
|
-- 回写状态
|
||||||
|
redis.call('HMSET', key, 'tokens', tokens, 'last_refill', now)
|
||||||
|
redis.call('EXPIRE', key, ttl)
|
||||||
|
|
||||||
|
return {allowed, tostring(retry_after)}
|
||||||
|
`
|
||||||
|
|
||||||
|
// RedisLimiter Redis 令牌桶限流器。
|
||||||
|
type RedisLimiter struct {
|
||||||
|
client *redis.Client
|
||||||
|
config config.RateLimitConfig
|
||||||
|
script *redis.Script
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRedisLimiter 创建 Redis 限流器。
|
||||||
|
func NewRedisLimiter(client *redis.Client, cfg config.RateLimitConfig) *RedisLimiter {
|
||||||
|
return &RedisLimiter{
|
||||||
|
client: client,
|
||||||
|
config: cfg,
|
||||||
|
script: redis.NewScript(luaScript),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow 实现 Limiter 接口。
|
||||||
|
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
cfg := l.getBucketConfig(key)
|
||||||
|
|
||||||
|
now := float64(time.Now().UnixNano()) / 1e9 // 秒,浮点
|
||||||
|
ttl := 600 // key 过期时间 10 分钟
|
||||||
|
|
||||||
|
result, err := l.script.Run(ctx, l.client, []string{key},
|
||||||
|
cfg.Capacity, cfg.Rate, now, ttl).Result()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
// Redis 错误时降级:允许请求(fail-open 策略)
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// 解析返回值
|
||||||
|
vals, ok := result.([]interface{})
|
||||||
|
if !ok || len(vals) != 2 {
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
allowed, _ := vals[0].(int64)
|
||||||
|
retryAfterStr, _ := vals[1].(string)
|
||||||
|
retryAfterSec, _ := strconv.ParseFloat(retryAfterStr, 64)
|
||||||
|
|
||||||
|
if allowed == 1 {
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
retryAfter := time.Duration(retryAfterSec*1000) * time.Millisecond
|
||||||
|
return false, retryAfter
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop 实现 Limiter 接口(Redis 不需要清理资源)。
|
||||||
|
func (l *RedisLimiter) Stop() {
|
||||||
|
// Redis 客户端由外部管理,这里不需要操作
|
||||||
|
}
|
||||||
|
|
||||||
|
// getBucketConfig 根据 key 获取桶配置。
|
||||||
|
func (l *RedisLimiter) getBucketConfig(key string) config.BucketConfig {
|
||||||
|
// 简化实现:默认使用 query 配置
|
||||||
|
return l.config.Query
|
||||||
|
}
|
||||||
|
|
||||||
|
// KeyPrefix 返回限流 key 的前缀。
|
||||||
|
func KeyPrefix() string {
|
||||||
|
return "ratelimit:"
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatKey 格式化限流 key。
|
||||||
|
func FormatKey(userID, action string) string {
|
||||||
|
return fmt.Sprintf("%s%s:%s", KeyPrefix(), userID, action)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 编译期接口检查
|
||||||
|
var _ Limiter = (*RedisLimiter)(nil)
|
||||||
228
backend/internal/ratelimit/redis_bucket_test.go
Normal file
228
backend/internal/ratelimit/redis_bucket_test.go
Normal file
@@ -0,0 +1,228 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/alicebob/miniredis/v2"
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setupMiniRedis 创建一个内存 Redis 实例用于测试。
|
||||||
|
func setupMiniRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) {
|
||||||
|
mr, err := miniredis.Run()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
client := redis.NewClient(&redis.Options{
|
||||||
|
Addr: mr.Addr(),
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
client.Close()
|
||||||
|
mr.Close()
|
||||||
|
})
|
||||||
|
|
||||||
|
return mr, client
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_FirstRequest(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 5, Rate: 0.2},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
|
||||||
|
|
||||||
|
assert.True(t, allowed)
|
||||||
|
assert.Equal(t, time.Duration(0), retryAfter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_ConsumeUntilEmpty(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 3, Rate: 0.2},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "user1:query"
|
||||||
|
|
||||||
|
// 连续消耗 3 个令牌
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
allowed, _ := limiter.Allow(ctx, key)
|
||||||
|
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 第 4 个请求应被拒绝
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||||
|
assert.False(t, allowed)
|
||||||
|
assert.Greater(t, retryAfter, time.Duration(0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_DifferentKeys(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// user1 消耗 2 个令牌
|
||||||
|
allowed, _ := limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
|
||||||
|
// user1 第 3 个被拒绝
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.False(t, allowed)
|
||||||
|
|
||||||
|
// user2 应该不受影响
|
||||||
|
allowed, _ = limiter.Allow(ctx, "user2:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_RefillAfterWait(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 2, Rate: 10.0}, // 每秒 10 个令牌
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "user1:query"
|
||||||
|
|
||||||
|
// 消耗 2 个令牌
|
||||||
|
limiter.Allow(ctx, key)
|
||||||
|
limiter.Allow(ctx, key)
|
||||||
|
|
||||||
|
// 真实等待 150ms(Lua 脚本使用系统时间)
|
||||||
|
time.Sleep(150 * time.Millisecond)
|
||||||
|
|
||||||
|
// 应该补充了至少 1 个令牌
|
||||||
|
allowed, _ := limiter.Allow(ctx, key)
|
||||||
|
assert.True(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_CapacityLimit(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 3, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "user1:query"
|
||||||
|
|
||||||
|
// 真实等待让桶"溢出"
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
|
||||||
|
// 但最多只能消耗 capacity 个令牌
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
allowed, _ := limiter.Allow(ctx, key)
|
||||||
|
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 第 4 个应被拒绝
|
||||||
|
allowed, _ := limiter.Allow(ctx, key)
|
||||||
|
assert.False(t, allowed)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_ZeroRate(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 1, Rate: 0.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "user1:query"
|
||||||
|
|
||||||
|
// 第一个通过
|
||||||
|
allowed, _ := limiter.Allow(ctx, key)
|
||||||
|
assert.True(t, allowed)
|
||||||
|
|
||||||
|
// 第二个被拒绝,retryAfter 应该很大
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||||
|
assert.False(t, allowed)
|
||||||
|
assert.Greater(t, retryAfter, 1*time.Hour)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_KeyTTL(t *testing.T) {
|
||||||
|
mr, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 5, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
key := "user1:query"
|
||||||
|
|
||||||
|
// 第一次请求
|
||||||
|
limiter.Allow(ctx, key)
|
||||||
|
|
||||||
|
// 验证 key 已设置 TTL
|
||||||
|
ttl := mr.TTL(key)
|
||||||
|
assert.Greater(t, ttl, time.Duration(0))
|
||||||
|
assert.LessOrEqual(t, ttl, 600*time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Allow_FailOpen(t *testing.T) {
|
||||||
|
mr, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// 关闭 Redis 模拟故障
|
||||||
|
mr.Close()
|
||||||
|
|
||||||
|
// 应该 fail-open(允许请求)
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
|
||||||
|
assert.True(t, allowed)
|
||||||
|
assert.Equal(t, time.Duration(0), retryAfter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRedisLimiter_Stop(t *testing.T) {
|
||||||
|
_, client := setupMiniRedis(t)
|
||||||
|
|
||||||
|
cfg := config.RateLimitConfig{
|
||||||
|
Enabled: true,
|
||||||
|
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||||
|
}
|
||||||
|
limiter := NewRedisLimiter(client, cfg)
|
||||||
|
|
||||||
|
// Stop 应该不会 panic(即使多次调用)
|
||||||
|
limiter.Stop()
|
||||||
|
limiter.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatKey(t *testing.T) {
|
||||||
|
key := FormatKey("user123", "query")
|
||||||
|
assert.Equal(t, "ratelimit:user123:query", key)
|
||||||
|
}
|
||||||
@@ -3,6 +3,7 @@ package ws
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -17,6 +18,7 @@ import (
|
|||||||
"github.com/hhs/camtalk/internal/logger"
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
"github.com/hhs/camtalk/internal/orchestrator"
|
"github.com/hhs/camtalk/internal/orchestrator"
|
||||||
|
"github.com/hhs/camtalk/internal/ratelimit"
|
||||||
"github.com/hhs/camtalk/internal/session"
|
"github.com/hhs/camtalk/internal/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -93,19 +95,19 @@ func (w *WSClient) SendError(err models.WsError) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ServeWS 处理 WebSocket 升级请求。
|
// ServeWS 处理 WebSocket 升级请求。
|
||||||
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config, tokenMgr *auth.TokenManager) gin.HandlerFunc {
|
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter) gin.HandlerFunc {
|
||||||
upgrader := newUpgrader(cfg)
|
upgrader := newUpgrader(cfg)
|
||||||
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
|
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
|
||||||
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
|
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
|
||||||
version := cfg.App.Version
|
version := cfg.App.Version
|
||||||
|
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr)
|
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr, limiter)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
|
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
|
||||||
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager) {
|
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter) {
|
||||||
|
|
||||||
// --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) ---
|
// --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) ---
|
||||||
token := c.Query("token")
|
token := c.Query("token")
|
||||||
@@ -225,6 +227,18 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
}
|
}
|
||||||
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
||||||
|
|
||||||
|
// 限流检查
|
||||||
|
if limiter != nil {
|
||||||
|
key := fmt.Sprintf("%s:query", userID)
|
||||||
|
allowed, retryAfter := limiter.Allow(context.Background(), key)
|
||||||
|
if !allowed {
|
||||||
|
logger.Log.Warnw("rate limited", "user_id", userID, "retry_after", retryAfter)
|
||||||
|
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
|
||||||
|
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 刷新会话 TTL
|
// 刷新会话 TTL
|
||||||
if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil {
|
if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil {
|
||||||
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
|
||||||
|
|||||||
@@ -148,7 +148,7 @@ func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Se
|
|||||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||||
Session: config.SessionConfig{MaxHistory: 20},
|
Session: config.SessionConfig{MaxHistory: 20},
|
||||||
}
|
}
|
||||||
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr, nil))
|
||||||
|
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
|
|
||||||
@@ -591,7 +591,7 @@ func setupTestServerEx(t *testing.T, orch orchestrator.Orchestrator) (*httptest.
|
|||||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||||
Session: config.SessionConfig{MaxHistory: 20},
|
Session: config.SessionConfig{MaxHistory: 20},
|
||||||
}
|
}
|
||||||
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr, nil))
|
||||||
|
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
return srv, tokenMgr, sessionMgr
|
return srv, tokenMgr, sessionMgr
|
||||||
@@ -642,7 +642,7 @@ func TestWS_AuthExpiredToken(t *testing.T) {
|
|||||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||||
Session: config.SessionConfig{MaxHistory: 20},
|
Session: config.SessionConfig{MaxHistory: 20},
|
||||||
}
|
}
|
||||||
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr))
|
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr, nil))
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
defer srv.Close()
|
defer srv.Close()
|
||||||
|
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ services:
|
|||||||
env_file:
|
env_file:
|
||||||
- /opt/camtalk/.env
|
- /opt/camtalk/.env
|
||||||
environment:
|
environment:
|
||||||
|
# 运行环境(强制生产环境)
|
||||||
|
- APP_ENV=prod
|
||||||
# 三级存储配置(敏感信息通过 env_file 注入)
|
# 三级存储配置(敏感信息通过 env_file 注入)
|
||||||
- CAMTALK_STORAGE_REDIS_ENABLED=${CAMTALK_STORAGE_REDIS_ENABLED:-true}
|
- CAMTALK_STORAGE_REDIS_ENABLED=${CAMTALK_STORAGE_REDIS_ENABLED:-true}
|
||||||
- CAMTALK_STORAGE_PERSISTENCE_ENABLED=${CAMTALK_STORAGE_PERSISTENCE_ENABLED:-true}
|
- CAMTALK_STORAGE_PERSISTENCE_ENABLED=${CAMTALK_STORAGE_PERSISTENCE_ENABLED:-true}
|
||||||
|
|||||||
@@ -426,7 +426,7 @@ ws.onerror = (error) => {
|
|||||||
### 3. 传输安全
|
### 3. 传输安全
|
||||||
|
|
||||||
- **HTTPS 强制**:生产环境必须使用 HTTPS
|
- **HTTPS 强制**:生产环境必须使用 HTTPS
|
||||||
- **CORS 限制**:配置 `AllowedOrigins` 限制允许的域名
|
- **同源反代**:通过 Nginx 反向代理(生产)或 Vite proxy(开发)统一前后端到同一域名,浏览器层面无跨域问题
|
||||||
- **HttpOnly Cookie**:refresh_token 存储在 httpOnly Cookie 中,防止 XSS 攻击
|
- **HttpOnly Cookie**:refresh_token 存储在 httpOnly Cookie 中,防止 XSS 攻击
|
||||||
|
|
||||||
### 4. 防攻击策略
|
### 4. 防攻击策略
|
||||||
|
|||||||
Reference in New Issue
Block a user