feat: 集成限流器到服务 #178

Merged
huanghaosheng merged 12 commits from develop into v2 2026-06-21 13:25:06 +08:00
23 changed files with 1449 additions and 58 deletions

232
CLAUDE.md
View File

@@ -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-plusLLM、MiMo ASRSTT、MiMo TTSTTS。仅通过 Go 网关访问,浏览器不直连。 3. **云端 AI 服务** —— 通过 OpenAI 兼容接口可灵活切换。默认DashScope qwen3-vl-plusLLM、MiMo ASRSTT、MiMo TTSTTS。仅通过 Go 网关访问,浏览器不直连。
**关键模式**AI 编排基于 Eino Graph 声明式 DAG`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`LLM token 通过 Callback 实时推送TTS 逐句合成并行推送,最小化感知延迟。 **关键模式**AI 编排基于 Eino Graph 声明式 DAG6 节点线性流水线:`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`LLM token 通过 Callback 实时推送TTS 逐句合成并行推送,最小化感知延迟。
**存储**三级存储架构TieredManager—— L1 Memory → L2 Redis → L3 PostgreSQL自动降级。Repository 接口模式UserRepository、MessageRepository、SessionRepositoryPostgreSQL + 内存双实现。 ### 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-plusOpenAI 兼容协议) |
| 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 → 回填 L1L3 通过 L1 的 `FindByID` 降级读取
- **写路径**L1 同步写入 → L2 同步写(失败 soft-warn→ L3 异步 goroutine 写(使用 `context.Background()` 防止请求取消丢失)
- **降级**:后台协程每 30 秒 ping RedisRedis 不可用时自动跳过 L2 操作;恢复后自动重新启用
- **TTL**Session 默认 30 分钟MaxHistory 20 条L1 后台协程每分钟清理过期 session
Repository 接口模式:`UserRepository``MessageRepository``SessionRepository`,均有 PostgreSQL 和内存双实现。
### 鉴权
JWT 双 token 轮转认证HMAC-SHA256
- Access Token默认 120 分钟 TTLBearer header 传递
- Refresh Token默认 7 天 TTL带 jtiUUIDHash 存储在 Redis/PostgreSQL
- 轮转Refresh 时旧 token hash 删除,新 pair 生成;若 JWT 有效但 DB hash 缺失 → 判定为重放攻击 → 吊销该用户所有 refresh token
- `CachedUserRepository``backend/internal/store/cached_user.go`装饰器模式Redis 缓存 refresh token hashRead-Through / Write-ThroughRedis 故障软降级
## 技术栈 ## 技术栈
| 层级 | 技术 | | 层级 | 技术 |
|------|------| |------|------|
| 前端 | 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` | 启用 Redistrue/false |
| `CAMTALK_STORAGE_PERSISTENCE_ENABLED` | 启用 PostgreSQLtrue/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` hook640x480, facingMode: environment |
| `MicManager` | 麦克风音频采集 | | `MicManager` | 麦克风音频采集`useMicrophone` hook16kHz 单声道) |
| `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/info3 秒自动消失)
- `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/PostgreSQL30 分钟 TTL | | Session Manager | 会话状态、对话历史三级存储Memory/Redis/PostgreSQL30 分钟 TTL |
| Eino 编排层 | 基于 Eino Graph 的声明式 AI 编排(7 节点 DAGStream 模式,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 多 providerLLM 通过 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/SessionRepositoryPG + 内存 + 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/` 下的接口文档。

View File

@@ -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
View File

@@ -4,8 +4,6 @@ bin/
# 环境配置 # 环境配置
.env .env
config.dev.yaml
config.prod.yaml
# 临时文件 # 临时文件
tmp/ tmp/

View File

@@ -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{

View 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 # 控制台格式,易读

View 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 格式,便于日志收集和分析

View File

@@ -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

View File

@@ -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

View File

@@ -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=

View File

@@ -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)
} }

View File

@@ -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
} }

View File

@@ -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 过期时间(分钟),默认 100807天 RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 100807天
} }
// 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 显式绑定敏感信息环境变量。

View 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)

View 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()
}

View 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()
}

View 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()
}
}

View 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"))
}

View 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] = ttlkey 过期时间,秒)
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)

View 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)
// 真实等待 150msLua 脚本使用系统时间)
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)
}

View File

@@ -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)

View File

@@ -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()

View File

@@ -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}

View File

@@ -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. 防攻击策略