Compare commits
223 Commits
943f36e9ec
...
develop
| Author | SHA1 | Date | |
|---|---|---|---|
| b1c3958e9b | |||
| f9a01269d2 | |||
| 74f1a8eb34 | |||
| 4dd4138197 | |||
| 5ef059eb40 | |||
| 9523fe6650 | |||
| a219934f4b | |||
| 72a8d4803e | |||
| dcf53a3783 | |||
| b5ec6551bf | |||
| 79d321c78a | |||
| 57bd8c72b8 | |||
| d640b5b41b | |||
| f1b9e4fcc1 | |||
| 18aa5f1949 | |||
| a16628119a | |||
| 3888941f85 | |||
| 0a816049fa | |||
| 140a178993 | |||
| 31c1720c00 | |||
| c6093c01c1 | |||
| 3267f35edb | |||
| a309c269e0 | |||
| 22930f0080 | |||
| c00b8a83d9 | |||
| 36007babcb | |||
| f01c84d89e | |||
| c3808ff38c | |||
| 2740bea755 | |||
| 69b9694cfe | |||
| 70ee212ea2 | |||
| f08f547e76 | |||
| 03a19befe1 | |||
| c4b68b77d1 | |||
| f2a883bda8 | |||
| 4cc0713459 | |||
| 2180751a9b | |||
| 964c5c967e | |||
| dcd08d031c | |||
| 6bca4e0d47 | |||
| 881c3f9853 | |||
| 2881941be6 | |||
| 1b99d24fc1 | |||
| f9e68b96b9 | |||
| c58c6b59a5 | |||
| e24fb2cc19 | |||
| 7a705744d0 | |||
| 583c33727a | |||
| 6cbabb63bb | |||
| b4fbf8625b | |||
| 239f8f9877 | |||
| d430e6e5b2 | |||
| edc66625ba | |||
| 9dce107a84 | |||
| 1c7dd708a0 | |||
| d0e4bdaeec | |||
| 949a707e0f | |||
| 55f7f183a7 | |||
| 6b4b033df3 | |||
| 76d331c885 | |||
| ad700743ef | |||
| 6ab4776e08 | |||
| c17798ec67 | |||
| 8a43f4406a | |||
| 065673fae2 | |||
| 03c27e7790 | |||
| 0a59173476 | |||
| 939e43acd0 | |||
| 87c3e7a8dd | |||
| d6e9555a97 | |||
| 9c763ec12a | |||
| 1bab02ae84 | |||
| 89d7b7c17c | |||
| 34d498510e | |||
| 104b28efd3 | |||
| a02a8bc374 | |||
| 0d99d06f06 | |||
| 8b18953010 | |||
| 0108ef2064 | |||
| e0dc8272a5 | |||
| 11c3955bd6 | |||
| a66ab764d9 | |||
| 51117b43f6 | |||
| 492fb06c08 | |||
| 6967dd7b2e | |||
| eea5c07eaa | |||
| 1079e22699 | |||
| ad5d90e344 | |||
| c094fe0867 | |||
| 032de796c8 | |||
| 8b4acb3ce7 | |||
| 9e5f691056 | |||
| 03127aa01a | |||
| 99fcd6bc29 | |||
| d0f5f5c94d | |||
| 361c5d07d3 | |||
| 910e71b6f0 | |||
| 19645be04e | |||
| d53265755a | |||
| 252cdcc8e7 | |||
| 515d7ae034 | |||
| 311330cea1 | |||
| 7adf81c6e5 | |||
| ea00939c13 | |||
| b74fb3564d | |||
| 3b6226394b | |||
| a6df8c9131 | |||
| 023c834074 | |||
| 3e00e39e8a | |||
| 9ad486d117 | |||
| 6af26ffc91 | |||
| 7ee6918015 | |||
| 2fb23852ef | |||
| ed95ce56a8 | |||
| e151c5b665 | |||
| dffd8bd4a5 | |||
| 3140a660e4 | |||
| 23e0e22a12 | |||
| ba7c5ed5ea | |||
| cc7a333b6d | |||
| f8af2f0ccc | |||
| d6059ae397 | |||
| d4398e55fe | |||
| f5440c9e2d | |||
| 190f40908f | |||
| ddb90c58ef | |||
| 340c26b7a6 | |||
| ca07188eda | |||
| 15a50f855b | |||
| e37dc7074d | |||
| f86c2560cd | |||
| 90c5ad7724 | |||
| 45e2c4f37b | |||
| b7d0edb6da | |||
| 159bf278e6 | |||
| 542126df90 | |||
| e56600e408 | |||
| 04215e7a53 | |||
| 9ba4eb4825 | |||
| bc7eb7409c | |||
| 0a8a20f3f8 | |||
| e20f58b4ae | |||
| ab07e01adf | |||
| 0447fdacac | |||
| 12cf59f1e8 | |||
| 4f735e6d29 | |||
| 32aea44f3b | |||
| 4d51bcefb0 | |||
| fe5ac20a1f | |||
| 2a13ee9c89 | |||
| df35ff73b5 | |||
| 92e6636d05 | |||
| 12e2d24f37 | |||
| ae100c4a75 | |||
| 4a5905307c | |||
| 2a7d4c74d4 | |||
| 6bb4773e07 | |||
| 97e125234b | |||
| 39073b7673 | |||
| 85d47c1fc4 | |||
| 4ff8cec312 | |||
| 3572b867c0 | |||
| dfed964f76 | |||
| 4651ae185b | |||
| 1c65433d40 | |||
| 590c4592ea | |||
| 15cd157f45 | |||
| 970f10a274 | |||
| d78cdb509d | |||
| ea70d2efc6 | |||
| a2a28a9f56 | |||
| 898e30b526 | |||
| 87c54b80c0 | |||
| 7288f443f1 | |||
| 488b1e62ba | |||
| 50c84fce88 | |||
| ae09fba400 | |||
| f8a79a3b0b | |||
| ae9a27c300 | |||
| 1f7cd407a8 | |||
| 6f81212997 | |||
| 402ad8c949 | |||
| 52015fa6c6 | |||
| 97706ea197 | |||
| d53de4f33f | |||
| 3dc2015a91 | |||
| de78d60959 | |||
| 93d5a90495 | |||
| 5910e02a66 | |||
| 88ee1e2548 | |||
| b8fd6a9330 | |||
| c2c9a6d392 | |||
| 795ace75d6 | |||
| dd846876ba | |||
| 898d1474ec | |||
| 582f68f68b | |||
| 7d328f2552 | |||
| eb1b90445f | |||
| e5537aaa1e | |||
| 556666c046 | |||
| 4235c5cae0 | |||
| 41f393c09d | |||
| 04dac0b673 | |||
| 4c7430d6f4 | |||
| 765cb34019 | |||
| 9576884619 | |||
| 4ffd84510e | |||
| 4b731b5ac0 | |||
| fd5c7712f8 | |||
| f38fbf0527 | |||
| b10a508356 | |||
| c0b4eeda46 | |||
| e6481e0faa | |||
| a04275cc76 | |||
| dca37f3e48 | |||
| 16302af7d2 | |||
| 5d8cacf16d | |||
| dbfdf3c3e5 | |||
| 54454ae2d7 | |||
| 720c2e2b5b | |||
| c2bae4e3b7 | |||
| 838493145f | |||
| 14242eb896 |
21
.env.example
21
.env.example
@@ -1,21 +0,0 @@
|
|||||||
# CamTalk 环境变量模板
|
|
||||||
# 复制为 .env 并填入实际值:cp .env.example .env
|
|
||||||
# .env 已在 .gitignore 中,不会提交到版本控制
|
|
||||||
|
|
||||||
# ---- AI 服务 API Key ----
|
|
||||||
CAMTALK_AI_LLM_API_KEY=sk-xxx
|
|
||||||
CAMTALK_AI_STT_API_KEY=
|
|
||||||
CAMTALK_AI_TTS_API_KEY=
|
|
||||||
|
|
||||||
# ---- 可选覆盖(默认值见 config.yaml)----
|
|
||||||
# CAMTALK_AI_LLM_MODEL=gpt-4o
|
|
||||||
# CAMTALK_AI_LLM_ENDPOINT=https://api.openai.com/v1
|
|
||||||
# CAMTALK_AI_LLM_TIMEOUT=10
|
|
||||||
# CAMTALK_AI_STT_ENDPOINT=https://api.xiaomimimo.com/v1
|
|
||||||
# CAMTALK_AI_TTS_ENDPOINT=https://api.openai.com/v1
|
|
||||||
# CAMTALK_AI_TTS_VOICE=alloy
|
|
||||||
# CAMTALK_AI_TTS_SPEED=1.0
|
|
||||||
# CAMTALK_AI_TTS_TIMEOUT=5
|
|
||||||
|
|
||||||
# ---- 应用 ----
|
|
||||||
# APP_ENV=dev
|
|
||||||
@@ -2,20 +2,34 @@ name: Deploy
|
|||||||
|
|
||||||
on:
|
on:
|
||||||
push:
|
push:
|
||||||
branches: [main]
|
branches: [main, v2]
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
deploy:
|
deploy:
|
||||||
runs-on: aliyun
|
runs-on: aliyun
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout
|
- name: Deploy
|
||||||
uses: "http://8.161.227.145:3000/huanghaosheng/checkout@releases/v4"
|
|
||||||
|
|
||||||
- name: Install Docker CLI
|
|
||||||
run: apk add --no-cache docker-cli docker-cli-compose
|
|
||||||
|
|
||||||
- name: Build and Deploy
|
|
||||||
run: |
|
run: |
|
||||||
chmod +x deploy.sh
|
sed -i 's/dl-cdn.alpinelinux.org/mirrors.aliyun.com/g' /etc/apk/repositories
|
||||||
./deploy.sh build
|
apk add --no-cache rsync docker-cli docker-cli-compose
|
||||||
./deploy.sh restart
|
GIT_URL="http://8.161.227.145:3000/XEngineers/CamTalk.git"
|
||||||
|
if [ -d /root/camtalk/.git ]; then
|
||||||
|
cd /root/camtalk
|
||||||
|
git fetch "$GIT_URL" ${GITHUB_REF_NAME} --depth=1
|
||||||
|
git reset --hard FETCH_HEAD
|
||||||
|
else
|
||||||
|
rm -rf /tmp/camtalk-deploy
|
||||||
|
git clone --depth=1 --branch ${GITHUB_REF_NAME} \
|
||||||
|
http://8.161.227.145:3000/XEngineers/CamTalk.git /tmp/camtalk-deploy
|
||||||
|
mkdir -p /root/camtalk
|
||||||
|
rsync -a --delete \
|
||||||
|
--exclude='.env' \
|
||||||
|
--exclude='pgdata' \
|
||||||
|
--exclude='redisdata' \
|
||||||
|
/tmp/camtalk-deploy/ /root/camtalk/
|
||||||
|
rm -rf /tmp/camtalk-deploy
|
||||||
|
fi
|
||||||
|
|
||||||
|
chmod +x /root/camtalk/deploy.sh
|
||||||
|
/root/camtalk/deploy.sh build
|
||||||
|
/root/camtalk/deploy.sh restart
|
||||||
7
.gitignore
vendored
7
.gitignore
vendored
@@ -22,8 +22,9 @@ Thumbs.db
|
|||||||
# ---- Playwright MCP ----
|
# ---- Playwright MCP ----
|
||||||
.playwright-mcp/
|
.playwright-mcp/
|
||||||
|
|
||||||
# ---- 截图 ----
|
|
||||||
*.png
|
|
||||||
|
|
||||||
# ---- Obsidian ----
|
# ---- Obsidian ----
|
||||||
.obsidian/
|
.obsidian/
|
||||||
|
.claudian/
|
||||||
|
修改过程笔记/
|
||||||
|
学习复盘/
|
||||||
|
docs/follow-up/
|
||||||
|
|||||||
178
CLAUDE.md
178
CLAUDE.md
@@ -1,106 +1,128 @@
|
|||||||
# CLAUDE.md
|
# CLAUDE.md
|
||||||
|
|
||||||
本文件为 Claude Code (claude.ai/code) 在本仓库中工作时提供指引。
|
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
||||||
|
|
||||||
## 项目概述
|
CamTalk — 多模态实时 AI 视觉对话助手(摄像头 + 麦克风 + 视觉 + 语音 AI)
|
||||||
|
|
||||||
CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。项目目前处于设计文档阶段,源代码正在逐步构建。
|
> **文档优先原则:** 开发前先读 `docs/` 设计文档,以文档为准;若代码与文档不一致,优先更新文档(尤其接口文档)。详细设计见 `docs/01-13` 系列文档。**注意**:`docs/Eino/` 框架文档内容庞大(~75 个文件),仅在需要了解 Eino Graph/节点/Callback 等框架细节时才读取。
|
||||||
|
|
||||||
> **文档优先原则:** 执行任何开发任务前,先读取 `docs/` 下的相关设计文档(架构、接口、技术选型等),以文档为最高依据。代码实现应与文档一致;若有偏差,优先更新文档(尤其是接口文档)。
|
## 常用命令
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# === 前端(frontend/ 目录)===
|
||||||
|
npm run dev # Vite 开发服务器(http://localhost:5173,代理 /ws 和 /api 到 :8080)
|
||||||
|
npm run build # 生产构建(tsc -b && vite build,输出到 dist/)
|
||||||
|
npm run lint # ESLint 代码检查
|
||||||
|
npm run preview # 预览生产构建
|
||||||
|
|
||||||
|
# === 后端(backend/ 目录)===
|
||||||
|
go run ./cmd/server # 启动服务(监听 :8080,启动时自动执行数据库迁移)
|
||||||
|
golangci-lint run # Go 代码检查
|
||||||
|
|
||||||
|
# 后端测试
|
||||||
|
go test ./... # 单元测试
|
||||||
|
go test -tags=integration ./... # 集成测试(需要 PostgreSQL)
|
||||||
|
go test -v -run TestXxx ./path/ # 运行单个测试
|
||||||
|
|
||||||
|
# === Docker 部署 ===
|
||||||
|
./deploy.sh build # 构建 Docker 镜像
|
||||||
|
./deploy.sh up # 启动服务(4 容器:frontend/backend/postgres/redis)
|
||||||
|
./deploy.sh down # 停止服务
|
||||||
|
./deploy.sh logs # 查看日志(可加服务名:./deploy.sh logs backend)
|
||||||
|
./deploy.sh status # 查看服务状态
|
||||||
|
```
|
||||||
|
|
||||||
## 架构
|
## 架构
|
||||||
|
|
||||||
三层系统:
|
三层系统:前端(React + Vite)→ Go 网关(Gin + WebSocket + Eino Graph AI 编排)→ AI 服务(DashScope LLM, MiMo STT/TTS)
|
||||||
|
|
||||||
1. **浏览器客户端**(React 18 + TypeScript, Vite)—— 媒体采集、边缘预处理(VAD 通过 `@ricky0123/vad-web`、关键帧检测通过 Canvas 像素比较)、UI 渲染。核心 Hook:`useVisionSession()`
|
**AI 编排流水线**(Eino Graph 7 节点 DAG):`STT → History → ChatModel → Msg2Str → Splitter → TTS → Done`。LLM token 通过 Callback 实时推送,TTS 逐句并行合成。
|
||||||
2. **Go 网关**(Gin, gorilla/websocket, Viper, Zap)—— WebSocket 服务器、会话管理、AI 编排。每个 WebSocket 连接一个 goroutine。
|
|
||||||
3. **云端 AI 服务** —— 通过 OpenAI 兼容接口可灵活切换。默认:GPT-4o(LLM)、Deepgram(STT)、OpenAI TTS。仅通过 Go 网关访问,浏览器不直连。
|
|
||||||
|
|
||||||
**关键模式**:LLM 文本流和 TTS 音频流并行推送给客户端,以最小化感知延迟。
|
**会话存储**(TieredManager):L1 Memory → L2 Redis → L3 PostgreSQL 三级存储,30 分钟 TTL,Redis 故障自动降级。
|
||||||
|
|
||||||
**存储**:MVP 阶段使用进程内存(`MemoryManager`),Redis 实现已就绪可通过配置切换,PostgreSQL 为规划中。Repository 接口模式(`HistoryRepository`、`UsageRepository`),MVP 用内存实现。
|
**鉴权**:JWT 双 token 轮转(Access 120min + Refresh 7d),重放攻击检测(DB hash 校验),Redis 缓存装饰器。
|
||||||
|
|
||||||
## 技术栈
|
## 技术栈
|
||||||
|
|
||||||
| 层级 | 技术 |
|
前端:React 18 + TypeScript + Vite,VAD(@ricky0123/vad-web),ONNX Runtime,国际化(zh-CN / en-US / ja-JP)
|
||||||
|
后端:Go 1.25+, Gin, WebSocket, Viper, Zap, CloudWeGo Eino Graph
|
||||||
|
AI:DashScope qwen3-vl-plus, MiMo ASR/TTS(可切换 Deepgram/OpenAI TTS)
|
||||||
|
存储:PostgreSQL 15 + Redis 7
|
||||||
|
CI/CD:Gitea Actions(`.gitea/workflows/deploy.yml`),push main/v2 自动构建部署到自建 aliyun runner
|
||||||
|
前端测试:**暂无**(package.json 无 test 脚本,无测试框架配置)
|
||||||
|
|
||||||
|
## 配置体系
|
||||||
|
|
||||||
|
配置优先级:**环境变量 > `config.{APP_ENV}.yaml` > `config.yaml` > 代码默认值**
|
||||||
|
|
||||||
|
配置文件位于 `backend/config/`:
|
||||||
|
- `config.yaml` — 基础配置(dev 默认值)
|
||||||
|
- `config.dev.yaml` — 开发环境覆盖(可选)
|
||||||
|
- `config.prod.yaml` — 生产环境覆盖(可选)
|
||||||
|
|
||||||
|
环境切换:`APP_ENV=dev|prod`(dev 默认,prod 启用限流 + 严格 CORS + Release 模式)
|
||||||
|
|
||||||
|
敏感信息(API Key、JWT Secret、数据库密码)**只能通过环境变量或 `.env` 文件注入**,不写入 YAML 配置文件。核心环境变量(参考 `backend/.env.example`):
|
||||||
|
|
||||||
|
| 变量 | 说明 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| 前端 | React 18, TypeScript, Vite, @ricky0123/vad-web |
|
| `CAMTALK_AI_LLM_API_KEY` | LLM API Key(DashScope) |
|
||||||
| 后端 | Go, Gin, gorilla/websocket, Viper, Zap |
|
| `CAMTALK_AI_STT_API_KEY` | STT API Key(MiMo/Deepgram) |
|
||||||
| LLM | GPT-4o(默认,通过 OpenAI 兼容接口可切换) |
|
| `CAMTALK_AI_TTS_API_KEY` | TTS API Key(MiMo/OpenAI) |
|
||||||
| STT | Deepgram(默认) / MiMo ASR |
|
| `CAMTALK_AUTH_JWT_SECRET` | JWT 签名密钥 |
|
||||||
| TTS | OpenAI TTS(默认) / MiMo TTS |
|
| `CAMTALK_STORAGE_DSN` | PostgreSQL 连接字符串 |
|
||||||
|
| `CAMTALK_REDIS_ADDR` | Redis 地址 |
|
||||||
|
| `CAMTALK_REDIS_PASSWORD` | Redis 密码 |
|
||||||
|
|
||||||
## 构建与运行命令
|
## 数据库迁移
|
||||||
|
|
||||||
```bash
|
迁移 SQL 文件位于 `backend/migrations/`(`001_*.up.sql` 等),通过 Go `//go:embed` 嵌入二进制(见 `backend/migrations/embed.go`)。应用启动时**自动执行**未应用的迁移,无需手动运行迁移命令。迁移通过 `schema_migrations` 表追踪执行状态。
|
||||||
# 前端
|
|
||||||
cd frontend && npm install
|
|
||||||
npm run dev # Vite 开发服务器
|
|
||||||
npm run build # 生产构建
|
|
||||||
npm run lint # ESLint 检查
|
|
||||||
npm run test # Vitest 测试
|
|
||||||
|
|
||||||
# 后端
|
回滚脚本为同目录下的 `*.down.sql` 文件,需手动执行。
|
||||||
cd backend && go mod download
|
|
||||||
go run ./cmd/server # 启动网关,监听 :8080
|
|
||||||
go build -o bin/camtalk ./cmd/server
|
|
||||||
go test ./... # 运行所有测试
|
|
||||||
go test -run TestName ./path # 运行单个测试
|
|
||||||
go vet ./... # 静态分析
|
|
||||||
```
|
|
||||||
|
|
||||||
基础设施:MVP 使用进程内存管理会话状态。Redis 已实现可通过配置切换,PostgreSQL 为规划中。
|
## 协议与 API
|
||||||
|
|
||||||
## WebSocket 协议
|
**WebSocket**:`ws://localhost:8080/ws?token=<jwt>&conversation_id=<uuid>`
|
||||||
|
- 客户端消息:`query`(图像/音频 Base64), `config`, `interrupt`, `ping`
|
||||||
|
- 服务端消息:`connected`, `stt_result`, `llm_chunk`, `llm_done`, `tts_audio`, `error`, `pong`
|
||||||
|
- 心跳:客户端 30s ping,服务端 60s 超时断连;重连:指数退避 1s→30s
|
||||||
|
- 实现:`CamTalkWebSocket` 单例(`frontend/src/lib/websocket.ts`),订阅模式,自动重连
|
||||||
|
- WebSocket 地址自动从当前页面协议/主机推导,也可通过 `VITE_WS_URL` 环境变量显式指定(如 `wss://api.example.com/ws`)
|
||||||
|
|
||||||
端点:`ws://localhost:8080/ws`
|
**REST API**:`/api/auth/*`(注册/登录/刷新/登出),`/api/conversations/*`(CRUD + 消息分页),`/api/scenarios/*`(用户自定义情景 CRUD),`/api/health`
|
||||||
|
|
||||||
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。完整契约见 `docs/03-接口文档.md`。
|
**错误码**:`INVALID_MESSAGE`, `SESSION_NOT_FOUND`, `RATE_LIMITED`, `IMAGE_TOO_LARGE`, `LLM_TIMEOUT`, `STT/TTS/LLM_ERROR`, `INVALID_TOKEN` 等
|
||||||
|
|
||||||
**客户端 → 服务端**:`query`(图像 Base64 + 音频 Base64)、`config`、`interrupt`、`ping`
|
## 关键文件路径
|
||||||
**服务端 → 客户端**:`connected`、`stt_result`、`llm_chunk`、`llm_done`、`tts_audio`、`error`、`pong`
|
|
||||||
|
|
||||||
**心跳**:客户端每 30 秒 ping,服务端 60 秒无 ping 断开连接。
|
**后端核心**:
|
||||||
**重连**:指数退避 + 抖动 —— 1s, 2s, 4s, 8s… 最大 30s。
|
- `backend/cmd/server/main.go` — 入口,依赖注入与启动流程(存储→AI 服务→Graph→路由→Server)
|
||||||
|
- `backend/internal/eino/` — Eino Graph 编排层(graph.go 构建、adapter.go 适配、callback.go 推送、state.go 状态、nodes_*.go 各节点实现)
|
||||||
|
- `backend/internal/session/tiered.go` — 三级会话存储(TieredManager)
|
||||||
|
- `backend/internal/store/` — 持久化层(Repository 接口 + PG 实现 + Redis 缓存装饰器)
|
||||||
|
- Repository 模式:接口定义在 `user.go`/`session.go`/`message.go`,PG 实现在 `*_pg.go`,Redis 缓存装饰器在 `cached_user.go`
|
||||||
|
- `backend/internal/ws/handler.go` — WebSocket 连接管理(升级→认证→收发循环→清理)
|
||||||
|
- `backend/internal/ai/` — AI 服务抽象层(llm/stt/tts 各子目录,统一 `Service` 接口)
|
||||||
|
- `backend/internal/auth/` — JWT/bcrypt/中间件
|
||||||
|
- `backend/internal/ratelimit/` — 令牌桶限流(内存/Redis 两种后端)
|
||||||
|
- `backend/migrations/` — 嵌入式 SQL 迁移文件(embed.go + *.sql)
|
||||||
|
|
||||||
## REST API(辅助)
|
**前端核心**:
|
||||||
|
- `frontend/src/hooks/useVisionSession.ts` — 核心会话 Hook(~500 行,编排整个采集→发送→接收→播放流程)
|
||||||
- `GET /api/health` — 健康检查(版本、运行时间、活跃会话数)
|
- `frontend/src/lib/websocket.ts` — WebSocket 客户端单例(心跳/重连/订阅模式)
|
||||||
- `POST /api/sessions` — 创建会话(可选,MVP 在 WS 连接时自动创建)
|
- `frontend/src/lib/auth.tsx` — AuthProvider(JWT 自动刷新 + React Context)
|
||||||
- `DELETE /api/sessions/{id}` — 销毁会话
|
- `frontend/src/lib/api.ts` — REST 客户端(401 拦截 + token 刷新)
|
||||||
|
- `frontend/src/lib/ttsPlayer.ts` — 流式 TTS 音频播放队列
|
||||||
## 错误码
|
- `frontend/src/lib/i18n/` — 国际化(zh-CN / en-US / ja-JP)
|
||||||
|
- `frontend/src/components/` — UI 组件(LandingPage/CameraManager/MicManager/WebSocketManager/ChatPanel/SessionSidebar/ConfigPanel/VideoPreview 等)
|
||||||
`INVALID_MESSAGE`、`SESSION_NOT_FOUND`、`RATE_LIMITED`、`IMAGE_TOO_LARGE`、`AUDIO_TOO_SHORT`、`LLM_TIMEOUT`、`LLM_ERROR`、`STT_ERROR`、`TTS_ERROR`、`INTERNAL_ERROR`
|
- `frontend/vite.config.ts` — VAD 模型文件自动复制 + ONNX WASM MIME 处理 + 代理配置
|
||||||
|
|
||||||
## 前端组件结构
|
|
||||||
|
|
||||||
| 组件 | 职责 |
|
|
||||||
|------|------|
|
|
||||||
| `CameraManager` | 摄像头流采集 |
|
|
||||||
| `MicManager` | 麦克风音频采集 |
|
|
||||||
| `EdgeProcessor` | VAD + 关键帧检测(Canvas 像素比较) |
|
|
||||||
| `WebSocketManager` | WebSocket 连接生命周期管理 |
|
|
||||||
| `ChatPanel` | 消息展示 |
|
|
||||||
| `VideoPreview` | 摄像头画面预览 |
|
|
||||||
|
|
||||||
## 后端模块结构
|
|
||||||
|
|
||||||
| 模块 | 职责 |
|
|
||||||
|------|------|
|
|
||||||
| WebSocket Handler | 连接管理、单播消息推送 |
|
|
||||||
| Session Manager | 会话状态、对话历史(Memory/Redis,30 分钟 TTL) |
|
|
||||||
| AI Orchestrator | STT→LLM→TTS 流式并行管道编排 |
|
|
||||||
| AI Service Layer | AI 服务抽象层(STT/LLM/TTS 多 provider) |
|
|
||||||
| REST API | 健康检查、会话管理(Gin 路由) |
|
|
||||||
| Models | 数据模型定义 |
|
|
||||||
| Model Router | 按请求选择 AI 模型(规划中) |
|
|
||||||
| Rate Limiter | 按用户的令牌桶速率限制(规划中) |
|
|
||||||
|
|
||||||
## 编码规范
|
## 编码规范
|
||||||
|
|
||||||
- **Go**:遵循标准 Go 规范。所有 AI 调用使用 `context.Context` 做取消/超时。并发 map 访问使用 `sync.RWMutex`。结构体标签用 `json:"snake_case"`。
|
- **Go**:标准规范,`context.Context` 超时控制,`sync.RWMutex` 并发保护,`json:"snake_case"` 标签,编译期接口检查 `var _ Interface = (*Impl)(nil)`
|
||||||
- **TypeScript**:严格模式。所有数据模型用接口定义。WebSocket 消息类型用可辨识联合类型(`type` 字段)。
|
- **TypeScript**:严格模式,接口定义数据模型,WebSocket 消息用可辨识联合类型(`type` 字段区分)
|
||||||
- **提交信息**:Conventional Commits 格式,描述用中文。示例:`feat: 添加 WebSocket 连接管理`、`fix: 修复心跳超时判断`、`docs: 更新接口文档`
|
- **存储层模式**:Repository 接口 + PostgreSQL 实现 + Redis 缓存装饰器(`CachedUserRepository` 包装模式)
|
||||||
- **禁止自动 push**:除非用户明确要求。
|
- **CORS**:禁止后端代码/配置文件配置 CORS,统一由代理层处理(开发环境 Vite proxy,生产环境 Nginx)
|
||||||
- **文档优先**:实现功能前先读取 `docs/` 下的相关设计文档。实现与文档不一致时,优先更新 `docs/` 下的接口文档。
|
- **提交信息**:Conventional Commits,中文描述(如 `feat: 添加 WebSocket 心跳`)
|
||||||
|
- **禁止自动 push**:除非用户明确要求
|
||||||
|
- **文档优先**:开发前先读 `docs/` 设计文档,代码与文档不一致时优先更新文档
|
||||||
|
|||||||
592
README.md
592
README.md
@@ -1,137 +1,561 @@
|
|||||||
# CamTalk
|
# CamTalk
|
||||||
|
|
||||||
多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。
|
<div align="center">
|
||||||
|
|
||||||
## 架构
|
**多模态实时 AI 视觉对话助手**
|
||||||
|
|
||||||
三层系统,前端做轻量预处理,后端做智能编排,云端 AI 服务按需调用:
|
用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应
|
||||||
|
|
||||||
```
|
[](https://opensource.org/licenses/MIT)
|
||||||
浏览器客户端 Go 网关 :8080 云端 AI 服务
|
[](https://go.dev/)
|
||||||
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
|
[](https://react.dev/)
|
||||||
│ 媒体采集 │ │ WebSocket Handler│ │ STT(语音识别) │
|
[](https://www.typescriptlang.org/)
|
||||||
│ VAD 语音检测 │ WebSocket│ Session Manager │ HTTP │ LLM(多模态推理) │
|
|
||||||
│ 关键帧检测 │ ◄──────► │ AI Orchestrator │ ◄──────► │ TTS(语音合成) │
|
[路演视频](https://www.bilibili.com/video/BV1dDJK6cE5S/) • [在线体验](https://camtalk.goanchor.top) • [文档](docs/README.md)
|
||||||
│ UI 渲染 │ │ REST API │ │ │
|
|
||||||
└─────────────────┘ └─────────────────┘ └─────────────────┘
|
</div>
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
|
||||||
|
<!-- > ⚠️ **在线体验提示**:由于演示环境使用 HTTP 协议,需配置 Chrome 允许非 HTTPS 下访问摄像头/麦克风:
|
||||||
|
>
|
||||||
|
> 1. 访问 `chrome://flags/#unsafely-treat-insecure-origin-as-secure`
|
||||||
|
> 2. 启用该选项,并在输入框填入 `http://8.161.227.145:9000`
|
||||||
|
> 3. 点击 **Relaunch** 重启浏览器
|
||||||
|
|
||||||
|
 -->
|
||||||
|
|
||||||
|
## ✨ 核心特性
|
||||||
|
|
||||||
|
- 🎥 **多模态理解**:摄像头视觉 + 麦克风语音双输入,AI 理解完整场景
|
||||||
|
- 🗣️ **自然对话**:基于 VAD 的端到端语音交互,低延迟流式响应
|
||||||
|
- 🚀 **实时推送**:LLM 文本流 + TTS 音频流并行推送,感知延迟 < 0.5 秒
|
||||||
|
- 🎭 **情景模式**:自由对话、面试官、英语老师等多场景支持
|
||||||
|
- 💾 **对话历史**:自动保存会话,支持搜索、重命名、删除、时间分组
|
||||||
|
- 🔐 **安全认证**:JWT 双 token 轮转 + Refresh Token Rotation 防重放
|
||||||
|
- 📊 **三级存储**:Memory → Redis → PostgreSQL 自动降级,保障可靠性
|
||||||
|
- 🌐 **国际化**:支持中文、英文、日文界面
|
||||||
|
|
||||||
|
## 🏗️ 系统架构
|
||||||
|
|
||||||
|
CamTalk 采用**三层架构**:前端轻量预处理 → Go 网关智能编排 → 云端 AI 按需调用
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph TB
|
||||||
|
subgraph Browser["🌐 浏览器客户端"]
|
||||||
|
UI["React UI 渲染"]
|
||||||
|
VAD["VAD 语音检测"]
|
||||||
|
Media["媒体采集"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Gateway["⚙️ Go 网关 (Eino Graph)"]
|
||||||
|
WS["WebSocket Handler"]
|
||||||
|
Auth["JWT 认证"]
|
||||||
|
Session["会话管理 (三级存储)"]
|
||||||
|
Orch["AI 编排器 (7节点DAG)"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph AI["☁️ 云端 AI 服务"]
|
||||||
|
STT["STT (MiMo/Deepgram)"]
|
||||||
|
LLM["LLM (qwen3-vl-plus)"]
|
||||||
|
TTS["TTS (MiMo/OpenAI)"]
|
||||||
|
end
|
||||||
|
|
||||||
|
Browser <-->|"WebSocket<br/>(JWT + query/config)"| Gateway
|
||||||
|
Orch --> STT
|
||||||
|
Orch --> LLM
|
||||||
|
Orch --> TTS
|
||||||
```
|
```
|
||||||
|
|
||||||
**关键模式**:LLM 文本流和 TTS 音频流并行推送,用户先看到文字、紧接着听到语音,感知延迟 < 0.5 秒。
|
### AI 编排流水线(Eino Graph)
|
||||||
|
|
||||||
## 技术栈
|
基于 [CloudWeGo Eino](https://github.com/cloudwego/eino) 框架的声明式 7 节点 DAG:
|
||||||
|
|
||||||
| 层级 | 技术 |
|
```
|
||||||
|------|------|
|
START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END
|
||||||
| 前端 | React 18, TypeScript, Vite, @ricky0123/vad-web |
|
```
|
||||||
| 后端 | Go, Gin, gorilla/websocket, Viper, Zap |
|
|
||||||
| STT | Deepgram(默认) / MiMo ASR |
|
|
||||||
| LLM | GPT-4o(默认,通过 OpenAI 兼容接口可切换) |
|
|
||||||
| TTS | OpenAI TTS(默认) / MiMo TTS |
|
|
||||||
|
|
||||||
## 项目结构
|
**核心优势**:
|
||||||
|
- **流式处理**:ChatModel 逐 token 推送,Callback AOP 机制实时转发客户端
|
||||||
|
- **句子级 TTS**:Splitter 实时切分句子,TTS 逐句并行合成,无需等待完整回复
|
||||||
|
- **类型安全**:Go 泛型 + 编译期检查,Graph 拓扑错误在编译时发现
|
||||||
|
|
||||||
|
## 🛠️ 技术栈
|
||||||
|
|
||||||
|
<table>
|
||||||
|
<tr>
|
||||||
|
<td><b>层级</b></td>
|
||||||
|
<td><b>技术选型</b></td>
|
||||||
|
<td><b>说明</b></td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>前端</b></td>
|
||||||
|
<td>React 18 + TypeScript + Vite</td>
|
||||||
|
<td>组件化开发,类型安全,快速热更新</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>VAD</b></td>
|
||||||
|
<td>@ricky0123/vad-web (ONNX Runtime)</td>
|
||||||
|
<td>浏览器端语音活动检测,零延迟</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>后端</b></td>
|
||||||
|
<td>Go 1.25+ + Gin + gorilla/websocket</td>
|
||||||
|
<td>高并发 goroutine,长连接管理</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>AI 编排</b></td>
|
||||||
|
<td>CloudWeGo Eino Graph</td>
|
||||||
|
<td>声明式 DAG,Stream 模式,Callback AOP</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>STT</b></td>
|
||||||
|
<td>MiMo ASR(默认)/ Deepgram</td>
|
||||||
|
<td>实时语音识别,多语言支持</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>LLM</b></td>
|
||||||
|
<td>DashScope qwen3-vl-plus</td>
|
||||||
|
<td>多模态推理(通过 eino-ext OpenAI 接入)</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>TTS</b></td>
|
||||||
|
<td>MiMo TTS(默认)/ OpenAI TTS</td>
|
||||||
|
<td>自然语音合成</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>存储</b></td>
|
||||||
|
<td>PostgreSQL 15 + Redis 7</td>
|
||||||
|
<td>三级存储架构:Memory → Redis → PG</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>认证</b></td>
|
||||||
|
<td>JWT (HS256) + bcrypt</td>
|
||||||
|
<td>双 token 轮转 + Refresh Token Rotation</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>配置</b></td>
|
||||||
|
<td>Viper + godotenv</td>
|
||||||
|
<td>YAML + .env + 环境变量覆盖</td>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><b>日志</b></td>
|
||||||
|
<td>Zap</td>
|
||||||
|
<td>高性能结构化日志 + Trace ID 追踪</td>
|
||||||
|
</tr>
|
||||||
|
</table>
|
||||||
|
|
||||||
|
## 📁 项目结构
|
||||||
|
|
||||||
```
|
```
|
||||||
CamTalk/
|
CamTalk/
|
||||||
├── frontend/ # 浏览器客户端
|
├── frontend/ # 🌐 浏览器客户端
|
||||||
│ └── src/
|
│ └── src/
|
||||||
│ ├── components/ # UI 组件
|
│ ├── components/ # UI 组件
|
||||||
|
│ │ ├── LandingPage/ # 登录着陆页 + LoginModal
|
||||||
│ │ ├── CameraManager/ # 摄像头流采集
|
│ │ ├── CameraManager/ # 摄像头流采集
|
||||||
│ │ ├── MicManager/ # 麦克风音频采集
|
│ │ ├── MicManager/ # 麦克风音频采集 + VAD
|
||||||
│ │ ├── EdgeProcessor/ # VAD + 关键帧检测
|
│ │ ├── WebSocketManager/ # WS 连接生命周期
|
||||||
│ │ ├── WebSocketManager/ # WS 连接管理
|
│ │ ├── ChatPanel/ # 消息展示 + 流式回复
|
||||||
│ │ ├── ChatPanel/ # 消息展示
|
│ │ ├── SessionSidebar/ # 对话历史侧边栏
|
||||||
│ │ ├── VideoPreview/ # 摄像头画面预览
|
│ │ └── ConfigPanel/ # 配置面板(主题/TTS/语言/场景)
|
||||||
│ │ ├── ConfigPanel/ # 配置面板
|
|
||||||
│ │ └── Toast/ # 通知提示
|
|
||||||
│ ├── hooks/ # 自定义 Hooks
|
│ ├── hooks/ # 自定义 Hooks
|
||||||
│ │ ├── useVisionSession.ts # 核心会话 Hook
|
│ │ ├── useVisionSession.ts # 核心会话 Hook (~500 行)
|
||||||
|
│ │ ├── useSessionList.ts # 对话列表管理
|
||||||
│ │ └── useObservationMode.ts # 观察模式
|
│ │ └── useObservationMode.ts # 观察模式
|
||||||
│ ├── lib/ # 工具库
|
│ ├── lib/ # 工具库
|
||||||
│ │ ├── websocket.ts # WebSocket 连接管理
|
│ │ ├── websocket.ts # WebSocket 单例(心跳/重连/订阅)
|
||||||
│ │ ├── audio.ts # 音频编码
|
│ │ ├── api.ts # REST 客户端(401拦截+刷新)
|
||||||
│ │ ├── ttsPlayer.ts # TTS 播放器
|
│ │ ├── auth.tsx # AuthProvider(JWT 自动刷新)
|
||||||
│ │ └── sampling.ts # 采样策略
|
│ │ ├── ttsPlayer.ts # TTS 流式播放队列
|
||||||
|
│ │ └── i18n/ # 国际化(zh-CN/en-US/ja-JP)
|
||||||
│ └── types/ # TypeScript 类型定义
|
│ └── types/ # TypeScript 类型定义
|
||||||
├── backend/ # Go 网关
|
├── backend/ # ⚙️ Go 网关
|
||||||
│ ├── cmd/server/ # 入口
|
│ ├── cmd/server/ # 服务入口(main.go)
|
||||||
│ └── internal/
|
│ └── internal/
|
||||||
|
│ ├── eino/ # 🔥 Eino Graph 编排层(7节点DAG)
|
||||||
|
│ │ ├── graph.go # Graph 构建与编译
|
||||||
|
│ │ ├── adapter.go # EinoOrchestrator 适配器
|
||||||
|
│ │ ├── callback.go # LLM token 推送回调
|
||||||
|
│ │ ├── state.go # 跨节点状态管理
|
||||||
|
│ │ └── nodes_*.go # STT/History/Splitter/TTS/Done 节点
|
||||||
|
│ ├── session/ # 会话管理(TieredManager 三级存储)
|
||||||
|
│ ├── store/ # 持久化层(Repository 接口 + PG/内存实现)
|
||||||
|
│ │ ├── user_pg.go # PostgreSQL 实现
|
||||||
|
│ │ └── cached_user.go # Redis 缓存装饰器
|
||||||
|
│ ├── auth/ # 认证(JWT/bcrypt/中间件)
|
||||||
│ ├── ai/ # AI 服务抽象层
|
│ ├── ai/ # AI 服务抽象层
|
||||||
│ │ ├── llm/ # LLM 服务(OpenAI 兼容)
|
│ │ ├── llm/ # LLM 提示词与场景
|
||||||
│ │ ├── stt/ # STT 服务(Deepgram/MiMo)
|
│ │ ├── stt/ # STT 服务(MiMo/Deepgram)
|
||||||
│ │ └── tts/ # TTS 服务(OpenAI/MiMo)
|
│ │ └── tts/ # TTS 服务(MiMo/OpenAI)
|
||||||
│ ├── orchestrator/ # AI 编排器(STT→LLM→TTS 管道)
|
|
||||||
│ ├── session/ # 会话管理(Memory/Redis)
|
|
||||||
│ ├── ws/ # WebSocket Handler
|
│ ├── ws/ # WebSocket Handler
|
||||||
│ ├── api/ # REST API
|
│ ├── api/ # REST API(Auth/Conversation)
|
||||||
│ ├── config/ # 配置管理
|
│ ├── config/ # 配置管理(Viper)
|
||||||
│ ├── models/ # 数据模型
|
│ └── logger/ # 日志(Zap + Trace ID)
|
||||||
│ ├── errors/ # 错误码
|
├── migrations/ # 📊 数据库迁移(嵌入式 SQL)
|
||||||
│ └── logger/ # 日志
|
├── docs/ # 📚 设计文档
|
||||||
├── docs/ # 设计文档
|
│ ├── 01-架构设计.md
|
||||||
└── CLAUDE.md # Claude Code 指引
|
│ ├── 02-接口文档.md
|
||||||
|
│ ├── 08-Eino框架与编排设计.md
|
||||||
|
│ ├── 10-鉴权体系.md
|
||||||
|
│ └── 13-日志追踪.md
|
||||||
|
├── deploy.sh # 🐳 部署脚本(Docker Compose)
|
||||||
|
├── docker-compose.yml # 容器编排配置
|
||||||
|
└── CLAUDE.md # 🤖 Claude Code 开发指引
|
||||||
```
|
```
|
||||||
|
|
||||||
## 快速开始
|
## 🚀 快速开始
|
||||||
|
|
||||||
### 前置条件
|
### 前置条件
|
||||||
|
|
||||||
- Node.js >= 18
|
- **Node.js** >= 18
|
||||||
- Go >= 1.24
|
- **Go** >= 1.25
|
||||||
|
- **PostgreSQL** >= 15(可选 Docker)
|
||||||
|
- **Redis** >= 7(可选,用于缓存加速)
|
||||||
|
|
||||||
### 前端
|
### 本地开发
|
||||||
|
|
||||||
|
#### 1. 克隆项目
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
cd frontend
|
git clone https://github.com/yourusername/CamTalk.git
|
||||||
npm install
|
cd CamTalk
|
||||||
npm run dev # Vite 开发服务器 http://localhost:5173
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### 后端
|
#### 2. 配置环境变量
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 复制环境变量模板
|
||||||
|
cp backend/.env.example backend/.env
|
||||||
|
|
||||||
|
# 编辑 .env 文件,填入以下必需配置:
|
||||||
|
# - CAMTALK_AUTH_JWT_SECRET(使用 openssl rand -hex 32 生成)
|
||||||
|
# - CAMTALK_STORAGE_DSN(PostgreSQL 连接字符串)
|
||||||
|
# - CAMTALK_AI_LLM_API_KEY(DashScope API Key)
|
||||||
|
# - CAMTALK_AI_STT_API_KEY(MiMo/Deepgram API Key)
|
||||||
|
# - CAMTALK_AI_TTS_API_KEY(MiMo/OpenAI API Key)
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 3. 启动后端
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
cd backend
|
cd backend
|
||||||
|
|
||||||
|
# 安装依赖
|
||||||
go mod download
|
go mod download
|
||||||
go run ./cmd/server # 启动网关 :8080
|
|
||||||
```
|
|
||||||
|
|
||||||
### 配置
|
# 运行数据库迁移(自动创建表)
|
||||||
|
go run ./cmd/server migrate
|
||||||
|
|
||||||
后端配置文件位于 `backend/config.yaml`,支持环境变量覆盖(前缀 `CAMTALK_`)。
|
# 启动服务(监听 :8080)
|
||||||
|
|
||||||
```bash
|
|
||||||
# 最小启动(需要至少一个 AI 服务的 API Key)
|
|
||||||
cd backend
|
|
||||||
CAMTALK_AI_LLM_API_KEY=sk-xxx \
|
|
||||||
CAMTALK_AI_STT_API_KEY=xxx \
|
|
||||||
go run ./cmd/server
|
go run ./cmd/server
|
||||||
```
|
```
|
||||||
|
|
||||||
配置优先级:环境变量 > `config.{env}.yaml` > `config.yaml` > `.env`
|
#### 4. 启动前端
|
||||||
|
|
||||||
## WebSocket 协议
|
```bash
|
||||||
|
cd frontend
|
||||||
|
|
||||||
连接地址:`ws://localhost:8080/ws`
|
# 安装依赖
|
||||||
|
npm install
|
||||||
|
|
||||||
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。
|
# 启动开发服务器(http://localhost:5173)
|
||||||
|
npm run dev
|
||||||
|
```
|
||||||
|
|
||||||
**客户端 → 服务端**:`query`、`config`、`interrupt`、`ping`
|
#### 5. 访问应用
|
||||||
**服务端 → 客户端**:`connected`、`stt_result`、`llm_chunk`、`llm_done`、`tts_audio`、`error`、`pong`
|
|
||||||
|
|
||||||
完整协议见 [docs/03-接口文档.md](docs/03-接口文档.md)。
|
打开浏览器访问 [http://localhost:5173](http://localhost:5173),注册账号后即可开始使用。
|
||||||
|
|
||||||
## 文档
|
#### 6. 代码检查与测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 安装 Go 代码检查工具
|
||||||
|
go install github.com/golangci/golangci-lint/cmd/golangci-lint@latest
|
||||||
|
|
||||||
|
# 运行后端代码检查
|
||||||
|
cd backend
|
||||||
|
golangci-lint run
|
||||||
|
|
||||||
|
# 后端单元测试
|
||||||
|
go test ./...
|
||||||
|
|
||||||
|
# 后端集成测试(需要 PostgreSQL)
|
||||||
|
go test -tags=integration ./...
|
||||||
|
|
||||||
|
# 前端代码检查
|
||||||
|
cd frontend
|
||||||
|
npm run lint
|
||||||
|
|
||||||
|
# 前端测试
|
||||||
|
npm test
|
||||||
|
```
|
||||||
|
|
||||||
|
### 远程部署
|
||||||
|
|
||||||
|
#### 方式一:Docker Compose(推荐)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. 克隆代码到服务器
|
||||||
|
git clone https://github.com/yourusername/CamTalk.git
|
||||||
|
cd CamTalk
|
||||||
|
|
||||||
|
# 2. 配置环境变量
|
||||||
|
cp backend/.env.example backend/.env
|
||||||
|
# 编辑 .env 文件,填入生产环境配置
|
||||||
|
|
||||||
|
# 3. 一键部署(frontend + backend + postgres + redis)
|
||||||
|
./deploy.sh up
|
||||||
|
|
||||||
|
# 4. 查看日志
|
||||||
|
./deploy.sh logs
|
||||||
|
|
||||||
|
# 5. 停止服务
|
||||||
|
./deploy.sh down
|
||||||
|
```
|
||||||
|
|
||||||
|
部署完成后访问 [http://localhost:9000](http://localhost:9000)
|
||||||
|
|
||||||
|
#### 方式二:手动部署
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. 构建前端
|
||||||
|
cd frontend
|
||||||
|
npm install
|
||||||
|
npm run build # 输出到 dist/
|
||||||
|
|
||||||
|
# 2. 构建后端
|
||||||
|
cd backend
|
||||||
|
go build -o camtalk ./cmd/server
|
||||||
|
|
||||||
|
# 3. 配置 Nginx
|
||||||
|
# 参考 nginx.conf.example 配置反向代理
|
||||||
|
|
||||||
|
# 4. 启动服务
|
||||||
|
APP_ENV=prod ./camtalk
|
||||||
|
|
||||||
|
# 5. 使用 systemd 管理(可选)
|
||||||
|
sudo systemctl enable camtalk
|
||||||
|
sudo systemctl start camtalk
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 环境变量检查清单
|
||||||
|
|
||||||
|
部署前确保已配置以下环境变量:
|
||||||
|
|
||||||
|
- ✅ `CAMTALK_AUTH_JWT_SECRET`(使用 `openssl rand -hex 32` 生成)
|
||||||
|
- ✅ `CAMTALK_STORAGE_DSN`(PostgreSQL 连接字符串)
|
||||||
|
- ✅ `CAMTALK_AI_LLM_API_KEY`(DashScope API Key)
|
||||||
|
- ✅ `CAMTALK_AI_STT_API_KEY`(STT 服务 API Key)
|
||||||
|
- ✅ `CAMTALK_AI_TTS_API_KEY`(TTS 服务 API Key)
|
||||||
|
- ✅ `APP_ENV=prod`(启用生产环境配置)
|
||||||
|
|
||||||
|
### 配置优先级
|
||||||
|
|
||||||
|
```
|
||||||
|
环境变量 > config.{APP_ENV}.yaml > config.yaml > .env
|
||||||
|
```
|
||||||
|
|
||||||
|
通过 `APP_ENV=prod` 切换生产环境配置(启用限流 + 严格 CORS)
|
||||||
|
|
||||||
|
## 📡 WebSocket 协议
|
||||||
|
|
||||||
|
连接地址:`ws://localhost:8080/ws?token=<jwt>&conversation_id=<uuid>`
|
||||||
|
|
||||||
|
所有消息为 JSON 文本帧,统一信封格式:
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface BaseMessage {
|
||||||
|
type: string;
|
||||||
|
request_id?: string;
|
||||||
|
timestamp?: number;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 客户端 → 服务端
|
||||||
|
|
||||||
|
| 消息类型 | 说明 | 示例 |
|
||||||
|
|---------|------|------|
|
||||||
|
| `query` | 发送视觉+语音查询 | `{type: "query", image: "base64...", audio: "base64..."}` |
|
||||||
|
| `config` | 更新会话配置 | `{type: "config", scenario: "interviewer", language: "en"}` |
|
||||||
|
| `interrupt` | 中断当前响应 | `{type: "interrupt", request_id: "xxx"}` |
|
||||||
|
| `ping` | 心跳保活 | `{type: "ping"}` |
|
||||||
|
|
||||||
|
### 服务端 → 客户端
|
||||||
|
|
||||||
|
| 消息类型 | 说明 | 触发时机 |
|
||||||
|
|---------|------|---------|
|
||||||
|
| `connected` | 连接成功 | WebSocket 握手后 |
|
||||||
|
| `stt_result` | STT 识别结果 | STT 节点完成 |
|
||||||
|
| `llm_chunk` | LLM 文本增量 | ChatModel 逐 token(Callback) |
|
||||||
|
| `llm_done` | LLM 推理完成 | Done 节点执行 |
|
||||||
|
| `tts_audio` | TTS 音频片段 | TTS 节点逐句合成 |
|
||||||
|
| `error` | 错误通知 | 任意节点失败 |
|
||||||
|
| `pong` | 心跳响应 | 响应 `ping` |
|
||||||
|
|
||||||
|
**心跳机制**:
|
||||||
|
- 客户端每 30 秒发送 `ping`
|
||||||
|
- 服务端 60 秒无消息自动断连
|
||||||
|
- 断连后自动重连(指数退避 1s → 30s)
|
||||||
|
|
||||||
|
完整协议定义见 [docs/02-接口文档.md](docs/02-接口文档.md)
|
||||||
|
|
||||||
|
## 🔐 认证体系
|
||||||
|
|
||||||
|
CamTalk 采用 **JWT 双 token 轮转 + Refresh Token Rotation** 安全机制:
|
||||||
|
|
||||||
|
### 双 Token 设计
|
||||||
|
|
||||||
|
| Token | 有效期 | 存储位置 | 用途 |
|
||||||
|
|-------|-------|---------|------|
|
||||||
|
| `access_token` | 120 分钟 | 前端内存(推荐)/ localStorage | 访问受保护资源 |
|
||||||
|
| `refresh_token` | 7 天 | httpOnly Cookie(推荐)/ localStorage | 刷新 access_token |
|
||||||
|
|
||||||
|
### Refresh Token Rotation
|
||||||
|
|
||||||
|
每次刷新 token 时:
|
||||||
|
1. 验证 `refresh_token` 签名和有效期
|
||||||
|
2. 查询数据库中的 SHA256 哈希
|
||||||
|
3. **如果哈希不存在** → 检测到 token 复用 → **吊销该用户所有 token**
|
||||||
|
4. 删除旧 refresh_token,生成新 token pair
|
||||||
|
5. 返回新 access_token + refresh_token
|
||||||
|
|
||||||
|
**防重放攻击**:旧 refresh_token 立即失效,复用时触发全局吊销,强制所有设备重新登录。
|
||||||
|
|
||||||
|
### REST API 端点
|
||||||
|
|
||||||
|
- `POST /api/auth/register` — 用户注册
|
||||||
|
- `POST /api/auth/login` — 用户登录
|
||||||
|
- `POST /api/auth/refresh` — 刷新 token
|
||||||
|
- `POST /api/auth/logout` — 登出(需认证)
|
||||||
|
- `GET /api/conversations` — 获取对话列表(需认证)
|
||||||
|
- `POST /api/conversations` — 创建对话(需认证)
|
||||||
|
- `GET /api/health` — 健康检查
|
||||||
|
|
||||||
|
详细设计见 [docs/10-鉴权体系.md](docs/10-鉴权体系.md)
|
||||||
|
|
||||||
|
## 💾 三级存储架构
|
||||||
|
|
||||||
|
**TieredManager** 实现会话状态的三级存储,平衡性能与可靠性:
|
||||||
|
|
||||||
|
```
|
||||||
|
┌─────────────┐
|
||||||
|
│ L1 Memory │ ← 微秒级读写,进程内缓存
|
||||||
|
├─────────────┤
|
||||||
|
│ L2 Redis │ ← 毫秒级访问,跨实例共享
|
||||||
|
├─────────────┤
|
||||||
|
│ L3 PostgreSQL│ ← 持久化存储,数据可靠性
|
||||||
|
└─────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
**特性**:
|
||||||
|
- ✅ **自动降级**:Redis 故障时自动切换到 Memory + PostgreSQL 模式
|
||||||
|
- ✅ **灵活配置**:支持单级(Memory)、双级(Memory + PG)、完整三级
|
||||||
|
- ✅ **TTL 管理**:会话默认 30 分钟过期,自动清理
|
||||||
|
- ✅ **写穿透**:数据先写 L1,异步同步到 L2/L3
|
||||||
|
|
||||||
|
## 📊 数据库设计
|
||||||
|
|
||||||
|
系统使用 PostgreSQL 存储持久化数据:
|
||||||
|
|
||||||
|
### 核心表
|
||||||
|
|
||||||
|
| 表名 | 说明 | 关键字段 |
|
||||||
|
|------|------|---------|
|
||||||
|
| `users` | 用户账户 | `id (UUID)`, `username (UNIQUE)`, `password_hash (bcrypt)` |
|
||||||
|
| `sessions` | 对话会话 | `id (UUID)`, `user_id (FK)`, `title`, `config (JSONB)` |
|
||||||
|
| `messages` | 消息记录 | `id (BIGSERIAL)`, `session_id (FK)`, `role`, `content`, `tokens_used` |
|
||||||
|
| `refresh_tokens` | 刷新令牌 | `token_hash (PK, SHA256)`, `user_id (FK)`, `expires_at` |
|
||||||
|
|
||||||
|
**关系**:`users 1:N sessions 1:N messages`,`users 1:N refresh_tokens`
|
||||||
|
|
||||||
|
**迁移管理**:使用嵌入式 SQL 文件(`backend/migrations/`),应用启动时自动执行。
|
||||||
|
|
||||||
|
## 🛡️ 安全特性
|
||||||
|
|
||||||
|
- 🔒 **密码安全**:bcrypt (cost=10) 哈希,自动生成盐值
|
||||||
|
- 🔑 **Token 安全**:JWT HS256 签名,refresh_token SHA256 哈希存储
|
||||||
|
- 🚫 **防重放攻击**:Refresh Token Rotation + 复用检测自动吊销
|
||||||
|
- 🌐 **传输安全**:生产环境强制 HTTPS,开发环境 Vite proxy 同源代理
|
||||||
|
- 🚦 **限流保护**:令牌桶算法(生产环境启用),防暴力破解
|
||||||
|
- 🔍 **日志追踪**:全链路 Trace ID,请求/响应/错误统一记录
|
||||||
|
|
||||||
|
## 🌍 部署架构
|
||||||
|
|
||||||
|
```
|
||||||
|
┌─────────────┐
|
||||||
|
│ Nginx │ ← 反向代理(静态资源 + API + WebSocket)
|
||||||
|
└──────┬──────┘
|
||||||
|
│
|
||||||
|
┌──────┴───────────────────┐
|
||||||
|
│ Go Gateway 集群 │
|
||||||
|
│ ├─ Gateway-1 │
|
||||||
|
│ ├─ Gateway-2 │
|
||||||
|
│ └─ Gateway-N │
|
||||||
|
└───┬────────────┬─────────┘
|
||||||
|
│ │
|
||||||
|
┌───┴────┐ ┌───┴────────┐
|
||||||
|
│ Redis │ │ PostgreSQL │
|
||||||
|
└────────┘ └────────────┘
|
||||||
|
│
|
||||||
|
┌───┴────────────────────┐
|
||||||
|
│ 外部 AI 服务 │
|
||||||
|
│ ├─ DashScope (LLM) │
|
||||||
|
│ ├─ MiMo (STT/TTS) │
|
||||||
|
│ └─ Deepgram (可选) │
|
||||||
|
└───────────────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
**跨域策略**:Nginx 统一反代前后端到同一域名,无跨域问题。
|
||||||
|
|
||||||
|
**水平扩展**:Gateway 无状态设计,会话状态存储在 Redis/PostgreSQL,支持多实例部署。
|
||||||
|
|
||||||
|
## 📖 文档
|
||||||
|
|
||||||
|
### 核心设计文档
|
||||||
|
|
||||||
| 文档 | 内容 |
|
| 文档 | 内容 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| [01-项目概述](docs/01-项目概述.md) | 项目目标与核心挑战 |
|
| [01-架构设计](docs/01-架构设计.md) | 三层架构、技术栈、数据库设计、部署方案 |
|
||||||
| [02-系统架构](docs/02-系统架构.md) | 三层架构、技术栈、部署方案 |
|
| [02-接口文档](docs/02-接口文档.md) | WebSocket 协议、REST API、AI 服务层、编排器、配置管理 |
|
||||||
| [03-接口文档](docs/03-接口文档.md) | WebSocket 协议、REST API、配置管理 |
|
| [08-Eino框架与编排设计](docs/08-Eino框架与编排设计.md) | Eino Graph 7 节点 DAG、节点实现、流式处理、Callback AOP |
|
||||||
| [04-技术选型](docs/04-技术选型.md) | AI 服务栈、持久化层、前端边缘处理选型 |
|
| [10-鉴权体系](docs/10-鉴权体系.md) | JWT 双 token 轮转、Refresh Token Rotation、密码安全、中间件 |
|
||||||
| [05-用户故事](docs/05-用户故事.md) | 用户场景与优先级 |
|
| [11-令牌桶限流](docs/11-令牌桶限流.md) | 限流算法、配置策略、生产环境保护 |
|
||||||
| [06-语音交互](docs/06-语音交互.md) | VAD → STT → LLM → TTS 全链路 |
|
| [13-日志追踪](docs/13-日志追踪.md) | Zap 日志、Trace ID 全链路追踪、日志级别 |
|
||||||
| [07-视觉理解](docs/07-视觉理解.md) | 帧采样、关键帧检测、多模态输入 |
|
|
||||||
| [08-成本控制](docs/08-成本控制.md) | 采样策略、端云协同、模型分级 |
|
|
||||||
|
|
||||||
## License
|
### 功能文档
|
||||||
|
|
||||||
[MIT](LICENSE) © XEngineers
|
| 文档 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| [03-技术选型](docs/03-技术选型.md) | AI 服务栈、持久化层、前端边缘处理选型 |
|
||||||
|
| [04-用户故事](docs/04-用户故事.md) | 用户场景与优先级 |
|
||||||
|
| [05-语音交互](docs/05-语音交互.md) | VAD → STT → LLM → TTS 全链路 |
|
||||||
|
| [06-视觉理解](docs/06-视觉理解.md) | 帧采样、关键帧检测、多模态输入 |
|
||||||
|
| [07-成本控制](docs/07-成本控制.md) | 采样策略、端云协同、模型分级 |
|
||||||
|
| [09-情景切换](docs/09-情景切换.md) | 情景模式设计与实现 |
|
||||||
|
| [12-自定义情景](docs/12-自定义情景.md) | 用户自定义情景功能(规划中) |
|
||||||
|
|
||||||
|
|
||||||
|
## 🐛 问题反馈
|
||||||
|
|
||||||
|
遇到问题?请提交 [Issue](https://github.com/yourusername/CamTalk/issues),并提供以下信息:
|
||||||
|
|
||||||
|
- 操作系统版本
|
||||||
|
- Go / Node.js 版本
|
||||||
|
- 错误日志(后端日志 + 浏览器控制台)
|
||||||
|
- 复现步骤
|
||||||
|
|
||||||
|
## 📝 版权声明
|
||||||
|
|
||||||
|
MIT License © 2024 XEngineers
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
|
||||||
|
**Built with ❤️ using Go, React, and AI**
|
||||||
|
|
||||||
|
[⬆️ 回到顶部](#camtalk)
|
||||||
|
|
||||||
|
</div>
|
||||||
|
|||||||
32
backend/.env.example
Normal file
32
backend/.env.example
Normal file
@@ -0,0 +1,32 @@
|
|||||||
|
# 运行环境
|
||||||
|
# dev:本地开发环境(debug 日志、关闭限流、允许所有 CORS)
|
||||||
|
# prod:生产环境(info 日志、启用限流、严格 CORS 白名单)
|
||||||
|
# 本地开发保持 dev,生产部署会被 docker-compose.yml 覆盖为 prod
|
||||||
|
APP_ENV=dev
|
||||||
|
|
||||||
|
# AI 服务 API Key
|
||||||
|
CAMTALK_AI_STT_API_KEY=sk-your-stt-key
|
||||||
|
CAMTALK_AI_LLM_API_KEY=sk-your-llm-key
|
||||||
|
CAMTALK_AI_TTS_API_KEY=sk-your-tts-key
|
||||||
|
|
||||||
|
# JWT 认证
|
||||||
|
CAMTALK_AUTH_JWT_SECRET=your-jwt-secret-here
|
||||||
|
|
||||||
|
# 三级存储配置
|
||||||
|
# L2: Redis(热数据分布式会话层)
|
||||||
|
CAMTALK_STORAGE_REDIS_ENABLED=true
|
||||||
|
# 开发时填写远程服务器地址,部署时 docker-compose 会覆盖为容器内网地址
|
||||||
|
CAMTALK_REDIS_ADDR=your-remote-server:6379
|
||||||
|
CAMTALK_REDIS_PASSWORD=your-redis-password
|
||||||
|
|
||||||
|
# L3: PostgreSQL(冷数据持久化层)
|
||||||
|
CAMTALK_STORAGE_PERSISTENCE_ENABLED=true
|
||||||
|
CAMTALK_STORAGE_PERSISTENCE_DRIVER=postgres
|
||||||
|
POSTGRES_USER=camtalk
|
||||||
|
POSTGRES_PASSWORD=your-postgres-password
|
||||||
|
# 开发时填写远程服务器地址,部署时 docker-compose 会覆盖为容器内网地址
|
||||||
|
CAMTALK_STORAGE_DSN=postgres://camtalk:your-postgres-password@your-remote-server:5432/camtalk?sslmode=disable
|
||||||
|
|
||||||
|
# 可选覆盖(默认值见 config.yaml)
|
||||||
|
# CAMTALK_SERVER_PORT=8080
|
||||||
|
# CAMTALK_LOG_LEVEL=info
|
||||||
3
backend/.gitignore
vendored
3
backend/.gitignore
vendored
@@ -3,8 +3,7 @@
|
|||||||
bin/
|
bin/
|
||||||
|
|
||||||
# 环境配置
|
# 环境配置
|
||||||
config.dev.yaml
|
.env
|
||||||
config.prod.yaml
|
|
||||||
|
|
||||||
# 临时文件
|
# 临时文件
|
||||||
tmp/
|
tmp/
|
||||||
|
|||||||
@@ -8,11 +8,19 @@ ENV GOPROXY=https://goproxy.cn,https://goproxy.io,direct
|
|||||||
|
|
||||||
# 先复制依赖清单,利用 Docker 缓存层
|
# 先复制依赖清单,利用 Docker 缓存层
|
||||||
COPY go.mod go.sum ./
|
COPY go.mod go.sum ./
|
||||||
RUN go mod download
|
|
||||||
|
# --mount=type=cache 复用 Go module 缓存,依赖不变时跳过下载
|
||||||
|
RUN --mount=type=cache,target=/go/pkg/mod \
|
||||||
|
go mod download
|
||||||
|
|
||||||
# 复制源码并构建
|
# 复制源码并构建
|
||||||
COPY . .
|
COPY . .
|
||||||
RUN CGO_ENABLED=0 GOOS=linux go build -o /camtalk ./cmd/server
|
|
||||||
|
# 复用 module 缓存 + 编译缓存;-ldflags="-s -w" 裁剪符号表减小 ~30% 体积
|
||||||
|
RUN --mount=type=cache,target=/go/pkg/mod \
|
||||||
|
--mount=type=cache,target=/root/.cache/go-build \
|
||||||
|
CGO_ENABLED=0 GOOS=linux \
|
||||||
|
go build -ldflags="-s -w" -o /camtalk ./cmd/server
|
||||||
|
|
||||||
# ---- 运行阶段 ----
|
# ---- 运行阶段 ----
|
||||||
FROM alpine:3.20
|
FROM alpine:3.20
|
||||||
@@ -21,9 +29,9 @@ RUN apk add --no-cache ca-certificates tzdata
|
|||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# 复制二进制和配置
|
# 复制二进制和配置文件(敏感配置通过 docker-compose env_file 注入覆盖)
|
||||||
COPY --from=builder /camtalk .
|
COPY --from=builder /camtalk .
|
||||||
COPY config.yaml .
|
COPY config/ ./config/
|
||||||
|
|
||||||
EXPOSE 8080
|
EXPOSE 8080
|
||||||
|
|
||||||
|
|||||||
@@ -10,17 +10,20 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/api"
|
"github.com/hhs/camtalk/internal/api"
|
||||||
"github.com/hhs/camtalk/internal/auth"
|
"github.com/hhs/camtalk/internal/auth"
|
||||||
"github.com/hhs/camtalk/internal/ai/llm"
|
|
||||||
"github.com/hhs/camtalk/internal/ai/stt"
|
"github.com/hhs/camtalk/internal/ai/stt"
|
||||||
"github.com/hhs/camtalk/internal/ai/tts"
|
"github.com/hhs/camtalk/internal/ai/tts"
|
||||||
"github.com/hhs/camtalk/internal/config"
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
eino "github.com/hhs/camtalk/internal/eino"
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
"github.com/hhs/camtalk/internal/orchestrator"
|
"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/trace"
|
||||||
"github.com/hhs/camtalk/internal/ws"
|
"github.com/hhs/camtalk/internal/ws"
|
||||||
migrations "github.com/hhs/camtalk/migrations"
|
migrations "github.com/hhs/camtalk/migrations"
|
||||||
)
|
)
|
||||||
@@ -32,8 +35,8 @@ var Version string
|
|||||||
var startTime = time.Now()
|
var startTime = time.Now()
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
// 加载配置
|
// 加载配置(工作目录用于定位 .env 和 config.yaml)
|
||||||
cfg, err := config.Load()
|
cfg, err := config.Load(".")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic("failed to load config: " + err.Error())
|
panic("failed to load config: " + err.Error())
|
||||||
}
|
}
|
||||||
@@ -47,20 +50,27 @@ func main() {
|
|||||||
"addr", cfg.Server.Addr(),
|
"addr", cfg.Server.Addr(),
|
||||||
)
|
)
|
||||||
|
|
||||||
// 初始化存储层(条件初始化 PostgreSQL)
|
// 初始化存储层(三级存储架构:L1 内存 → L2 Redis → L3 PostgreSQL)
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
var userRepo store.UserRepository
|
var userRepo store.UserRepository
|
||||||
var msgRepo store.MessageRepository
|
var msgRepo store.MessageRepository
|
||||||
var sessRepo store.SessionRepository
|
var sessRepo store.SessionRepository
|
||||||
|
var pool *pgxpool.Pool // 数据库连接池
|
||||||
|
|
||||||
if cfg.Storage.Driver == "postgres" {
|
// L3: PostgreSQL(冷数据持久化层)
|
||||||
if cfg.Storage.DSN == "" {
|
dsn := cfg.Storage.Persistence.DSN
|
||||||
logger.Log.Fatalw("storage.dsn is required when storage.driver is postgres",
|
if dsn == "" {
|
||||||
|
dsn = cfg.Storage.DSN // 兼容旧配置
|
||||||
|
}
|
||||||
|
if cfg.Storage.Persistence.Enabled && cfg.Storage.Persistence.Driver == "postgres" {
|
||||||
|
if dsn == "" {
|
||||||
|
logger.Log.Fatalw("storage.persistence.dsn is required when persistence is enabled",
|
||||||
"hint", "set CAMTALK_STORAGE_DSN environment variable")
|
"hint", "set CAMTALK_STORAGE_DSN environment variable")
|
||||||
}
|
}
|
||||||
pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN)
|
var err error
|
||||||
|
pool, err = store.NewPostgresPool(ctx, dsn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Log.Fatalw("failed to connect to postgres", "error", err)
|
logger.Log.Fatalw("failed to connect to postgres", "error", err)
|
||||||
}
|
}
|
||||||
@@ -74,27 +84,76 @@ func main() {
|
|||||||
userRepo = store.NewPgUserRepository(pool)
|
userRepo = store.NewPgUserRepository(pool)
|
||||||
msgRepo = store.NewPgMessageRepository(pool)
|
msgRepo = store.NewPgMessageRepository(pool)
|
||||||
sessRepo = store.NewPgSessionRepository(pool)
|
sessRepo = store.NewPgSessionRepository(pool)
|
||||||
logger.Log.Infow("postgres storage initialized", "driver", cfg.Storage.Driver)
|
logger.Log.Infow("L3 PostgreSQL storage initialized", "driver", cfg.Storage.Persistence.Driver)
|
||||||
} else {
|
} else {
|
||||||
userRepo = store.NewMemUserRepository()
|
userRepo = store.NewMemUserRepository()
|
||||||
logger.Log.Info("using in-memory storage")
|
logger.Log.Info("using in-memory user storage")
|
||||||
}
|
}
|
||||||
|
|
||||||
// 初始化 Session Manager
|
// L2: Redis(热数据分布式会话层)
|
||||||
|
var rdb *redis.Client
|
||||||
|
var redisMgr *session.RedisManager
|
||||||
|
if cfg.Storage.Redis.Enabled {
|
||||||
|
rdb = redis.NewClient(&redis.Options{
|
||||||
|
Addr: cfg.Redis.Addr,
|
||||||
|
Password: cfg.Redis.Password,
|
||||||
|
DB: cfg.Redis.DB,
|
||||||
|
})
|
||||||
|
// 验证 Redis 连接
|
||||||
|
if err := rdb.Ping(ctx).Err(); err != nil {
|
||||||
|
logger.Log.Fatalw("failed to connect to redis", "error", err)
|
||||||
|
}
|
||||||
|
redisMgr = session.NewRedisManager(
|
||||||
|
rdb,
|
||||||
|
time.Duration(cfg.Session.TTL)*time.Minute,
|
||||||
|
cfg.Session.MaxHistory,
|
||||||
|
)
|
||||||
|
// 包装 userRepo 为带 Redis 缓存的版本(refresh token 二级缓存)
|
||||||
|
userRepo = store.NewCachedUserRepository(userRepo, rdb, time.Duration(cfg.Auth.RefreshTTL)*time.Minute)
|
||||||
|
logger.Log.Infow("L2 Redis storage initialized",
|
||||||
|
"addr", cfg.Redis.Addr,
|
||||||
|
"db", cfg.Redis.DB,
|
||||||
|
"cached_user_repo", true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 初始化 Session Manager(三级存储)
|
||||||
var sessionMgr session.Manager
|
var sessionMgr session.Manager
|
||||||
var sessionOpts []session.Option
|
if cfg.Storage.Redis.Enabled {
|
||||||
if msgRepo != nil {
|
// L1 + L2 + L3 三级存储
|
||||||
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
|
var tieredOpts []session.TieredOption
|
||||||
|
if sessRepo != nil {
|
||||||
|
tieredOpts = append(tieredOpts, session.WithTieredSessionRepository(sessRepo))
|
||||||
|
}
|
||||||
|
if msgRepo != nil {
|
||||||
|
tieredOpts = append(tieredOpts, session.WithTieredMessageRepository(msgRepo))
|
||||||
|
}
|
||||||
|
tieredMgr := session.NewTieredManager(
|
||||||
|
time.Duration(cfg.Session.TTL)*time.Minute,
|
||||||
|
cfg.Session.MaxHistory,
|
||||||
|
redisMgr,
|
||||||
|
tieredOpts...,
|
||||||
|
)
|
||||||
|
sessionMgr = tieredMgr
|
||||||
|
defer tieredMgr.Stop()
|
||||||
|
logger.Log.Info("session manager initialized with L1+L2+L3 tiered storage")
|
||||||
|
} else {
|
||||||
|
// L1 + L3 两级存储(无 Redis)
|
||||||
|
var sessionOpts []session.Option
|
||||||
|
if msgRepo != nil {
|
||||||
|
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
|
||||||
|
}
|
||||||
|
if sessRepo != nil {
|
||||||
|
sessionOpts = append(sessionOpts, session.WithSessionRepository(sessRepo))
|
||||||
|
}
|
||||||
|
memMgr := session.NewMemoryManager(
|
||||||
|
time.Duration(cfg.Session.TTL)*time.Minute,
|
||||||
|
cfg.Session.MaxHistory,
|
||||||
|
sessionOpts...,
|
||||||
|
)
|
||||||
|
sessionMgr = memMgr
|
||||||
|
defer memMgr.Stop()
|
||||||
|
logger.Log.Info("session manager initialized with L1+L3 storage (Redis disabled)")
|
||||||
}
|
}
|
||||||
if sessRepo != nil {
|
|
||||||
sessionOpts = append(sessionOpts, session.WithSessionRepository(sessRepo))
|
|
||||||
}
|
|
||||||
sessionMgr = session.NewMemoryManager(
|
|
||||||
time.Duration(cfg.Session.TTL)*time.Minute,
|
|
||||||
cfg.Session.MaxHistory,
|
|
||||||
sessionOpts...,
|
|
||||||
)
|
|
||||||
defer sessionMgr.(*session.MemoryManager).Stop()
|
|
||||||
|
|
||||||
// 初始化 AI 服务
|
// 初始化 AI 服务
|
||||||
logger.Log.Infow("initializing AI services",
|
logger.Log.Infow("initializing AI services",
|
||||||
@@ -116,9 +175,6 @@ func main() {
|
|||||||
sttService = stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, cfg.AI.STT.Timeout, logger.Log)
|
sttService = stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, cfg.AI.STT.Timeout, logger.Log)
|
||||||
logger.Log.Infow("STT service initialized", "provider", "deepgram", "model", cfg.AI.STT.Model)
|
logger.Log.Infow("STT service initialized", "provider", "deepgram", "model", cfg.AI.STT.Model)
|
||||||
}
|
}
|
||||||
llmService := llm.NewOpenAIService(cfg.AI.LLM.APIKey, cfg.AI.LLM.Model, cfg.AI.LLM.Endpoint, cfg.AI.LLM.Timeout, cfg.AI.LLM.HTTPClientTimeout, logger.Log)
|
|
||||||
logger.Log.Infow("LLM service initialized", "provider", cfg.AI.LLM.Provider, "model", cfg.AI.LLM.Model, "endpoint", cfg.AI.LLM.Endpoint, "timeout", cfg.AI.LLM.Timeout)
|
|
||||||
|
|
||||||
var ttsService tts.Service
|
var ttsService tts.Service
|
||||||
switch strings.ToLower(cfg.AI.TTS.Provider) {
|
switch strings.ToLower(cfg.AI.TTS.Provider) {
|
||||||
case "mimo", "xiaomi":
|
case "mimo", "xiaomi":
|
||||||
@@ -129,8 +185,16 @@ func main() {
|
|||||||
logger.Log.Infow("TTS service initialized", "provider", "openai", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "speed", cfg.AI.TTS.Speed)
|
logger.Log.Infow("TTS service initialized", "provider", "openai", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "speed", cfg.AI.TTS.Speed)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 初始化 Orchestrator
|
// 初始化 Eino Graph + Orchestrator
|
||||||
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr, cfg)
|
var userScenarioRepo store.UserScenarioRepository
|
||||||
|
if pool != nil {
|
||||||
|
userScenarioRepo = store.NewPostgresUserScenarioRepo(pool)
|
||||||
|
}
|
||||||
|
pipelineGraph, err := eino.NewPipelineGraph(ctx, cfg, sttService, ttsService, sessionMgr, userScenarioRepo)
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Fatalw("failed to create eino pipeline graph", "error", err)
|
||||||
|
}
|
||||||
|
orch := eino.NewEinoOrchestrator(pipelineGraph, sessionMgr, cfg.AI.LLM.Model)
|
||||||
|
|
||||||
// 初始化认证服务
|
// 初始化认证服务
|
||||||
tokenMgr := auth.NewTokenManager(
|
tokenMgr := auth.NewTokenManager(
|
||||||
@@ -140,13 +204,32 @@ 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
r := gin.New()
|
r := gin.New()
|
||||||
r.Use(gin.Recovery())
|
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
|
||||||
|
r.Use(trace.GinLogger()) // 第二层:记录请求
|
||||||
|
r.Use(trace.GinRecovery()) // 第三层:panic 恢复
|
||||||
|
|
||||||
// REST API
|
// REST API
|
||||||
apiGroup := r.Group("/api")
|
apiGroup := r.Group("/api")
|
||||||
@@ -160,14 +243,29 @@ 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)
|
||||||
|
|
||||||
|
// UserScenario REST 端点
|
||||||
|
if pool != nil {
|
||||||
|
userScenarioRepo := store.NewPostgresUserScenarioRepo(pool)
|
||||||
|
userScenarioHandler := api.NewUserScenarioHandler(userScenarioRepo)
|
||||||
|
scenarioGroup := apiGroup.Group("/scenarios")
|
||||||
|
scenarioGroup.Use(auth.AuthMiddleware(tokenMgr))
|
||||||
|
{
|
||||||
|
scenarioGroup.GET("", userScenarioHandler.List)
|
||||||
|
scenarioGroup.POST("", userScenarioHandler.Create)
|
||||||
|
scenarioGroup.GET("/:id", userScenarioHandler.Get)
|
||||||
|
scenarioGroup.PATCH("/:id", userScenarioHandler.Update)
|
||||||
|
scenarioGroup.DELETE("/:id", userScenarioHandler.Delete)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// WebSocket
|
// WebSocket
|
||||||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter, userScenarioRepo))
|
||||||
|
|
||||||
// HTTP Server
|
// HTTP Server
|
||||||
srv := &http.Server{
|
srv := &http.Server{
|
||||||
|
|||||||
@@ -1,42 +0,0 @@
|
|||||||
# config.yaml — 默认配置
|
|
||||||
app:
|
|
||||||
env: dev
|
|
||||||
|
|
||||||
server:
|
|
||||||
host: "0.0.0.0"
|
|
||||||
port: 8080
|
|
||||||
read_timeout: 30
|
|
||||||
write_timeout: 30
|
|
||||||
|
|
||||||
redis:
|
|
||||||
addr: "localhost:6379"
|
|
||||||
password: ""
|
|
||||||
db: 0
|
|
||||||
|
|
||||||
ai:
|
|
||||||
stt:
|
|
||||||
provider: mimo
|
|
||||||
model: mimo-v2.5-asr
|
|
||||||
endpoint: "https://api.xiaomimimo.com/v1"
|
|
||||||
api_key: "sk-c3jhv58rr5djhxw398w2rrij5tfpnpdgxqq1bojagshzviah"
|
|
||||||
llm:
|
|
||||||
provider: dashscope
|
|
||||||
model: qwen3-vl-plus
|
|
||||||
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
|
||||||
api_key: "sk-ws-H.REHELLY.C4s3.MEUCIQCRee37XWEKp2szaxVLFDtR1rxNNsf372zMvCR0Xl6UvQIgZgvhRTvaa1FmhbCQJgaHu4Jny29AQkn01-3hX9CWBOg"
|
|
||||||
timeout: 30
|
|
||||||
tts:
|
|
||||||
provider: mimo
|
|
||||||
model: mimo-v2.5-tts
|
|
||||||
voice: mimo_default
|
|
||||||
speed: 1.0
|
|
||||||
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
|
||||||
api_key: "tp-c9e7scwfx94qvqyhpnahnw8uaiya01za2qzvg4xe24rp3xiv"
|
|
||||||
timeout: 5
|
|
||||||
|
|
||||||
storage:
|
|
||||||
driver: memory
|
|
||||||
|
|
||||||
log:
|
|
||||||
level: info
|
|
||||||
format: console
|
|
||||||
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 格式,便于日志收集和分析
|
||||||
80
backend/config/config.yaml
Normal file
80
backend/config/config.yaml
Normal file
@@ -0,0 +1,80 @@
|
|||||||
|
# CamTalk 后端配置
|
||||||
|
|
||||||
|
app:
|
||||||
|
env: dev # dev / prod,可通过 APP_ENV 环境变量覆盖
|
||||||
|
|
||||||
|
server:
|
||||||
|
host: "0.0.0.0"
|
||||||
|
port: 8080
|
||||||
|
read_timeout: 30 # 秒
|
||||||
|
write_timeout: 30 # 秒
|
||||||
|
shutdown_timeout: 10 # 优雅关闭超时(秒)
|
||||||
|
heartbeat_interval: 30 # 心跳检查间隔(秒)
|
||||||
|
heartbeat_timeout: 60 # 心跳超时断开(秒)
|
||||||
|
allowed_origins: [] # CORS 白名单,空=允许所有
|
||||||
|
|
||||||
|
session:
|
||||||
|
ttl: 30 # 会话过期时间(分钟)
|
||||||
|
max_history: 20 # 对话历史上限(条)
|
||||||
|
|
||||||
|
ai:
|
||||||
|
stt:
|
||||||
|
provider: mimo # mimo / deepgram
|
||||||
|
model: mimo-v2.5-asr
|
||||||
|
endpoint: "https://api.xiaomimimo.com/v1"
|
||||||
|
timeout: 5 # STT 请求超时(秒)
|
||||||
|
http_client_timeout: 30 # HTTP 客户端超时(秒)
|
||||||
|
llm:
|
||||||
|
provider: dashscope # dashscope / openai
|
||||||
|
model: qwen3-vl-plus
|
||||||
|
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
timeout: 30 # LLM 请求超时(秒)
|
||||||
|
http_client_timeout: 60 # HTTP 客户端超时(秒)
|
||||||
|
tts:
|
||||||
|
provider: mimo # mimo / openai
|
||||||
|
model: mimo-v2.5-tts
|
||||||
|
voice: mimo_default
|
||||||
|
speed: 1.0
|
||||||
|
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||||
|
timeout: 5 # TTS 请求超时(秒)
|
||||||
|
http_client_timeout: 30 # HTTP 客户端超时(秒)
|
||||||
|
output_format: mp3 # 输出格式:mp3 / wav
|
||||||
|
sample_rate: 24000 # 输出采样率
|
||||||
|
|
||||||
|
storage:
|
||||||
|
# 三级存储架构:L1 内存 → L2 Redis → L3 PostgreSQL
|
||||||
|
redis:
|
||||||
|
enabled: true # 是否启用 Redis(L2 热数据层)
|
||||||
|
persistence:
|
||||||
|
enabled: true # 是否启用持久化(L3 冷数据层)
|
||||||
|
driver: postgres # postgres
|
||||||
|
# dsn 通过环境变量 CAMTALK_STORAGE_DSN 设置
|
||||||
|
|
||||||
|
redis:
|
||||||
|
addr: "localhost:6379"
|
||||||
|
password: ""
|
||||||
|
db: 0
|
||||||
|
|
||||||
|
auth:
|
||||||
|
# jwt_secret 通过环境变量 CAMTALK_AUTH_JWT_SECRET 设置
|
||||||
|
access_ttl: 120 # Access Token 过期时间(分钟)
|
||||||
|
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:
|
||||||
|
level: info # debug / info / warn / error
|
||||||
|
format: console # console / json
|
||||||
@@ -3,23 +3,36 @@ module github.com/hhs/camtalk
|
|||||||
go 1.25.0
|
go 1.25.0
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
github.com/alicebob/miniredis/v2 v2.38.0
|
||||||
|
github.com/cloudwego/eino v0.9.9
|
||||||
|
github.com/cloudwego/eino-ext/components/model/openai v0.1.13
|
||||||
github.com/gin-gonic/gin v1.10.0
|
github.com/gin-gonic/gin v1.10.0
|
||||||
|
github.com/golang-jwt/jwt/v5 v5.3.1
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/gorilla/websocket v1.5.3
|
github.com/gorilla/websocket v1.5.3
|
||||||
github.com/jackc/pgx/v5 v5.10.0
|
github.com/jackc/pgx/v5 v5.10.0
|
||||||
|
github.com/joho/godotenv v1.5.1
|
||||||
|
github.com/oklog/ulid/v2 v2.1.1
|
||||||
github.com/redis/go-redis/v9 v9.20.1
|
github.com/redis/go-redis/v9 v9.20.1
|
||||||
github.com/spf13/viper v1.21.0
|
github.com/spf13/viper v1.21.0
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
go.uber.org/zap v1.28.0
|
go.uber.org/zap v1.28.0
|
||||||
|
golang.org/x/crypto v0.31.0
|
||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/bytedance/sonic v1.11.6 // indirect
|
github.com/bahlo/generic-list-go v0.2.0 // indirect
|
||||||
github.com/bytedance/sonic/loader v0.1.1 // indirect
|
github.com/buger/jsonparser v1.1.1 // indirect
|
||||||
|
github.com/bytedance/gopkg v0.1.3 // indirect
|
||||||
|
github.com/bytedance/sonic v1.15.0 // indirect
|
||||||
|
github.com/bytedance/sonic/loader v0.5.0 // indirect
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
github.com/cloudwego/base64x v0.1.4 // indirect
|
github.com/cloudwego/base64x v0.1.6 // indirect
|
||||||
github.com/cloudwego/iasm v0.2.0 // indirect
|
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.17 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
|
github.com/eino-contrib/jsonschema v1.0.3 // indirect
|
||||||
|
github.com/evanphx/json-patch v0.5.2 // indirect
|
||||||
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
||||||
github.com/gabriel-vasile/mimetype v1.4.3 // indirect
|
github.com/gabriel-vasile/mimetype v1.4.3 // indirect
|
||||||
github.com/gin-contrib/sse v0.1.0 // indirect
|
github.com/gin-contrib/sse v0.1.0 // indirect
|
||||||
@@ -28,19 +41,25 @@ require (
|
|||||||
github.com/go-playground/validator/v10 v10.20.0 // indirect
|
github.com/go-playground/validator/v10 v10.20.0 // indirect
|
||||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
||||||
github.com/goccy/go-json v0.10.2 // indirect
|
github.com/goccy/go-json v0.10.2 // indirect
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1 // indirect
|
github.com/goph/emperror v0.17.2 // indirect
|
||||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||||
github.com/json-iterator/go v1.1.12 // indirect
|
github.com/json-iterator/go v1.1.12 // indirect
|
||||||
github.com/klauspost/cpuid/v2 v2.2.10 // indirect
|
github.com/klauspost/cpuid/v2 v2.2.10 // indirect
|
||||||
github.com/leodido/go-urn v1.4.0 // indirect
|
github.com/leodido/go-urn v1.4.0 // indirect
|
||||||
|
github.com/mailru/easyjson v0.7.7 // indirect
|
||||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||||
|
github.com/meguminnnnnnnnn/go-openai v0.1.2 // indirect
|
||||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
|
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
|
||||||
github.com/modern-go/reflect2 v1.0.2 // indirect
|
github.com/modern-go/reflect2 v1.0.2 // indirect
|
||||||
|
github.com/nikolalohinski/gonja v1.5.3 // indirect
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||||
|
github.com/pkg/errors v0.9.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/sagikazarmark/locafero v0.11.0 // indirect
|
github.com/sagikazarmark/locafero v0.11.0 // indirect
|
||||||
|
github.com/sirupsen/logrus v1.9.3 // indirect
|
||||||
|
github.com/slongfield/pyfmt v0.0.0-20220222012616-ea85ff4c361f // indirect
|
||||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
|
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
|
||||||
github.com/spf13/afero v1.15.0 // indirect
|
github.com/spf13/afero v1.15.0 // indirect
|
||||||
github.com/spf13/cast v1.10.0 // indirect
|
github.com/spf13/cast v1.10.0 // indirect
|
||||||
@@ -49,11 +68,14 @@ require (
|
|||||||
github.com/subosito/gotenv v1.6.0 // indirect
|
github.com/subosito/gotenv v1.6.0 // indirect
|
||||||
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
|
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
|
||||||
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/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
|
||||||
golang.org/x/arch v0.8.0 // indirect
|
golang.org/x/arch v0.11.0 // indirect
|
||||||
golang.org/x/crypto v0.23.0 // indirect
|
golang.org/x/exp v0.0.0-20230713183714-613f0c0eb8a1 // indirect
|
||||||
golang.org/x/net v0.25.0 // indirect
|
golang.org/x/net v0.25.0 // indirect
|
||||||
golang.org/x/sync v0.17.0 // indirect
|
golang.org/x/sync v0.17.0 // indirect
|
||||||
golang.org/x/sys v0.30.0 // indirect
|
golang.org/x/sys v0.30.0 // indirect
|
||||||
|
|||||||
135
backend/go.sum
135
backend/go.sum
@@ -1,30 +1,60 @@
|
|||||||
|
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/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg=
|
||||||
|
github.com/bitly/go-simplejson v0.5.0/go.mod h1:cXHtHw4XUPsvGaxgjIAn8PhEWG9NfngEKAMDJEczWVA=
|
||||||
|
github.com/bmizerany/assert v0.0.0-20160611221934-b7ed37b82869/go.mod h1:Ekp36dRnpXw/yCqJaO+ZrUyxD+3VXMFFr56k5XYrpB4=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
|
||||||
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
|
||||||
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
|
||||||
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
|
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
|
||||||
github.com/bytedance/sonic v1.11.6 h1:oUp34TzMlL+OY1OUWxHqsdkgC/Zfc85zGqw9siXjrc0=
|
github.com/buger/jsonparser v1.1.1 h1:2PnMjfWD7wBILjqQbt530v576A/cAbQvEW9gGIpYMUs=
|
||||||
github.com/bytedance/sonic v1.11.6/go.mod h1:LysEHSvpvDySVdC2f87zGWf6CIKJcAvqab1ZaiQtds4=
|
github.com/buger/jsonparser v1.1.1/go.mod h1:6RYKKt7H4d4+iWqouImQ9R2FZql3VbhNgx27UK13J/0=
|
||||||
github.com/bytedance/sonic/loader v0.1.1 h1:c+e5Pt1k/cy5wMveRDyk2X4B9hF4g7an8N3zCYjJFNM=
|
github.com/bugsnag/bugsnag-go v1.4.0/go.mod h1:2oa8nejYd4cQ/b0hMIopN0lCRxU0bueqREvZLWFrtK8=
|
||||||
github.com/bytedance/sonic/loader v0.1.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU=
|
github.com/bugsnag/panicwrap v1.2.0/go.mod h1:D/8v3kj0zr8ZAKg1AQ6crr+5VwKN5eIywRkfhyM/+dE=
|
||||||
|
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
|
||||||
|
github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM=
|
||||||
|
github.com/bytedance/mockey v1.3.0 h1:ONLRdvhqmCfr9rTasUB8ZKCfvbdD2tohOg4u+4Q/ed0=
|
||||||
|
github.com/bytedance/mockey v1.3.0/go.mod h1:1BPHF9sol5R1ud/+0VEHGQq/+i2lN+GTsr3O2Q9IENY=
|
||||||
|
github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE=
|
||||||
|
github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k=
|
||||||
|
github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE=
|
||||||
|
github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo=
|
||||||
|
github.com/certifi/gocertifi v0.0.0-20190105021004-abcd57078448/go.mod h1:GJKEexRPVJrBSOjoqN5VNOIKJ5Q3RViH6eu3puDRwx4=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
github.com/cloudwego/base64x v0.1.4 h1:jwCgWpFanWmN8xoIUHa2rtzmkd5J2plF/dnLS6Xd/0Y=
|
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
|
||||||
github.com/cloudwego/base64x v0.1.4/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w=
|
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
|
||||||
github.com/cloudwego/iasm v0.2.0 h1:1KNIy1I1H9hNNFEEH3DVnI4UujN+1zjpuk6gwHLTssg=
|
github.com/cloudwego/eino v0.9.9 h1:x63hvRif6ANPh9YEPoTIrp1potEeoLQFAjOclKaX/Kg=
|
||||||
github.com/cloudwego/iasm v0.2.0/go.mod h1:8rXZaNYT2n95jn+zTI1sDr+IgcD2GVs0nlbbQPiEFhY=
|
github.com/cloudwego/eino v0.9.9/go.mod h1:OBD1mrkfkt/pJa4rkg1P0VnaMeOVl7l8IAdEqY//3IQ=
|
||||||
|
github.com/cloudwego/eino-ext/components/model/openai v0.1.13 h1:5XHRTiTD5bt9KQrMHcfvuWNklEC3tpm3XHejdozt9vM=
|
||||||
|
github.com/cloudwego/eino-ext/components/model/openai v0.1.13/go.mod h1:mgIoqYYOc0eECCqvLbEYpOJrQNTNxkwXzSJzFU+v5sQ=
|
||||||
|
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.17 h1:EeVcR1TslRA2IdNW1h/2LaGbPlffwGhQm99jM3zWZiI=
|
||||||
|
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.17/go.mod h1:Zkcx6DPTR2NfWmtSXbhItswGw6hqUezNPhNcke0pOG8=
|
||||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
|
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||||
|
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||||
|
github.com/eino-contrib/jsonschema v1.0.3 h1:2Kfsm1xlMV0ssY2nuxshS4AwbLFuqmPmzIjLVJ1Fsp0=
|
||||||
|
github.com/eino-contrib/jsonschema v1.0.3/go.mod h1:cpnX4SyKjWjGC7iN2EbhxaTdLqGjCi0e9DxpLYxddD4=
|
||||||
|
github.com/evanphx/json-patch v0.5.2 h1:xVCHIVMUu1wtM/VkR9jVZ45N3FhZfYMMYGorLCR8P3k=
|
||||||
|
github.com/evanphx/json-patch v0.5.2/go.mod h1:ZWS5hhDbVDyob71nXKNL0+PWn6ToqBHMikGIFbs31qQ=
|
||||||
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
|
||||||
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
|
||||||
|
github.com/fsnotify/fsnotify v1.4.7/go.mod h1:jwhsz4b93w/PPRr/qN1Yymfu8t87LnFCMoQvtojpjFo=
|
||||||
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
|
||||||
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0=
|
||||||
github.com/gabriel-vasile/mimetype v1.4.3 h1:in2uUcidCuFcDKtdcBxlR0rJ1+fsokWf+uqxgUFjbI0=
|
github.com/gabriel-vasile/mimetype v1.4.3 h1:in2uUcidCuFcDKtdcBxlR0rJ1+fsokWf+uqxgUFjbI0=
|
||||||
github.com/gabriel-vasile/mimetype v1.4.3/go.mod h1:d8uq/6HKRL6CGdk+aubisF/M5GcPfT7nKyLpA0lbSSk=
|
github.com/gabriel-vasile/mimetype v1.4.3/go.mod h1:d8uq/6HKRL6CGdk+aubisF/M5GcPfT7nKyLpA0lbSSk=
|
||||||
|
github.com/getsentry/raven-go v0.2.0/go.mod h1:KungGk8q33+aIAZUIVWZDr2OfAEBsO49PX4NzFV5kcQ=
|
||||||
github.com/gin-contrib/sse v0.1.0 h1:Y/yl/+YNO8GZSjAhjMsSuLt29uWRFHdHYUb5lYOV9qE=
|
github.com/gin-contrib/sse v0.1.0 h1:Y/yl/+YNO8GZSjAhjMsSuLt29uWRFHdHYUb5lYOV9qE=
|
||||||
github.com/gin-contrib/sse v0.1.0/go.mod h1:RHrZQHXnP2xjPF+u1gW/2HnVO7nvIa9PG3Gm+fLHvGI=
|
github.com/gin-contrib/sse v0.1.0/go.mod h1:RHrZQHXnP2xjPF+u1gW/2HnVO7nvIa9PG3Gm+fLHvGI=
|
||||||
github.com/gin-gonic/gin v1.10.0 h1:nTuyha1TYqgedzytsKYqna+DfLos46nTv2ygFy86HFU=
|
github.com/gin-gonic/gin v1.10.0 h1:nTuyha1TYqgedzytsKYqna+DfLos46nTv2ygFy86HFU=
|
||||||
github.com/gin-gonic/gin v1.10.0/go.mod h1:4PMNQiOhvDRa013RKVbsiNwoyezlm2rm0uX/T7kzp5Y=
|
github.com/gin-gonic/gin v1.10.0/go.mod h1:4PMNQiOhvDRa013RKVbsiNwoyezlm2rm0uX/T7kzp5Y=
|
||||||
|
github.com/go-check/check v0.0.0-20180628173108-788fd7840127 h1:0gkP6mzaMqkmpcJYCFOLkIBwI7xFExG03bbkOkCvUPI=
|
||||||
|
github.com/go-check/check v0.0.0-20180628173108-788fd7840127/go.mod h1:9ES+weclKsC9YodN5RgxqK/VD9HM9JsCSh7rNhMZE98=
|
||||||
github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s=
|
github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s=
|
||||||
github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4=
|
github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4=
|
||||||
github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA=
|
github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA=
|
||||||
@@ -37,15 +67,22 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L
|
|||||||
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
|
||||||
github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
|
github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
|
||||||
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
|
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
|
||||||
|
github.com/gofrs/uuid v3.2.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRxnplIgP/c0N/04lM=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||||
|
github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
|
||||||
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
|
||||||
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
|
||||||
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
||||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||||
|
github.com/goph/emperror v0.17.2 h1:yLapQcmEsO0ipe9p5TaN22djm3OFV/TfM/fcYP0/J18=
|
||||||
|
github.com/goph/emperror v0.17.2/go.mod h1:+ZbQ+fUNO/6FNiUo0ujtMjhgad9Xa6fQL9KhH4LNHic=
|
||||||
|
github.com/gopherjs/gopherjs v1.17.2 h1:fQnZVsXk8uxXIStYb0N4bGk7jeyTalG/wsZjQ25dO0g=
|
||||||
|
github.com/gopherjs/gopherjs v1.17.2/go.mod h1:pRRIvn/QzFLrKfvEz3qUuEhtE/zLCWfreZ6J5gM2i+k=
|
||||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
|
github.com/hpcloud/tail v1.0.0/go.mod h1:ab1qPbhIpdTxEkNHXyeSf5vhxWSCs/tWer42PpOxQnU=
|
||||||
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
||||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||||
@@ -54,35 +91,73 @@ github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0=
|
|||||||
github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||||
|
github.com/jessevdk/go-flags v1.4.0/go.mod h1:4FA24M0QyGHXBuZZK/XkWh8h0e1EYbRYJSGM75WSRxI=
|
||||||
|
github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=
|
||||||
|
github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4=
|
||||||
|
github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y=
|
||||||
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
|
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
|
||||||
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
|
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
|
||||||
github.com/klauspost/cpuid/v2 v2.0.9/go.mod h1:FInQzS24/EEf25PyTYn52gqo7WaD8xa0213Md/qVLRg=
|
github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7C0MuV77Wo=
|
||||||
|
github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU=
|
||||||
|
github.com/kardianos/osext v0.0.0-20190222173326-2bc1f35cddc0/go.mod h1:1NbS8ALrpOvjt0rHPNLyCIeMtbizbir8U//inJ+zuB8=
|
||||||
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
|
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
|
||||||
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
|
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
|
||||||
github.com/knz/go-libedit v1.10.1/go.mod h1:MZTVkCWyz0oBc7JOWP3wNAzd002ZbM/5hgShxwh4x8M=
|
github.com/konsorten/go-windows-terminal-sequences v1.0.1/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ=
|
||||||
|
github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo=
|
||||||
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
|
||||||
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
|
||||||
|
github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ=
|
||||||
|
github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
|
||||||
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
|
||||||
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
|
||||||
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
|
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
|
||||||
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
|
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
|
||||||
|
github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0=
|
||||||
|
github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc=
|
||||||
|
github.com/mattn/go-colorable v0.1.2 h1:/bC9yWikZXAL9uJdulbSfyVNIR3n3trXl+v8+1sx8mU=
|
||||||
|
github.com/mattn/go-colorable v0.1.2/go.mod h1:U0ppj6V5qS13XJ6of8GYAs25YV2eR4EVcfRqFIhoBtE=
|
||||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||||
|
github.com/meguminnnnnnnnn/go-openai v0.1.2 h1:iXombGGjqjBrmE9WaSidUhhi3YQhf42QTHvHLMkgvCA=
|
||||||
|
github.com/meguminnnnnnnnn/go-openai v0.1.2/go.mod h1:qs96ysDmxhE4BZoU45I43zcyfnaYxU3X+aRzLko/htY=
|
||||||
|
github.com/mgutz/ansi v0.0.0-20170206155736-9520e82c474b h1:j7+1HpAFS1zy5+Q4qx1fWh90gTKwiN4QCGoY9TWyyO4=
|
||||||
|
github.com/mgutz/ansi v0.0.0-20170206155736-9520e82c474b/go.mod h1:01TrycV0kFyexm33Z7vhZRXopbI8J3TDReVlkTgMUxE=
|
||||||
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
||||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg=
|
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg=
|
||||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
|
||||||
github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
|
github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
|
||||||
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
|
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
|
||||||
|
github.com/nikolalohinski/gonja v1.5.3 h1:GsA+EEaZDZPGJ8JtpeGN78jidhOlxeJROpqMT9fTj9c=
|
||||||
|
github.com/nikolalohinski/gonja v1.5.3/go.mod h1:RmjwxNiXAEqcq1HeK5SSMmqFJvKOfTfXhkJv6YBtPa4=
|
||||||
|
github.com/oklog/ulid/v2 v2.1.1 h1:suPZ4ARWLOJLegGFiZZ1dFAkqzhMjL3J1TzI+5wHz8s=
|
||||||
|
github.com/oklog/ulid/v2 v2.1.1/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNsTT1QQ=
|
||||||
|
github.com/onsi/ginkgo v1.6.0/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+WWjE=
|
||||||
|
github.com/onsi/ginkgo v1.8.0/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+WWjE=
|
||||||
|
github.com/onsi/gomega v1.5.0/go.mod h1:ex+gbHU/CVuBBDIJjb2X0qEXbFg53c61hWP/1CpauHY=
|
||||||
|
github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||||
|
github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
|
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
||||||
|
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||||
github.com/redis/go-redis/v9 v9.20.1 h1:sfCU6A8P3dXbKyWes02uxA2baehGux9dZHfEKtsTB1w=
|
github.com/redis/go-redis/v9 v9.20.1 h1:sfCU6A8P3dXbKyWes02uxA2baehGux9dZHfEKtsTB1w=
|
||||||
github.com/redis/go-redis/v9 v9.20.1/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
github.com/redis/go-redis/v9 v9.20.1/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
|
||||||
github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8=
|
github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8=
|
||||||
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
|
||||||
|
github.com/rollbar/rollbar-go v1.0.2/go.mod h1:AcFs5f0I+c71bpHlXNNDbOWJiKwjFDtISeXco0L5PKQ=
|
||||||
github.com/sagikazarmark/locafero v0.11.0 h1:1iurJgmM9G3PA/I+wWYIOw/5SyBtxapeHDcg+AAIFXc=
|
github.com/sagikazarmark/locafero v0.11.0 h1:1iurJgmM9G3PA/I+wWYIOw/5SyBtxapeHDcg+AAIFXc=
|
||||||
github.com/sagikazarmark/locafero v0.11.0/go.mod h1:nVIGvgyzw595SUSUE6tvCp3YYTeHs15MvlmU87WwIik=
|
github.com/sagikazarmark/locafero v0.11.0/go.mod h1:nVIGvgyzw595SUSUE6tvCp3YYTeHs15MvlmU87WwIik=
|
||||||
|
github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo=
|
||||||
|
github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ=
|
||||||
|
github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ=
|
||||||
|
github.com/slongfield/pyfmt v0.0.0-20220222012616-ea85ff4c361f h1:Z2cODYsUxQPofhpYRMQVwWz4yUVpHF+vPi+eUdruUYI=
|
||||||
|
github.com/slongfield/pyfmt v0.0.0-20220222012616-ea85ff4c361f/go.mod h1:JqzWyvTuI2X4+9wOHmKSQCYxybB/8j6Ko43qVmXDuZg=
|
||||||
|
github.com/smarty/assertions v1.15.0 h1:cR//PqUBUiQRakZWqBiFFQ9wb8emQGDb0HeGdqGByCY=
|
||||||
|
github.com/smarty/assertions v1.15.0/go.mod h1:yABtdzeQs6l1brC900WlRNwj6ZR55d7B+E8C6HtKdec=
|
||||||
|
github.com/smartystreets/goconvey v1.8.1 h1:qGjIddxOk4grTu9JPOU31tVfq3cNdBlNa5sSznIX1xY=
|
||||||
|
github.com/smartystreets/goconvey v1.8.1/go.mod h1:+/u4qLyY6x1jReYOp7GOM2FSt8aP9CzCZL03bI28W60=
|
||||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 h1:+jumHNA0Wrelhe64i8F6HNlS8pkoyMv5sreGx2Ry5Rw=
|
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 h1:+jumHNA0Wrelhe64i8F6HNlS8pkoyMv5sreGx2Ry5Rw=
|
||||||
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8/go.mod h1:3n1Cwaq1E1/1lhQhtRK2ts/ZwZEhjcQeJQ1RuC6Q/8U=
|
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8/go.mod h1:3n1Cwaq1E1/1lhQhtRK2ts/ZwZEhjcQeJQ1RuC6Q/8U=
|
||||||
github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I=
|
github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I=
|
||||||
@@ -94,15 +169,18 @@ github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3A
|
|||||||
github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU=
|
github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU=
|
||||||
github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY=
|
github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY=
|
||||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||||
|
github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||||
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
|
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
|
||||||
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
|
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
|
||||||
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
|
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
|
||||||
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
|
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
|
||||||
|
github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs=
|
||||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||||
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||||
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
||||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
|
||||||
|
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||||
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
|
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
|
||||||
@@ -111,30 +189,50 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS
|
|||||||
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
|
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
|
||||||
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
|
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
|
||||||
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
|
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
|
||||||
|
github.com/wk8/go-ordered-map/v2 v2.1.8 h1:5h/BUHu93oj4gIdvHHHGsScSTMijfx5PeYkE/fJgbpc=
|
||||||
|
github.com/wk8/go-ordered-map/v2 v2.1.8/go.mod h1:5nJHM5DyteebpVlHnWMV0rPz6Zp7+xBAnxjb1X5vnTw=
|
||||||
|
github.com/x-cray/logrus-prefixed-formatter v0.5.2 h1:00txxvfBM9muc0jiLIEAkAcIMJzfthRT6usrui8uGmg=
|
||||||
|
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/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=
|
||||||
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
|
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
|
||||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
|
go.uber.org/mock v0.4.0 h1:VcM4ZOtdbR4f6VXfiOpwpVJDL6lCReaZ6mw31wqh7KU=
|
||||||
|
go.uber.org/mock v0.4.0/go.mod h1:a6FSlNadKUHUa9IP5Vyt1zh4fC7uAwxMutEAscFbkZc=
|
||||||
go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ=
|
go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ=
|
||||||
go.uber.org/multierr v1.10.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
go.uber.org/multierr v1.10.0/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
|
||||||
go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
go.uber.org/zap v1.28.0 h1:IZzaP1Fv73/T/pBMLk4VutPl36uNC+OSUh3JLG3FIjo=
|
||||||
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
go.uber.org/zap v1.28.0/go.mod h1:rDLpOi171uODNm/mxFcuYWxDsqWSAVkFdX4XojSKg/Q=
|
||||||
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
|
||||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||||
golang.org/x/arch v0.0.0-20210923205945-b76863e36670/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8=
|
golang.org/x/arch v0.11.0 h1:KXV8WWKCXm6tRpLirl2szsO5j/oOODwZf4hATmGVNs4=
|
||||||
golang.org/x/arch v0.8.0 h1:3wRIsP3pM4yUptoR96otTUOXI367OS0+c9eeRi9doIc=
|
golang.org/x/arch v0.11.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
|
||||||
golang.org/x/arch v0.8.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
|
golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4=
|
||||||
golang.org/x/crypto v0.23.0 h1:dIJU/v2J8Mdglj/8rJ6UUOM3Zc9zLZxVZwwxMooUSAI=
|
golang.org/x/crypto v0.31.0 h1:ihbySMvVjLAeSH1IbfcRTkD/iNscyz8rGzjF/E5hV6U=
|
||||||
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
|
golang.org/x/crypto v0.31.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
|
||||||
|
golang.org/x/exp v0.0.0-20230713183714-613f0c0eb8a1 h1:MGwJjxBy0HJshjDNfLsYO8xppfqWlA5ZT9OhtUUhTNw=
|
||||||
|
golang.org/x/exp v0.0.0-20230713183714-613f0c0eb8a1/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||||
|
golang.org/x/net v0.0.0-20180906233101-161cd47e91fd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||||
golang.org/x/net v0.25.0 h1:d/OCCoBEUq33pjydKrGQhw7IlUPI2Oylr+8qLx49kac=
|
golang.org/x/net v0.25.0 h1:d/OCCoBEUq33pjydKrGQhw7IlUPI2Oylr+8qLx49kac=
|
||||||
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
|
||||||
|
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
||||||
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||||
|
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
|
golang.org/x/sys v0.0.0-20180909124046-d0be0721c37e/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
|
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
|
golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
|
||||||
golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||||
|
golang.org/x/term v0.28.0 h1:/Ts8HFuMR2E6IP/jlo7QVLZHggjKQbhu/7H0LJFr3Gg=
|
||||||
|
golang.org/x/term v0.28.0/go.mod h1:Sw/lC2IAUZ92udQNf3WodGtn4k/XoLyZoh8v/8uiwek=
|
||||||
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
|
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
|
||||||
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
|
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
|
||||||
google.golang.org/protobuf v1.34.1 h1:9ddQBjfCyZPOHPUiPxpYESBLc+T8P3E+Vo4IbKZgFWg=
|
google.golang.org/protobuf v1.34.1 h1:9ddQBjfCyZPOHPUiPxpYESBLc+T8P3E+Vo4IbKZgFWg=
|
||||||
@@ -142,8 +240,9 @@ google.golang.org/protobuf v1.34.1/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHh
|
|||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
|
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
|
||||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
|
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
|
||||||
|
gopkg.in/fsnotify.v1 v1.4.7/go.mod h1:Tz8NjZHkW78fSQdbUxIjBTcgA1z1m8ZHf0WmKUhAMys=
|
||||||
|
gopkg.in/tomb.v1 v1.0.0-20141024135613-dd632973f1e7/go.mod h1:dt/ZhP58zS4L8KSrWDmTeBkI65Dw0HsyUHuEVlX15mw=
|
||||||
|
gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
nullprogram.com/x/optparse v1.0.0/go.mod h1:KdyPE+Igbe0jQUrVfMqDMeJQIJZEuyV7pjYmp6pbG50=
|
|
||||||
rsc.io/pdf v0.1.1/go.mod h1:n8OzWcQ6Sp37PL01nO98y4iUCRdTGarVfzxY20ICaU4=
|
|
||||||
|
|||||||
@@ -1,239 +0,0 @@
|
|||||||
package llm
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/base64"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"go.uber.org/zap"
|
|
||||||
)
|
|
||||||
|
|
||||||
// OpenAIService 基于 OpenAI Chat Completions API 的 LLM 实现。
|
|
||||||
type OpenAIService struct {
|
|
||||||
apiKey string
|
|
||||||
model string
|
|
||||||
endpoint string
|
|
||||||
timeout time.Duration
|
|
||||||
logger *zap.SugaredLogger
|
|
||||||
client *http.Client
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewOpenAIService 创建 OpenAI LLM 服务。
|
|
||||||
// model、endpoint 由 config 层保证非空。
|
|
||||||
func NewOpenAIService(apiKey, model, endpoint string, timeoutSec, httpClientTimeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
|
|
||||||
timeout := time.Duration(timeoutSec) * time.Second
|
|
||||||
if timeout <= 0 {
|
|
||||||
timeout = 10 * time.Second
|
|
||||||
}
|
|
||||||
httpClientTimeout := time.Duration(httpClientTimeoutSec) * time.Second
|
|
||||||
if httpClientTimeout <= 0 {
|
|
||||||
httpClientTimeout = 60 * time.Second
|
|
||||||
}
|
|
||||||
return &OpenAIService{
|
|
||||||
apiKey: apiKey,
|
|
||||||
model: model,
|
|
||||||
endpoint: endpoint,
|
|
||||||
timeout: timeout,
|
|
||||||
logger: logger,
|
|
||||||
client: &http.Client{Timeout: httpClientTimeout},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- OpenAI API 请求/响应结构 ---
|
|
||||||
|
|
||||||
type chatRequest struct {
|
|
||||||
Model string `json:"model"`
|
|
||||||
Messages []chatMessage `json:"messages"`
|
|
||||||
Stream bool `json:"stream"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type chatMessage struct {
|
|
||||||
Role string `json:"role"`
|
|
||||||
Content []contentPart `json:"content"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type contentPart struct {
|
|
||||||
Type string `json:"type"`
|
|
||||||
Text string `json:"text"`
|
|
||||||
ImageURL *imageURL `json:"image_url,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type imageURL struct {
|
|
||||||
URL string `json:"url"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// streamDelta SSE 流式响应的单个 delta。
|
|
||||||
type streamDelta struct {
|
|
||||||
Choices []struct {
|
|
||||||
Delta struct {
|
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"delta"`
|
|
||||||
FinishReason *string `json:"finish_reason"`
|
|
||||||
} `json:"choices"`
|
|
||||||
Usage *struct {
|
|
||||||
PromptTokens int `json:"prompt_tokens"`
|
|
||||||
CompletionTokens int `json:"completion_tokens"`
|
|
||||||
TotalTokens int `json:"total_tokens"`
|
|
||||||
} `json:"usage"`
|
|
||||||
Model string `json:"model"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// ChatStream 实现 llm.Service。调用 OpenAI Chat Completions API 流式推理。
|
|
||||||
func (o *OpenAIService) ChatStream(ctx context.Context, req Request) (<-chan Chunk, error) {
|
|
||||||
// 构建请求
|
|
||||||
messages := o.buildMessages(req)
|
|
||||||
|
|
||||||
body := chatRequest{
|
|
||||||
Model: o.model,
|
|
||||||
Messages: messages,
|
|
||||||
Stream: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
payload, err := json.Marshal(body)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("llm: marshal request: %w", err)
|
|
||||||
}
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("llm: marshal request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 创建带超时的 context
|
|
||||||
ctx, cancel := context.WithTimeout(ctx, o.timeout)
|
|
||||||
|
|
||||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/chat/completions", bytes.NewReader(payload))
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
return nil, fmt.Errorf("llm: create request: %w", err)
|
|
||||||
}
|
|
||||||
httpReq.Header.Set("Content-Type", "application/json")
|
|
||||||
httpReq.Header.Set("Authorization", "Bearer "+o.apiKey)
|
|
||||||
|
|
||||||
resp, err := o.client.Do(httpReq)
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
return nil, fmt.Errorf("llm: send request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
|
||||||
cancel()
|
|
||||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
|
||||||
resp.Body.Close()
|
|
||||||
return nil, fmt.Errorf("llm: api error (status %d): %s", resp.StatusCode, string(bodyBytes))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 启动 goroutine 解析 SSE 流
|
|
||||||
ch := make(chan Chunk, 64)
|
|
||||||
go func() {
|
|
||||||
defer close(ch)
|
|
||||||
defer cancel()
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
o.parseSSEStream(resp.Body, ch)
|
|
||||||
}()
|
|
||||||
|
|
||||||
return ch, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseSSEStream 解析 SSE 流,将 delta 发送到 channel。
|
|
||||||
func (o *OpenAIService) parseSSEStream(body io.Reader, ch chan<- Chunk) {
|
|
||||||
scanner := bufio.NewScanner(body)
|
|
||||||
scanner.Buffer(make([]byte, 0, 64*1024), 256*1024)
|
|
||||||
|
|
||||||
var fullText strings.Builder
|
|
||||||
var lastModel string
|
|
||||||
|
|
||||||
for scanner.Scan() {
|
|
||||||
line := scanner.Text()
|
|
||||||
|
|
||||||
// SSE 格式:data: {...}
|
|
||||||
if !strings.HasPrefix(line, "data: ") {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
data := strings.TrimPrefix(line, "data: ")
|
|
||||||
if data == "[DONE]" {
|
|
||||||
// 流结束,发送最终 chunk
|
|
||||||
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var delta streamDelta
|
|
||||||
if err := json.Unmarshal([]byte(data), &delta); err != nil {
|
|
||||||
o.logger.Warnw("llm: unmarshal delta failed", "error", err, "data", data)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if delta.Model != "" {
|
|
||||||
lastModel = delta.Model
|
|
||||||
}
|
|
||||||
|
|
||||||
// 提取增量文本
|
|
||||||
if len(delta.Choices) > 0 {
|
|
||||||
content := delta.Choices[0].Delta.Content
|
|
||||||
if content != "" {
|
|
||||||
fullText.WriteString(content)
|
|
||||||
ch <- Chunk{Delta: content, Done: false, Model: lastModel}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 某些模型在最后一个 choice 中携带 usage
|
|
||||||
if delta.Choices[0].FinishReason != nil && delta.Usage != nil {
|
|
||||||
ch <- Chunk{
|
|
||||||
Delta: "",
|
|
||||||
Done: true,
|
|
||||||
Model: lastModel,
|
|
||||||
TokensUsed: &TokenUsage{
|
|
||||||
Prompt: delta.Usage.PromptTokens,
|
|
||||||
Completion: delta.Usage.CompletionTokens,
|
|
||||||
Total: delta.Usage.TotalTokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// scanner 结束但没收到 [DONE]
|
|
||||||
if err := scanner.Err(); err != nil {
|
|
||||||
o.logger.Warnw("llm: scan error", "error", err)
|
|
||||||
}
|
|
||||||
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildMessages 构建 OpenAI Chat API 的 messages 数组。
|
|
||||||
func (o *OpenAIService) buildMessages(req Request) []chatMessage {
|
|
||||||
var messages []chatMessage
|
|
||||||
|
|
||||||
// System prompt(情景覆盖优先)
|
|
||||||
messages = append(messages, chatMessage{
|
|
||||||
Role: "system",
|
|
||||||
Content: []contentPart{{Type: "text", Text: BuildSystemPrompt(req.Language, "", req.SystemPrompt)}},
|
|
||||||
})
|
|
||||||
|
|
||||||
// 历史消息
|
|
||||||
for _, msg := range req.History {
|
|
||||||
messages = append(messages, chatMessage{
|
|
||||||
Role: msg.Role,
|
|
||||||
Content: []contentPart{{Type: "text", Text: msg.Content}},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// 当前用户消息(图像 + 文本)
|
|
||||||
var parts []contentPart
|
|
||||||
if len(req.Image) > 0 {
|
|
||||||
b64 := base64.StdEncoding.EncodeToString(req.Image)
|
|
||||||
parts = append(parts, contentPart{
|
|
||||||
Type: "image_url",
|
|
||||||
ImageURL: &imageURL{URL: "data:image/jpeg;base64," + b64},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
parts = append(parts, contentPart{Type: "text", Text: req.Text})
|
|
||||||
messages = append(messages, chatMessage{Role: "user", Content: parts})
|
|
||||||
|
|
||||||
return messages
|
|
||||||
}
|
|
||||||
@@ -1,251 +0,0 @@
|
|||||||
package llm
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"go.uber.org/zap"
|
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
|
||||||
)
|
|
||||||
|
|
||||||
// mockLLMServer 创建模拟 OpenAI SSE 流式响应的 HTTP 服务器。
|
|
||||||
func mockLLMServer(t *testing.T, handler http.HandlerFunc) *httptest.Server {
|
|
||||||
t.Helper()
|
|
||||||
return httptest.NewServer(handler)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_Success(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
// 验证请求
|
|
||||||
if r.Method != http.MethodPost {
|
|
||||||
t.Errorf("method = %s, want POST", r.Method)
|
|
||||||
}
|
|
||||||
if !strings.Contains(r.URL.Path, "/chat/completions") {
|
|
||||||
t.Errorf("path = %s, should contain /chat/completions", r.URL.Path)
|
|
||||||
}
|
|
||||||
auth := r.Header.Get("Authorization")
|
|
||||||
if auth != "Bearer test-key" {
|
|
||||||
t.Errorf("Authorization = %q, want %q", auth, "Bearer test-key")
|
|
||||||
}
|
|
||||||
|
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
|
||||||
flusher, ok := w.(http.Flusher)
|
|
||||||
if !ok {
|
|
||||||
t.Fatal("ResponseWriter does not support Flusher")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 发送几个 delta
|
|
||||||
deltas := []string{"你好", "世界", "!"}
|
|
||||||
for _, d := range deltas {
|
|
||||||
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"%s\"}}],\"model\":\"gpt-4o\"}\n\n", d)
|
|
||||||
flusher.Flush()
|
|
||||||
}
|
|
||||||
|
|
||||||
// 发送 [DONE]
|
|
||||||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
||||||
flusher.Flush()
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, 60, zap.NewNop().Sugar())
|
|
||||||
|
|
||||||
ch, err := svc.ChatStream(context.Background(), Request{
|
|
||||||
Text: "这是什么?",
|
|
||||||
Language: "zh-CN",
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ChatStream() error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var chunks []Chunk
|
|
||||||
for c := range ch {
|
|
||||||
chunks = append(chunks, c)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 应该有 3 个文本 chunk + 1 个 Done chunk
|
|
||||||
if len(chunks) != 4 {
|
|
||||||
t.Fatalf("got %d chunks, want 4", len(chunks))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 验证文本内容
|
|
||||||
if chunks[0].Delta != "你好" {
|
|
||||||
t.Errorf("chunk[0].Delta = %q, want %q", chunks[0].Delta, "你好")
|
|
||||||
}
|
|
||||||
if chunks[1].Delta != "世界" {
|
|
||||||
t.Errorf("chunk[1].Delta = %q, want %q", chunks[1].Delta, "世界")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 验证最后一个 chunk 是 Done
|
|
||||||
last := chunks[len(chunks)-1]
|
|
||||||
if !last.Done {
|
|
||||||
t.Error("last chunk should be Done")
|
|
||||||
}
|
|
||||||
if last.Model != "gpt-4o" {
|
|
||||||
t.Errorf("last chunk Model = %q, want %q", last.Model, "gpt-4o")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_WithImage(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
|
||||||
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}],\"model\":\"gpt-4o\"}\n\n")
|
|
||||||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, 60, zap.NewNop().Sugar())
|
|
||||||
|
|
||||||
ch, err := svc.ChatStream(context.Background(), Request{
|
|
||||||
Image: []byte("fake-jpeg-data"),
|
|
||||||
Text: "描述图片",
|
|
||||||
Language: "zh-CN",
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ChatStream() error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 消费 channel
|
|
||||||
for range ch {
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_WithHistory(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
|
||||||
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}],\"model\":\"gpt-4o\"}\n\n")
|
|
||||||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, 60, zap.NewNop().Sugar())
|
|
||||||
|
|
||||||
ch, err := svc.ChatStream(context.Background(), Request{
|
|
||||||
Text: "继续",
|
|
||||||
Language: "zh-CN",
|
|
||||||
History: []models.Message{
|
|
||||||
{Role: "user", Content: "你好"},
|
|
||||||
{Role: "assistant", Content: "你好!有什么可以帮助你的吗?"},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ChatStream() error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for range ch {
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_APIError(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
w.WriteHeader(http.StatusUnauthorized)
|
|
||||||
fmt.Fprintf(w, `{"error":{"message":"Invalid API key"}}`)
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("bad-key", "gpt-4o", srv.URL, 10, 60, zap.NewNop().Sugar())
|
|
||||||
|
|
||||||
_, err := svc.ChatStream(context.Background(), Request{
|
|
||||||
Text: "test",
|
|
||||||
})
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("ChatStream() should return error for 401")
|
|
||||||
}
|
|
||||||
if !strings.Contains(err.Error(), "401") {
|
|
||||||
t.Errorf("error should mention 401, got: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_Timeout(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
// 模拟慢响应
|
|
||||||
time.Sleep(5 * time.Second)
|
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
|
||||||
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"late\"}}]}\n\n")
|
|
||||||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 1, 60, zap.NewNop().Sugar()) // 1s timeout
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
ch, err := svc.ChatStream(ctx, Request{Text: "test"})
|
|
||||||
if err != nil {
|
|
||||||
// 超时可能在建立连接时或读取时发生
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// 如果连接成功,消费 channel 应该超时
|
|
||||||
var gotContent bool
|
|
||||||
for c := range ch {
|
|
||||||
if c.Delta != "" {
|
|
||||||
gotContent = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if gotContent {
|
|
||||||
t.Error("should not receive content before timeout")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenAIService_ChatStream_UsageInResponse(t *testing.T) {
|
|
||||||
srv := mockLLMServer(t, func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "text/event-stream")
|
|
||||||
// 带 usage 的最后一个 chunk
|
|
||||||
fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":\"stop\"}],\"model\":\"gpt-4o\",\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":5,\"total_tokens\":15}}\n\n")
|
|
||||||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
||||||
})
|
|
||||||
defer srv.Close()
|
|
||||||
|
|
||||||
svc := NewOpenAIService("test-key", "gpt-4o", srv.URL, 10, 60, zap.NewNop().Sugar())
|
|
||||||
|
|
||||||
ch, err := svc.ChatStream(context.Background(), Request{Text: "test"})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ChatStream() error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var last Chunk
|
|
||||||
for c := range ch {
|
|
||||||
last = c
|
|
||||||
}
|
|
||||||
|
|
||||||
if !last.Done {
|
|
||||||
t.Error("last chunk should be Done")
|
|
||||||
}
|
|
||||||
if last.TokensUsed == nil {
|
|
||||||
t.Fatal("last chunk should have TokensUsed")
|
|
||||||
}
|
|
||||||
if last.TokensUsed.Total != 15 {
|
|
||||||
t.Errorf("TokensUsed.Total = %d, want 15", last.TokensUsed.Total)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildSystemPrompt(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
language string
|
|
||||||
detailLevel string
|
|
||||||
wantContain string
|
|
||||||
}{
|
|
||||||
{"chinese default", "zh-CN", "", "视觉助手"},
|
|
||||||
{"chinese high", "zh-CN", "high", "更详细"},
|
|
||||||
{"english default", "en", "", "visual assistant"},
|
|
||||||
{"english high", "en", "high", "detailed"},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
got := BuildSystemPrompt(tt.language, tt.detailLevel, "")
|
|
||||||
if !strings.Contains(got, tt.wantContain) {
|
|
||||||
t.Errorf("BuildSystemPrompt(%q, %q, \"\") should contain %q", tt.language, tt.detailLevel, tt.wantContain)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -2,54 +2,147 @@ package llm
|
|||||||
|
|
||||||
import "strings"
|
import "strings"
|
||||||
|
|
||||||
// scenarioPrompt 定义单个情景的多语言 system prompt。
|
// scenarioPrompt 定义单个情景的多语言 system prompt 和首句引导。
|
||||||
type scenarioPrompt struct {
|
type scenarioPrompt struct {
|
||||||
ZH string
|
ZH string
|
||||||
EN string
|
EN string
|
||||||
JA string
|
JA string
|
||||||
|
GreetingZH string // 首句引导(中文)
|
||||||
|
GreetingEN string // 首句引导(英文)
|
||||||
|
GreetingJA string // 首句引导(日文)
|
||||||
}
|
}
|
||||||
|
|
||||||
// scenarioPrompts 预置情景 → prompt 映射表。
|
// scenarioPrompts 预置情景 → prompt 映射表。
|
||||||
// key 为情景 ID(与前端 Scenario.id 对齐)。
|
// key 为情景 ID(与前端 Scenario.id 对齐)。
|
||||||
var scenarioPrompts = map[string]scenarioPrompt{
|
var scenarioPrompts = map[string]scenarioPrompt{
|
||||||
"interviewer": {
|
"interviewer": {
|
||||||
ZH: "你是一位资深面试官。你通过摄像头观察面试者,并根据他们的背景和表现提出面试问题。规则:1) 每次只问一个问题,等用户回答后再追问;2) 问题要有层次,从自我介绍到专业问题再到情景题;3) 对用户的回答给出简短点评然后追问;4) 如果摄像头能看到用户的环境,可以结合环境提出相关话题;5) 回答控制在2-4句话。",
|
ZH: `你是一位资深面试官。你通过摄像头观察面试者,并根据他们的背景和表现提出面试问题。
|
||||||
|
|
||||||
|
【角色定位】
|
||||||
|
- 你是面试官,不是助手或顾问
|
||||||
|
- 你的目标是评估候选人的能力
|
||||||
|
- 保持专业、客观、礼貌
|
||||||
|
|
||||||
|
【交互规则】
|
||||||
|
1. 每次只问一个问题,等用户回答后再追问
|
||||||
|
2. 问题要有层次:自我介绍 → 专业问题 → 情景题 → 反问环节
|
||||||
|
3. 对用户的回答给出简短点评(优点+不足),然后追问
|
||||||
|
4. 如果摄像头能看到用户的环境,可以结合环境提出相关话题
|
||||||
|
5. 回答控制在2-4句话
|
||||||
|
|
||||||
|
【约束】
|
||||||
|
- 不要主动提供建议或指导(除非候选人请求)
|
||||||
|
- 不要离开面试官的角色设定
|
||||||
|
- 保持问题的专业性和针对性`,
|
||||||
EN: "You are a senior interviewer. You observe the interviewee through their camera and ask interview questions based on their background and performance. Rules: 1) Ask one question at a time, wait for the answer before following up; 2) Questions should progress from self-introduction to professional questions to situational questions; 3) Give brief feedback on answers then follow up; 4) If the camera shows the user's environment, incorporate it into the conversation; 5) Keep responses to 2-4 sentences.",
|
EN: "You are a senior interviewer. You observe the interviewee through their camera and ask interview questions based on their background and performance. Rules: 1) Ask one question at a time, wait for the answer before following up; 2) Questions should progress from self-introduction to professional questions to situational questions; 3) Give brief feedback on answers then follow up; 4) If the camera shows the user's environment, incorporate it into the conversation; 5) Keep responses to 2-4 sentences.",
|
||||||
JA: "あなたはベテラン面接官です。カメラで面接者を見て、バックグラウンドと実績に基づいて面接質問をします。ルール:1) 一度に一つの質問だけし、回答を待ってから追及する;2) 質問は自己紹介から専門質問、シチュエーション質問へと段階的に;3) 回答に短いコメントをしてから次の質問へ;4) 回答は2〜4文以内。",
|
JA: "あなたはベテラン面接官です。カメラで面接者を見て、バックグラウンドと実績に基づいて面接質問をします。ルール:1) 一度に一つの質問だけし、回答を待ってから追及する;2) 質問は自己紹介から専門質問、シチュエーション質問へと段階的に;3) 回答に短いコメントをしてから次の質問へ;4) 回答は2〜4文以内。",
|
||||||
|
GreetingZH: "你好!我是今天的面试官。让我们先从自我介绍开始,请简单介绍一下你自己和你应聘的岗位。",
|
||||||
|
GreetingEN: "Hello! I'm your interviewer today. Let's start with a self-introduction. Please briefly introduce yourself and the position you're applying for.",
|
||||||
|
GreetingJA: "こんにちは!本日の面接官です。まず自己紹介から始めましょう。あなた自身と応募職種について簡単に教えてください。",
|
||||||
},
|
},
|
||||||
"english_teacher": {
|
"english_teacher": {
|
||||||
ZH: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material (e.g., 'I can see a bookshelf behind you. What's your favorite book?'); 6) Keep responses to 3-5 sentences.",
|
ZH: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material (e.g., 'I can see a bookshelf behind you. What's your favorite book?'); 6) Keep responses to 3-5 sentences.",
|
||||||
EN: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
EN: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
||||||
JA: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
JA: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
||||||
|
GreetingZH: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||||
|
GreetingEN: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||||
|
GreetingJA: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||||
},
|
},
|
||||||
"debate": {
|
"debate": {
|
||||||
ZH: "你是一位辩论赛对手。用户提出一个观点,你需要站在反方进行反驳。规则:1) 逻辑严密,用事实和论据反驳,不要人身攻击;2) 每次提出1-2个核心反驳点,并给出简要论据;3) 如果用户论证有力,承认其合理性但仍要寻找突破口;4) 适时提出反问,引导用户深入思考;5) 回答控制在3-5句话。",
|
ZH: `你是一位辩论赛对手。用户提出一个观点,你需要站在反方进行反驳。
|
||||||
|
|
||||||
|
【角色定位】
|
||||||
|
- 你是辩论对手,不是评委或顾问
|
||||||
|
- 你的目标是通过逻辑论证反驳对方观点
|
||||||
|
- 保持理性、严谨、尊重对手
|
||||||
|
|
||||||
|
【交互规则】
|
||||||
|
1. 逻辑严密,用事实和论据反驳,不要人身攻击
|
||||||
|
2. 每次提出1-2个核心反驳点,并给出简要论据
|
||||||
|
3. 如果用户论证有力,承认其合理性但仍要寻找突破口
|
||||||
|
4. 适时提出反问,引导用户深入思考
|
||||||
|
5. 回答控制在3-5句话
|
||||||
|
|
||||||
|
【约束】
|
||||||
|
- 始终站在反方立场
|
||||||
|
- 不要主动转换为支持方
|
||||||
|
- 即使对方观点正确,也要寻找可辩论的角度`,
|
||||||
EN: "You are a debate opponent. The user presents a viewpoint, and you argue against it. Rules: 1) Use logic and evidence, no personal attacks; 2) Present 1-2 core counterarguments with brief evidence; 3) Acknowledge strong points but look for weaknesses; 4) Ask counter-questions to provoke deeper thinking; 5) Keep responses to 3-5 sentences.",
|
EN: "You are a debate opponent. The user presents a viewpoint, and you argue against it. Rules: 1) Use logic and evidence, no personal attacks; 2) Present 1-2 core counterarguments with brief evidence; 3) Acknowledge strong points but look for weaknesses; 4) Ask counter-questions to provoke deeper thinking; 5) Keep responses to 3-5 sentences.",
|
||||||
JA: "あなたはディベートの相手です。ユーザーが提示した观点に対して反論します。ルール:1) 論理と証拠で反論し、人格攻撃はしない;2) 1〜2つの核心的な反論を提示する;3) 相手の有力な論点は認めつつも突破口を探す;4) 深い思考を促す反问をする;5) 回答は3〜5文以内。",
|
JA: "あなたはディベートの相手です。ユーザーが提示した观点に対して反論します。ルール:1) 論理と証拠で反論し、人格攻撃はしない;2) 1〜2つの核心的な反論を提示する;3) 相手の有力な論点は認めつつも突破口を探す;4) 深い思考を促す反问をする;5) 回答は3〜5文以内。",
|
||||||
|
GreetingZH: "你好!我是你的辩论对手。请提出一个你坚信的观点,我会站在反方立场与你辩论,帮你锻炼逻辑思维。",
|
||||||
|
GreetingEN: "Hello! I'm your debate opponent. Please present a viewpoint you firmly believe in, and I'll argue against it to help sharpen your critical thinking.",
|
||||||
|
GreetingJA: "こんにちは!あなたのディベート相手です。あなたが信じる观点を提示してください。反対の立場から論じて、論理的思考を鍛えます。",
|
||||||
},
|
},
|
||||||
"interpreter": {
|
"interpreter": {
|
||||||
ZH: "你是一名同声翻译员。将用户说的话实时翻译为目标语言。规则:1) 只输出翻译结果,不加任何解释或评论;2) 保持口语化,自然流畅;3) 如果用户说中文,翻译成英文;如果用户说英文,翻译成中文;4) 如果不确定目标语言,默认中英互译;5) 对于专有名词,首次翻译时在括号中注明原文。",
|
ZH: "你是一名同声翻译员。将用户说的话实时翻译为目标语言。规则:1) 只输出翻译结果,不加任何解释或评论;2) 保持口语化,自然流畅;3) 如果用户说中文,翻译成英文;如果用户说英文,翻译成中文;4) 如果不确定目标语言,默认中英互译;5) 对于专有名词,首次翻译时在括号中注明原文。",
|
||||||
EN: "You are a simultaneous interpreter. Translate what the user says in real-time. Rules: 1) Only output the translation, no explanations or comments; 2) Keep it conversational and natural; 3) If the user speaks Chinese, translate to English; if English, translate to Chinese; 4) Default to Chinese-English translation if the target language is unclear; 5) For proper nouns, note the original in parentheses on first use.",
|
EN: "You are a simultaneous interpreter. Translate what the user says in real-time. Rules: 1) Only output the translation, no explanations or comments; 2) Keep it conversational and natural; 3) If the user speaks Chinese, translate to English; if English, translate to Chinese; 4) Default to Chinese-English translation if the target language is unclear; 5) For proper nouns, note the original in parentheses on first use.",
|
||||||
JA: "あなたは同時通訳者です。ユーザーの発言をリアルタイムで翻訳します。ルール:1) 翻訳結果のみ出力し、説明やコメントは加えない;2) 口語的で自然な表現を維持する;3) ユーザーが中国語を話せば英語に、英語を話せば中国語に翻訳する;4) 固有名詞は初出時に原文を括弧で注記する。",
|
JA: "あなたは同時通訳者です。ユーザーの発言をリアルタイムで翻訳します。ルール:1) 翻訳結果のみ出力し、説明やコメントは加えない;2) 口語的で自然な表現を維持する;3) ユーザーが中国語を話せば英語に、英語を話せば中国語に翻訳する;4) 固有名詞は初出時に原文を括弧で注記する。",
|
||||||
|
GreetingZH: "我是你的同声翻译。请开始说话,我会实时将中文翻译成英文,或将英文翻译成中文。",
|
||||||
|
GreetingEN: "I'm your simultaneous interpreter. Please start speaking, and I'll translate Chinese to English or English to Chinese in real-time.",
|
||||||
|
GreetingJA: "私はあなたの同時通訳者です。お話しください。中国語を英語に、または英語を中国語にリアルタイムで翻訳します。",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetScenarioPrompt 根据情景 ID 和语言获取对应的 system prompt。
|
// GetScenarioPrompt 根据情景 ID 和语言获取对应的 system prompt。
|
||||||
|
// 支持系统预置情景和用户自建情景。
|
||||||
|
// customScenarios: 用户自建情景映射表(scenarioID → prompt),可为 nil
|
||||||
// 返回空字符串表示无此情景(使用默认 prompt)。
|
// 返回空字符串表示无此情景(使用默认 prompt)。
|
||||||
func GetScenarioPrompt(scenarioID, language string) string {
|
func GetScenarioPrompt(scenarioID, language string, customScenarios map[string]string) string {
|
||||||
if scenarioID == "" || scenarioID == "free_chat" {
|
if scenarioID == "" || scenarioID == "free_chat" {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
p, ok := scenarioPrompts[scenarioID]
|
|
||||||
if !ok {
|
// 1. 优先查找系统预置情景
|
||||||
|
if p, ok := scenarioPrompts[scenarioID]; ok {
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(language, "zh"):
|
||||||
|
return p.ZH
|
||||||
|
case strings.HasPrefix(language, "ja"):
|
||||||
|
return p.JA
|
||||||
|
default:
|
||||||
|
return p.EN
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. 查找用户自建情景
|
||||||
|
if customScenarios != nil {
|
||||||
|
if customPrompt, ok := customScenarios[scenarioID]; ok {
|
||||||
|
return customPrompt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. 默认空字符串
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetScenarioGreeting 根据情景 ID 和语言获取对应的首句引导。
|
||||||
|
// 支持系统预置情景和用户自建情景。
|
||||||
|
// customGreetings: 用户自建情景的首句引导映射表(scenarioID → greeting),可为 nil
|
||||||
|
// 返回空字符串表示无此情景或不需要引导(自由对话)。
|
||||||
|
func GetScenarioGreeting(scenarioID, language string, customGreetings map[string]string) string {
|
||||||
|
if scenarioID == "" || scenarioID == "free_chat" {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
switch {
|
|
||||||
case strings.HasPrefix(language, "zh"):
|
// 1. 优先查找系统预置情景
|
||||||
return p.ZH
|
if p, ok := scenarioPrompts[scenarioID]; ok {
|
||||||
case strings.HasPrefix(language, "ja"):
|
switch {
|
||||||
return p.JA
|
case strings.HasPrefix(language, "zh"):
|
||||||
default:
|
return p.GreetingZH
|
||||||
return p.EN
|
case strings.HasPrefix(language, "ja"):
|
||||||
|
return p.GreetingJA
|
||||||
|
default:
|
||||||
|
return p.GreetingEN
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 2. 查找用户自建情景
|
||||||
|
if customGreetings != nil {
|
||||||
|
if customGreeting, ok := customGreetings[scenarioID]; ok {
|
||||||
|
return customGreeting
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. 默认空字符串
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,6 +11,8 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
"github.com/hhs/camtalk/internal/util"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -108,7 +110,11 @@ func (m *MiMoService) SynthesizeStream(ctx context.Context, textStream <-chan st
|
|||||||
|
|
||||||
audio, err := m.synthesize(ctx, text, voice)
|
audio, err := m.synthesize(ctx, text, voice)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
m.logger.Warnw("mimo tts: synthesize failed", "error", err, "text", text)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("mimo tts: synthesize failed",
|
||||||
|
"error", err,
|
||||||
|
"text_len", len(text),
|
||||||
|
"text_preview", util.Truncate(text, 100))
|
||||||
// 静默跳过,不中断整个流
|
// 静默跳过,不中断整个流
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,8 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
"github.com/hhs/camtalk/internal/util"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -81,7 +83,11 @@ func (o *OpenAIService) SynthesizeStream(ctx context.Context, textStream <-chan
|
|||||||
|
|
||||||
audio, err := o.synthesize(ctx, text, voice, speed)
|
audio, err := o.synthesize(ctx, text, voice, speed)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
o.logger.Warnw("tts: synthesize failed", "error", err, "text", text)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("tts: synthesize failed",
|
||||||
|
"error", err,
|
||||||
|
"text_len", len(text),
|
||||||
|
"text_preview", util.Truncate(text, 100))
|
||||||
// 静默跳过,不中断整个流
|
// 静默跳过,不中断整个流
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ 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"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// AuthHandler 提供认证相关的 REST 端点。
|
// AuthHandler 提供认证相关的 REST 端点。
|
||||||
@@ -25,11 +27,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)
|
||||||
}
|
}
|
||||||
@@ -37,6 +54,9 @@ func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
|||||||
|
|
||||||
// Register POST /api/auth/register — 用户注册。
|
// Register POST /api/auth/register — 用户注册。
|
||||||
func (h *AuthHandler) Register(c *gin.Context) {
|
func (h *AuthHandler) Register(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
clientIP := c.ClientIP()
|
||||||
|
|
||||||
var req auth.RegisterRequest
|
var req auth.RegisterRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{
|
c.JSON(http.StatusBadRequest, gin.H{
|
||||||
@@ -56,15 +76,25 @@ func (h *AuthHandler) Register(c *gin.Context) {
|
|||||||
|
|
||||||
resp, err := h.authService.Register(c.Request.Context(), req)
|
resp, err := h.authService.Register(c.Request.Context(), req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Warnw("register failed",
|
||||||
|
"username", req.Username,
|
||||||
|
"client_ip", clientIP,
|
||||||
|
"error", err)
|
||||||
handleAuthError(c, err)
|
handleAuthError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("register success",
|
||||||
|
"username", req.Username,
|
||||||
|
"client_ip", clientIP)
|
||||||
c.JSON(http.StatusCreated, resp)
|
c.JSON(http.StatusCreated, resp)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Login POST /api/auth/login — 用户登录。
|
// Login POST /api/auth/login — 用户登录。
|
||||||
func (h *AuthHandler) Login(c *gin.Context) {
|
func (h *AuthHandler) Login(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
clientIP := c.ClientIP()
|
||||||
|
|
||||||
var req auth.LoginRequest
|
var req auth.LoginRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{
|
c.JSON(http.StatusBadRequest, gin.H{
|
||||||
@@ -84,15 +114,24 @@ func (h *AuthHandler) Login(c *gin.Context) {
|
|||||||
|
|
||||||
resp, err := h.authService.Login(c.Request.Context(), req)
|
resp, err := h.authService.Login(c.Request.Context(), req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Warnw("login failed",
|
||||||
|
"username", req.Username,
|
||||||
|
"client_ip", clientIP,
|
||||||
|
"error", err)
|
||||||
handleAuthError(c, err)
|
handleAuthError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("login success",
|
||||||
|
"username", req.Username,
|
||||||
|
"client_ip", clientIP)
|
||||||
c.JSON(http.StatusOK, resp)
|
c.JSON(http.StatusOK, resp)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Refresh POST /api/auth/refresh — 刷新令牌。
|
// Refresh POST /api/auth/refresh — 刷新令牌。
|
||||||
func (h *AuthHandler) Refresh(c *gin.Context) {
|
func (h *AuthHandler) Refresh(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
|
||||||
var req auth.RefreshRequest
|
var req auth.RefreshRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{
|
c.JSON(http.StatusBadRequest, gin.H{
|
||||||
@@ -112,15 +151,21 @@ func (h *AuthHandler) Refresh(c *gin.Context) {
|
|||||||
|
|
||||||
resp, err := h.authService.Refresh(c.Request.Context(), req)
|
resp, err := h.authService.Refresh(c.Request.Context(), req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Warnw("token refresh failed",
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
"error", err)
|
||||||
handleAuthError(c, err)
|
handleAuthError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("token refresh success",
|
||||||
|
"client_ip", c.ClientIP())
|
||||||
c.JSON(http.StatusOK, resp)
|
c.JSON(http.StatusOK, resp)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Logout POST /api/auth/logout — 登出(需要认证)。
|
// Logout POST /api/auth/logout — 登出(需要认证)。
|
||||||
func (h *AuthHandler) Logout(c *gin.Context) {
|
func (h *AuthHandler) Logout(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
userID := c.GetString(auth.ContextKeyUserID)
|
userID := c.GetString(auth.ContextKeyUserID)
|
||||||
|
|
||||||
var req struct {
|
var req struct {
|
||||||
@@ -143,6 +188,9 @@ func (h *AuthHandler) Logout(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err := h.authService.Logout(c.Request.Context(), userID, req.RefreshToken); err != nil {
|
if err := h.authService.Logout(c.Request.Context(), userID, req.RefreshToken); err != nil {
|
||||||
|
log.Errorw("logout failed",
|
||||||
|
"user_id", userID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to logout",
|
"message": "failed to logout",
|
||||||
@@ -150,6 +198,8 @@ func (h *AuthHandler) Logout(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("logout success",
|
||||||
|
"user_id", userID)
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"message": "logged out successfully",
|
"message": "logged out successfully",
|
||||||
})
|
})
|
||||||
@@ -158,8 +208,8 @@ func (h *AuthHandler) Logout(c *gin.Context) {
|
|||||||
// validateCredentials 校验用户名和密码格式。
|
// validateCredentials 校验用户名和密码格式。
|
||||||
// 返回空字符串表示校验通过,否则返回错误描述。
|
// 返回空字符串表示校验通过,否则返回错误描述。
|
||||||
func validateCredentials(username, password string) string {
|
func validateCredentials(username, password string) string {
|
||||||
if len(username) < 3 || len(username) > 64 {
|
if len(username) > 64 {
|
||||||
return "username must be 3-64 characters"
|
return "username must not exceed 64 characters"
|
||||||
}
|
}
|
||||||
if len(password) < 8 || len(password) > 72 {
|
if len(password) < 8 || len(password) > 72 {
|
||||||
return "password must be 8-72 characters"
|
return "password must be 8-72 characters"
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
"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/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConversationHandler 提供对话相关的 REST 端点。
|
// ConversationHandler 提供对话相关的 REST 端点。
|
||||||
@@ -47,6 +48,7 @@ func (h *ConversationHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
|||||||
|
|
||||||
// List GET /api/conversations — 获取当前用户的对话列表。
|
// List GET /api/conversations — 获取当前用户的对话列表。
|
||||||
func (h *ConversationHandler) List(c *gin.Context) {
|
func (h *ConversationHandler) List(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
userID := c.GetString(auth.ContextKeyUserID)
|
userID := c.GetString(auth.ContextKeyUserID)
|
||||||
|
|
||||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||||
@@ -61,6 +63,9 @@ func (h *ConversationHandler) List(c *gin.Context) {
|
|||||||
|
|
||||||
summaries, total, err := h.sessionMgr.ListByUser(c.Request.Context(), userID, page, size)
|
summaries, total, err := h.sessionMgr.ListByUser(c.Request.Context(), userID, page, size)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("list conversations failed",
|
||||||
|
"user_id", userID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to list conversations",
|
"message": "failed to list conversations",
|
||||||
@@ -83,6 +88,7 @@ type CreateConversationRequest struct {
|
|||||||
|
|
||||||
// Create POST /api/conversations — 创建新对话。
|
// Create POST /api/conversations — 创建新对话。
|
||||||
func (h *ConversationHandler) Create(c *gin.Context) {
|
func (h *ConversationHandler) Create(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
userID := c.GetString(auth.ContextKeyUserID)
|
userID := c.GetString(auth.ContextKeyUserID)
|
||||||
|
|
||||||
var req CreateConversationRequest
|
var req CreateConversationRequest
|
||||||
@@ -95,6 +101,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
|||||||
|
|
||||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), userID, cfg)
|
sessionID, err := h.sessionMgr.Create(c.Request.Context(), userID, cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("create conversation failed",
|
||||||
|
"user_id", userID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to create conversation",
|
"message": "failed to create conversation",
|
||||||
@@ -104,6 +113,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
|||||||
|
|
||||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("retrieve created conversation failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to retrieve created conversation",
|
"message": "failed to retrieve created conversation",
|
||||||
@@ -111,6 +123,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("conversation created",
|
||||||
|
"conversation_id", sess.ID,
|
||||||
|
"user_id", userID)
|
||||||
c.JSON(http.StatusCreated, gin.H{
|
c.JSON(http.StatusCreated, gin.H{
|
||||||
"id": sess.ID,
|
"id": sess.ID,
|
||||||
"title": sess.Title,
|
"title": sess.Title,
|
||||||
@@ -144,6 +159,7 @@ type UpdateTitleRequest struct {
|
|||||||
|
|
||||||
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
|
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
|
||||||
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
sessionID := c.Param("id")
|
sessionID := c.Param("id")
|
||||||
|
|
||||||
// 先校验归属
|
// 先校验归属
|
||||||
@@ -176,6 +192,9 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
|||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
log.Errorw("update title failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to update title",
|
"message": "failed to update title",
|
||||||
@@ -190,6 +209,7 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
|||||||
|
|
||||||
// Delete DELETE /api/conversations/:id — 删除对话。
|
// Delete DELETE /api/conversations/:id — 删除对话。
|
||||||
func (h *ConversationHandler) Delete(c *gin.Context) {
|
func (h *ConversationHandler) Delete(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
sessionID := c.Param("id")
|
sessionID := c.Param("id")
|
||||||
|
|
||||||
// 先校验归属
|
// 先校验归属
|
||||||
@@ -205,6 +225,9 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
|
|||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
log.Errorw("delete conversation failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to delete conversation",
|
"message": "failed to delete conversation",
|
||||||
@@ -221,6 +244,7 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
|
|||||||
// - limit: 返回消息数量上限,默认 50
|
// - limit: 返回消息数量上限,默认 50
|
||||||
// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
|
// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
|
||||||
func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
sessionID := c.Param("id")
|
sessionID := c.Param("id")
|
||||||
|
|
||||||
// 先校验归属
|
// 先校验归属
|
||||||
@@ -239,6 +263,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
|||||||
if h.msgRepo != nil {
|
if h.msgRepo != nil {
|
||||||
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
|
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("get messages failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to get messages",
|
"message": "failed to get messages",
|
||||||
@@ -246,6 +273,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
count, _ := h.msgRepo.GetMessageCount(c.Request.Context(), sessionID)
|
count, _ := h.msgRepo.GetMessageCount(c.Request.Context(), sessionID)
|
||||||
|
if messages == nil {
|
||||||
|
messages = []store.StoredMessage{}
|
||||||
|
}
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"total": count,
|
"total": count,
|
||||||
@@ -263,6 +293,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
|||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
log.Errorw("get messages failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": apperr.CodeInternalError,
|
"code": apperr.CodeInternalError,
|
||||||
"message": "failed to get messages",
|
"message": "failed to get messages",
|
||||||
@@ -284,6 +317,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
messages := allMessages[start:]
|
messages := allMessages[start:]
|
||||||
|
|
||||||
|
if messages == nil {
|
||||||
|
messages = []models.Message{}
|
||||||
|
}
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"total": total,
|
"total": total,
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
"github.com/hhs/camtalk/internal/session"
|
"github.com/hhs/camtalk/internal/session"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// SessionHandler 提供会话相关的 REST 端点。
|
// SessionHandler 提供会话相关的 REST 端点。
|
||||||
@@ -27,6 +28,8 @@ type CreateSessionRequest struct {
|
|||||||
|
|
||||||
// CreateSession POST /api/sessions — 创建新会话。
|
// CreateSession POST /api/sessions — 创建新会话。
|
||||||
func (h *SessionHandler) CreateSession(c *gin.Context) {
|
func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
|
||||||
var req CreateSessionRequest
|
var req CreateSessionRequest
|
||||||
// 请求体可选,解析失败不报错(使用默认配置)
|
// 请求体可选,解析失败不报错(使用默认配置)
|
||||||
_ = c.ShouldBindJSON(&req)
|
_ = c.ShouldBindJSON(&req)
|
||||||
@@ -38,6 +41,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
|||||||
|
|
||||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), "", cfg)
|
sessionID, err := h.sessionMgr.Create(c.Request.Context(), "", cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("create session failed",
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": "INTERNAL_ERROR",
|
"code": "INTERNAL_ERROR",
|
||||||
"message": "failed to create session",
|
"message": "failed to create session",
|
||||||
@@ -48,6 +53,9 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
|||||||
// 获取创建后的会话以返回 created_at
|
// 获取创建后的会话以返回 created_at
|
||||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("retrieve created session failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": "INTERNAL_ERROR",
|
"code": "INTERNAL_ERROR",
|
||||||
"message": "failed to retrieve created session",
|
"message": "failed to retrieve created session",
|
||||||
@@ -55,6 +63,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("session created",
|
||||||
|
"session_id", sess.ID)
|
||||||
c.JSON(http.StatusCreated, gin.H{
|
c.JSON(http.StatusCreated, gin.H{
|
||||||
"session_id": sess.ID,
|
"session_id": sess.ID,
|
||||||
"created_at": sess.CreatedAt,
|
"created_at": sess.CreatedAt,
|
||||||
@@ -63,6 +73,7 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
|||||||
|
|
||||||
// DestroySession DELETE /api/sessions/:id — 销毁会话。
|
// DestroySession DELETE /api/sessions/:id — 销毁会话。
|
||||||
func (h *SessionHandler) DestroySession(c *gin.Context) {
|
func (h *SessionHandler) DestroySession(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
sessionID := c.Param("id")
|
sessionID := c.Param("id")
|
||||||
|
|
||||||
err := h.sessionMgr.Destroy(c.Request.Context(), sessionID)
|
err := h.sessionMgr.Destroy(c.Request.Context(), sessionID)
|
||||||
@@ -74,6 +85,9 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
|
|||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
log.Errorw("destroy session failed",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"error", err)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{
|
c.JSON(http.StatusInternalServerError, gin.H{
|
||||||
"code": "INTERNAL_ERROR",
|
"code": "INTERNAL_ERROR",
|
||||||
"message": "failed to destroy session",
|
"message": "failed to destroy session",
|
||||||
@@ -81,6 +95,8 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infow("session destroyed",
|
||||||
|
"session_id", sessionID)
|
||||||
c.Status(http.StatusNoContent)
|
c.Status(http.StatusNoContent)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
208
backend/internal/api/user_scenario_handler.go
Normal file
208
backend/internal/api/user_scenario_handler.go
Normal file
@@ -0,0 +1,208 @@
|
|||||||
|
package api
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
MaxScenariosPerUser = 20 // 每个用户最多 20 个自建情景
|
||||||
|
MaxPromptLength = 2000 // Prompt 最大长度
|
||||||
|
)
|
||||||
|
|
||||||
|
// UserScenarioHandler 用户情景 API Handler。
|
||||||
|
type UserScenarioHandler struct {
|
||||||
|
repo store.UserScenarioRepository
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewUserScenarioHandler 创建用户情景 Handler。
|
||||||
|
func NewUserScenarioHandler(repo store.UserScenarioRepository) *UserScenarioHandler {
|
||||||
|
return &UserScenarioHandler{repo: repo}
|
||||||
|
}
|
||||||
|
|
||||||
|
// List 获取用户的所有自建情景。
|
||||||
|
// GET /api/scenarios
|
||||||
|
func (h *UserScenarioHandler) List(c *gin.Context) {
|
||||||
|
userID, exists := c.Get("user_id")
|
||||||
|
if !exists {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scenarios, err := h.repo.FindByUserID(c.Request.Context(), userID.(string))
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "error", err)
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "查询失败"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if scenarios == nil {
|
||||||
|
scenarios = []*models.UserScenario{}
|
||||||
|
}
|
||||||
|
|
||||||
|
c.JSON(http.StatusOK, models.UserScenarioListResponse{
|
||||||
|
Scenarios: scenarios,
|
||||||
|
Total: len(scenarios),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create 创建用户情景。
|
||||||
|
// POST /api/scenarios
|
||||||
|
func (h *UserScenarioHandler) Create(c *gin.Context) {
|
||||||
|
userID, exists := c.Get("user_id")
|
||||||
|
if !exists {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req models.CreateUserScenarioRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 检查用户是否已达上限
|
||||||
|
count, err := h.repo.CountByUserID(c.Request.Context(), userID.(string))
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Errorw("统计用户情景数量失败", "user_id", userID, "error", err)
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建失败"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if count >= MaxScenariosPerUser {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "已达创建上限(最多 20 个)"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 创建情景
|
||||||
|
scenario := &models.UserScenario{
|
||||||
|
UserID: userID.(string),
|
||||||
|
Name: req.Name,
|
||||||
|
Icon: req.Icon,
|
||||||
|
Description: req.Description,
|
||||||
|
Prompt: req.Prompt,
|
||||||
|
Greeting: req.Greeting,
|
||||||
|
Language: req.Language,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.repo.Create(c.Request.Context(), scenario); err != nil {
|
||||||
|
logger.Log.Errorw("创建用户情景失败", "user_id", userID, "error", err)
|
||||||
|
if err.Error() == "duplicate key value violates unique constraint" {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "情景名称已存在"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建失败"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Log.Infow("创建用户情景成功", "user_id", userID, "scenario_id", scenario.ID)
|
||||||
|
c.JSON(http.StatusCreated, scenario)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get 获取单个情景详情。
|
||||||
|
// GET /api/scenarios/:id
|
||||||
|
func (h *UserScenarioHandler) Get(c *gin.Context) {
|
||||||
|
userID, exists := c.Get("user_id")
|
||||||
|
if !exists {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scenarioID := c.Param("id")
|
||||||
|
scenario, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||||
|
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
c.JSON(http.StatusOK, scenario)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update 更新用户情景。
|
||||||
|
// PATCH /api/scenarios/:id
|
||||||
|
func (h *UserScenarioHandler) Update(c *gin.Context) {
|
||||||
|
userID, exists := c.Get("user_id")
|
||||||
|
if !exists {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scenarioID := c.Param("id")
|
||||||
|
|
||||||
|
// 查询并校验所有权
|
||||||
|
scenario, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||||
|
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req models.UpdateUserScenarioRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 更新字段
|
||||||
|
if req.Name != nil {
|
||||||
|
scenario.Name = *req.Name
|
||||||
|
}
|
||||||
|
if req.Icon != nil {
|
||||||
|
scenario.Icon = *req.Icon
|
||||||
|
}
|
||||||
|
if req.Description != nil {
|
||||||
|
scenario.Description = *req.Description
|
||||||
|
}
|
||||||
|
if req.Prompt != nil {
|
||||||
|
scenario.Prompt = *req.Prompt
|
||||||
|
}
|
||||||
|
if req.Greeting != nil {
|
||||||
|
scenario.Greeting = *req.Greeting
|
||||||
|
}
|
||||||
|
if req.Language != nil {
|
||||||
|
scenario.Language = *req.Language
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.repo.Update(c.Request.Context(), scenario); err != nil {
|
||||||
|
logger.Log.Errorw("更新用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "更新失败"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Log.Infow("更新用户情景成功", "user_id", userID, "scenario_id", scenarioID)
|
||||||
|
c.JSON(http.StatusOK, scenario)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete 删除用户情景。
|
||||||
|
// DELETE /api/scenarios/:id
|
||||||
|
func (h *UserScenarioHandler) Delete(c *gin.Context) {
|
||||||
|
userID, exists := c.Get("user_id")
|
||||||
|
if !exists {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scenarioID := c.Param("id")
|
||||||
|
|
||||||
|
// 查询并校验所有权
|
||||||
|
_, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||||
|
if err != nil {
|
||||||
|
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||||
|
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.repo.Delete(c.Request.Context(), scenarioID); err != nil {
|
||||||
|
logger.Log.Errorw("删除用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "删除失败"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Log.Infow("删除用户情景成功", "user_id", userID, "scenario_id", scenarioID)
|
||||||
|
c.Status(http.StatusNoContent)
|
||||||
|
}
|
||||||
@@ -15,10 +15,17 @@ var (
|
|||||||
ErrInvalidToken = errors.New("invalid or expired token")
|
ErrInvalidToken = errors.New("invalid or expired token")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// 令牌类型常量。
|
||||||
|
const (
|
||||||
|
TokenTypeAccess = "access"
|
||||||
|
TokenTypeRefresh = "refresh"
|
||||||
|
)
|
||||||
|
|
||||||
// Claims JWT 声明。
|
// Claims JWT 声明。
|
||||||
type Claims struct {
|
type Claims struct {
|
||||||
UserID string `json:"user_id"`
|
UserID string `json:"user_id"`
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
|
TokenType string `json:"token_type"`
|
||||||
jwt.RegisteredClaims
|
jwt.RegisteredClaims
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -45,8 +52,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
|
|
||||||
// access token
|
// access token
|
||||||
accessClaims := &Claims{
|
accessClaims := &Claims{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
Username: username,
|
Username: username,
|
||||||
|
TokenType: TokenTypeAccess,
|
||||||
RegisteredClaims: jwt.RegisteredClaims{
|
RegisteredClaims: jwt.RegisteredClaims{
|
||||||
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
|
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
|
||||||
IssuedAt: jwt.NewNumericDate(now),
|
IssuedAt: jwt.NewNumericDate(now),
|
||||||
@@ -62,8 +70,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
// refresh token(含唯一 token_id 用于 DB 关联)
|
// refresh token(含唯一 token_id 用于 DB 关联)
|
||||||
tokenID := uuid.New().String()
|
tokenID := uuid.New().String()
|
||||||
refreshClaims := &Claims{
|
refreshClaims := &Claims{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
Username: username,
|
Username: username,
|
||||||
|
TokenType: TokenTypeRefresh,
|
||||||
RegisteredClaims: jwt.RegisteredClaims{
|
RegisteredClaims: jwt.RegisteredClaims{
|
||||||
ID: tokenID,
|
ID: tokenID,
|
||||||
ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)),
|
ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)),
|
||||||
@@ -78,12 +87,26 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
|
|||||||
|
|
||||||
// ValidateAccess 校验 access token 并返回 Claims。
|
// ValidateAccess 校验 access token 并返回 Claims。
|
||||||
func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) {
|
func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) {
|
||||||
return tm.validate(tokenStr)
|
claims, err := tm.validate(tokenStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if claims.TokenType != TokenTypeAccess {
|
||||||
|
return nil, ErrInvalidToken
|
||||||
|
}
|
||||||
|
return claims, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ValidateRefresh 校验 refresh token 并返回 Claims。
|
// ValidateRefresh 校验 refresh token 并返回 Claims。
|
||||||
func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) {
|
func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) {
|
||||||
return tm.validate(tokenStr)
|
claims, err := tm.validate(tokenStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if claims.TokenType != TokenTypeRefresh {
|
||||||
|
return nil, ErrInvalidToken
|
||||||
|
}
|
||||||
|
return claims, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// validate 解析并校验 JWT。
|
// validate 解析并校验 JWT。
|
||||||
|
|||||||
@@ -103,6 +103,44 @@ func TestValidateRefresh_ExpiredToken(t *testing.T) {
|
|||||||
assert.ErrorIs(t, err, ErrInvalidToken)
|
assert.ErrorIs(t, err, ErrInvalidToken)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidateAccess_RejectsRefreshToken(t *testing.T) {
|
||||||
|
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
|
||||||
|
|
||||||
|
_, refresh, err := tm.GeneratePair("user-123", "alice")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// refresh token 不能通过 access 校验
|
||||||
|
_, err = tm.ValidateAccess(refresh)
|
||||||
|
assert.ErrorIs(t, err, ErrInvalidToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateRefresh_RejectsAccessToken(t *testing.T) {
|
||||||
|
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
|
||||||
|
|
||||||
|
access, _, err := tm.GeneratePair("user-123", "alice")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// access token 不能通过 refresh 校验
|
||||||
|
_, err = tm.ValidateRefresh(access)
|
||||||
|
assert.ErrorIs(t, err, ErrInvalidToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGeneratePair_TokenTypesAreCorrect(t *testing.T) {
|
||||||
|
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
|
||||||
|
|
||||||
|
access, refresh, err := tm.GeneratePair("user-123", "alice")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// 通过 validate(不做类型检查)验证 token_type 字段
|
||||||
|
accessClaims, err := tm.validate(access)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, TokenTypeAccess, accessClaims.TokenType)
|
||||||
|
|
||||||
|
refreshClaims, err := tm.validate(refresh)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, TokenTypeRefresh, refreshClaims.TokenType)
|
||||||
|
}
|
||||||
|
|
||||||
func TestGeneratePair_TokenClaimsContainCorrectExpiry(t *testing.T) {
|
func TestGeneratePair_TokenClaimsContainCorrectExpiry(t *testing.T) {
|
||||||
accessTTL := 15 * time.Minute
|
accessTTL := 15 * time.Minute
|
||||||
refreshTTL := 7 * 24 * time.Hour
|
refreshTTL := 7 * 24 * time.Hour
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// contextKey 用于在 Gin context 中存储 Claims 的 key。
|
// contextKey 用于在 Gin context 中存储 Claims 的 key。
|
||||||
@@ -17,8 +19,13 @@ const (
|
|||||||
// 校验成功后将 user_id 和 username 写入 Gin Context。
|
// 校验成功后将 user_id 和 username 写入 Gin Context。
|
||||||
func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
authHeader := c.GetHeader("Authorization")
|
authHeader := c.GetHeader("Authorization")
|
||||||
if authHeader == "" {
|
if authHeader == "" {
|
||||||
|
log.Warnw("auth rejected",
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
"path", c.Request.URL.Path,
|
||||||
|
"reason", "missing authorization header")
|
||||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||||
"code": "INVALID_TOKEN",
|
"code": "INVALID_TOKEN",
|
||||||
"message": "missing authorization header",
|
"message": "missing authorization header",
|
||||||
@@ -29,6 +36,10 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
|||||||
// 提取 Bearer token
|
// 提取 Bearer token
|
||||||
parts := strings.SplitN(authHeader, " ", 2)
|
parts := strings.SplitN(authHeader, " ", 2)
|
||||||
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
|
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
|
||||||
|
log.Warnw("auth rejected",
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
"path", c.Request.URL.Path,
|
||||||
|
"reason", "invalid authorization format")
|
||||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||||
"code": "INVALID_TOKEN",
|
"code": "INVALID_TOKEN",
|
||||||
"message": "invalid authorization format",
|
"message": "invalid authorization format",
|
||||||
@@ -38,6 +49,11 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
|||||||
|
|
||||||
claims, err := tokenMgr.ValidateAccess(parts[1])
|
claims, err := tokenMgr.ValidateAccess(parts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Warnw("auth rejected",
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
"path", c.Request.URL.Path,
|
||||||
|
"reason", "invalid or expired token",
|
||||||
|
"error", err)
|
||||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||||
"code": "INVALID_TOKEN",
|
"code": "INVALID_TOKEN",
|
||||||
"message": "invalid or expired token",
|
"message": "invalid or expired token",
|
||||||
|
|||||||
@@ -166,6 +166,9 @@ func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthRes
|
|||||||
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
|
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, store.ErrRefreshTokenNotFound) {
|
if errors.Is(err, store.ErrRefreshTokenNotFound) {
|
||||||
|
// JWT 校验已通过但 DB 中不存在 → token 已被 rotation 删除,属于复用行为
|
||||||
|
// 吊销该用户全部 refresh token,强制所有设备重新登录
|
||||||
|
_ = s.userRepo.DeleteUserRefreshTokens(ctx, claims.UserID)
|
||||||
return nil, ErrRefreshTokenUsed
|
return nil, ErrRefreshTokenUsed
|
||||||
}
|
}
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|||||||
@@ -188,3 +188,48 @@ func TestLogout_Success(t *testing.T) {
|
|||||||
})
|
})
|
||||||
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
|
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// --- Refresh Token 复用检测 ---
|
||||||
|
|
||||||
|
func TestRefresh_ReuseDetectedRevokesAllTokens(t *testing.T) {
|
||||||
|
svc, repo := newTestService(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// 注册,获得令牌对 A
|
||||||
|
regResp, err := svc.Register(ctx, auth.RegisterRequest{
|
||||||
|
Username: "eve",
|
||||||
|
Password: "password123",
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
tokenPairA_refresh := regResp.RefreshToken
|
||||||
|
|
||||||
|
// 再次登录,获得令牌对 B
|
||||||
|
loginResp, err := svc.Login(ctx, auth.LoginRequest{
|
||||||
|
Username: "eve",
|
||||||
|
Password: "password123",
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
tokenPairB_refresh := loginResp.RefreshToken
|
||||||
|
|
||||||
|
// 用令牌对 A 的 refresh token 正常刷新 → 成功
|
||||||
|
refreshResp, err := svc.Refresh(ctx, auth.RefreshRequest{
|
||||||
|
RefreshToken: tokenPairA_refresh,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, refreshResp.AccessToken)
|
||||||
|
|
||||||
|
// 用令牌对 A 的旧 refresh token 再次刷新 → 复用检测,应失败
|
||||||
|
_, err = svc.Refresh(ctx, auth.RefreshRequest{
|
||||||
|
RefreshToken: tokenPairA_refresh,
|
||||||
|
})
|
||||||
|
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
|
||||||
|
|
||||||
|
// 令牌对 B 的 refresh token 也应被吊销(全量吊销)
|
||||||
|
_, err = svc.Refresh(ctx, auth.RefreshRequest{
|
||||||
|
RefreshToken: tokenPairB_refresh,
|
||||||
|
})
|
||||||
|
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
|
||||||
|
|
||||||
|
// 确认 DB 中该用户已无 refresh token
|
||||||
|
_ = repo // repo 用于确认,但 MemUserRepository 无直接查询方法,通过 Refresh 失败已间接验证
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,22 +2,23 @@ package config
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
|
|
||||||
|
"github.com/joho/godotenv"
|
||||||
"github.com/spf13/viper"
|
"github.com/spf13/viper"
|
||||||
)
|
)
|
||||||
|
|
||||||
// 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 会话管理配置。
|
||||||
@@ -91,10 +92,23 @@ type TTSConfig struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type StorageConfig struct {
|
type StorageConfig struct {
|
||||||
|
Redis RedisStorageConfig `mapstructure:"redis"`
|
||||||
|
Persistence PersistenceConfig `mapstructure:"persistence"`
|
||||||
|
// Deprecated: 使用 Redis 和 Persistence 替代
|
||||||
Driver string `mapstructure:"driver"`
|
Driver string `mapstructure:"driver"`
|
||||||
DSN string `mapstructure:"dsn"`
|
DSN string `mapstructure:"dsn"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type RedisStorageConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type PersistenceConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
Driver string `mapstructure:"driver"`
|
||||||
|
DSN string `mapstructure:"dsn"`
|
||||||
|
}
|
||||||
|
|
||||||
type LogConfig struct {
|
type LogConfig struct {
|
||||||
Level string `mapstructure:"level"`
|
Level string `mapstructure:"level"`
|
||||||
Format string `mapstructure:"format"`
|
Format string `mapstructure:"format"`
|
||||||
@@ -107,90 +121,152 @@ type AuthConfig struct {
|
|||||||
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 10080(7天)
|
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 10080(7天)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml。
|
// RateLimitConfig 限流配置。
|
||||||
func Load() (*Config, error) {
|
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 > 默认值。
|
||||||
|
// workDir 为项目根目录或 backend 目录,用于定位 .env 和 config/config.yaml。
|
||||||
|
func Load(workDir string) (*Config, error) {
|
||||||
|
// 1. 加载 .env 文件(敏感信息)
|
||||||
|
envFile := filepath.Join(workDir, ".env")
|
||||||
|
_ = godotenv.Load(envFile) // 文件不存在也不报错
|
||||||
|
|
||||||
v := viper.New()
|
v := viper.New()
|
||||||
v.SetConfigName("config")
|
v.SetConfigName("config")
|
||||||
v.SetConfigType("yaml")
|
v.SetConfigType("yaml")
|
||||||
v.AddConfigPath(".")
|
v.AddConfigPath(filepath.Join(workDir, "config")) // 配置文件在 config/ 目录下
|
||||||
v.AddConfigPath("./config")
|
v.AddConfigPath(workDir) // 兼容旧路径
|
||||||
v.AddConfigPath("./backend")
|
|
||||||
v.AddConfigPath("..") // 兼容从 backend/cmd/ 启动
|
|
||||||
v.AddConfigPath("../..") // 兼容从 backend/cmd/server/ 启动
|
|
||||||
|
|
||||||
// 默认值
|
// 2. 设置默认值(与 config.yaml 保持一致,仅作为兜底)
|
||||||
v.SetDefault("app.env", "dev")
|
setDefaults(v)
|
||||||
v.SetDefault("app.version", "dev")
|
|
||||||
v.SetDefault("server.host", "0.0.0.0")
|
|
||||||
v.SetDefault("server.port", 8080)
|
|
||||||
v.SetDefault("server.read_timeout", 30)
|
|
||||||
v.SetDefault("server.write_timeout", 30)
|
|
||||||
v.SetDefault("server.heartbeat_interval", 30)
|
|
||||||
v.SetDefault("server.heartbeat_timeout", 60)
|
|
||||||
v.SetDefault("server.shutdown_timeout", 10)
|
|
||||||
v.SetDefault("session.ttl", 30)
|
|
||||||
v.SetDefault("session.max_history", 20)
|
|
||||||
v.SetDefault("redis.addr", "localhost:6379")
|
|
||||||
v.SetDefault("redis.db", 0)
|
|
||||||
v.SetDefault("ai.stt.provider", "deepgram")
|
|
||||||
v.SetDefault("ai.stt.model", "nova-2")
|
|
||||||
v.SetDefault("ai.stt.endpoint", "wss://api.deepgram.com/v1/listen")
|
|
||||||
v.SetDefault("ai.stt.timeout", 5)
|
|
||||||
v.SetDefault("ai.stt.http_client_timeout", 30)
|
|
||||||
v.SetDefault("ai.llm.provider", "openai")
|
|
||||||
v.SetDefault("ai.llm.model", "gpt-4o")
|
|
||||||
v.SetDefault("ai.llm.endpoint", "https://api.openai.com/v1")
|
|
||||||
v.SetDefault("ai.llm.timeout", 10)
|
|
||||||
v.SetDefault("ai.llm.http_client_timeout", 60)
|
|
||||||
v.SetDefault("ai.tts.provider", "openai")
|
|
||||||
v.SetDefault("ai.tts.model", "tts-1")
|
|
||||||
v.SetDefault("ai.tts.voice", "mimo_default")
|
|
||||||
v.SetDefault("ai.tts.speed", 1.0)
|
|
||||||
v.SetDefault("ai.tts.endpoint", "https://api.openai.com/v1")
|
|
||||||
v.SetDefault("ai.tts.timeout", 5)
|
|
||||||
v.SetDefault("ai.tts.http_client_timeout", 30)
|
|
||||||
v.SetDefault("ai.tts.output_format", "mp3")
|
|
||||||
v.SetDefault("ai.tts.sample_rate", 24000)
|
|
||||||
v.SetDefault("storage.driver", "memory")
|
|
||||||
v.SetDefault("storage.dsn", "")
|
|
||||||
v.SetDefault("log.level", "info")
|
|
||||||
v.SetDefault("log.format", "console")
|
|
||||||
v.SetDefault("auth.access_ttl", 15)
|
|
||||||
v.SetDefault("auth.refresh_ttl", 10080)
|
|
||||||
|
|
||||||
// 读取基础配置文件
|
// 3. 读取 config.yaml
|
||||||
_ = v.ReadInConfig() // 文件不存在不报错
|
if err := v.ReadInConfig(); err != nil {
|
||||||
|
return nil, fmt.Errorf("config: read config.yaml: %w", err)
|
||||||
// 根据 APP_ENV 覆盖
|
|
||||||
env := os.Getenv("APP_ENV")
|
|
||||||
if env == "" {
|
|
||||||
env = v.GetString("app.env")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 4. 合并环境专属配置 config.{env}.yaml(可选)
|
||||||
|
env := v.GetString("app.env")
|
||||||
if env != "" {
|
if env != "" {
|
||||||
v.SetConfigName("config." + env)
|
v.SetConfigName("config." + env)
|
||||||
_ = v.MergeInConfig()
|
_ = v.MergeInConfig() // 文件不存在也不报错
|
||||||
}
|
}
|
||||||
|
|
||||||
// 环境变量覆盖
|
// 5. 显式绑定敏感信息环境变量(不用 AutomaticEnv,避免隐式映射)
|
||||||
v.SetEnvPrefix("CAMTALK")
|
bindEnvVars(v)
|
||||||
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
|
|
||||||
v.AutomaticEnv()
|
|
||||||
|
|
||||||
var cfg Config
|
var cfg Config
|
||||||
if err := v.Unmarshal(&cfg); err != nil {
|
if err := v.Unmarshal(&cfg); err != nil {
|
||||||
return nil, fmt.Errorf("config unmarshal: %w", err)
|
return nil, fmt.Errorf("config: unmarshal: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
// 填充默认值
|
|
||||||
if cfg.Server.Host == "" {
|
|
||||||
cfg.Server.Host = "0.0.0.0"
|
|
||||||
}
|
|
||||||
if cfg.Server.Port == 0 {
|
|
||||||
cfg.Server.Port = 8080
|
|
||||||
}
|
|
||||||
if cfg.App.Env == "" {
|
|
||||||
cfg.App.Env = "dev"
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return &cfg, nil
|
return &cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// setDefaults 设置兜底默认值,与 config.yaml 保持一致。
|
||||||
|
func setDefaults(v *viper.Viper) {
|
||||||
|
// app
|
||||||
|
v.SetDefault("app.env", "dev")
|
||||||
|
v.SetDefault("app.version", "dev")
|
||||||
|
|
||||||
|
// server
|
||||||
|
v.SetDefault("server.host", "0.0.0.0")
|
||||||
|
v.SetDefault("server.port", 8080)
|
||||||
|
v.SetDefault("server.read_timeout", 30)
|
||||||
|
v.SetDefault("server.write_timeout", 30)
|
||||||
|
v.SetDefault("server.shutdown_timeout", 10)
|
||||||
|
v.SetDefault("server.heartbeat_interval", 30)
|
||||||
|
v.SetDefault("server.heartbeat_timeout", 60)
|
||||||
|
|
||||||
|
// session
|
||||||
|
v.SetDefault("session.ttl", 30)
|
||||||
|
v.SetDefault("session.max_history", 20)
|
||||||
|
|
||||||
|
// ai — 默认值与 config.yaml 一致(mimo/dashscope)
|
||||||
|
v.SetDefault("ai.stt.provider", "mimo")
|
||||||
|
v.SetDefault("ai.stt.model", "mimo-v2.5-asr")
|
||||||
|
v.SetDefault("ai.stt.endpoint", "https://api.xiaomimimo.com/v1")
|
||||||
|
v.SetDefault("ai.stt.timeout", 5)
|
||||||
|
v.SetDefault("ai.stt.http_client_timeout", 30)
|
||||||
|
|
||||||
|
v.SetDefault("ai.llm.provider", "dashscope")
|
||||||
|
v.SetDefault("ai.llm.model", "qwen3-vl-plus")
|
||||||
|
v.SetDefault("ai.llm.endpoint", "https://dashscope.aliyuncs.com/compatible-mode/v1")
|
||||||
|
v.SetDefault("ai.llm.timeout", 30)
|
||||||
|
v.SetDefault("ai.llm.http_client_timeout", 60)
|
||||||
|
|
||||||
|
v.SetDefault("ai.tts.provider", "mimo")
|
||||||
|
v.SetDefault("ai.tts.model", "mimo-v2.5-tts")
|
||||||
|
v.SetDefault("ai.tts.voice", "mimo_default")
|
||||||
|
v.SetDefault("ai.tts.speed", 1.0)
|
||||||
|
v.SetDefault("ai.tts.endpoint", "https://token-plan-cn.xiaomimimo.com/v1")
|
||||||
|
v.SetDefault("ai.tts.timeout", 5)
|
||||||
|
v.SetDefault("ai.tts.http_client_timeout", 30)
|
||||||
|
v.SetDefault("ai.tts.output_format", "mp3")
|
||||||
|
v.SetDefault("ai.tts.sample_rate", 24000)
|
||||||
|
|
||||||
|
// storage
|
||||||
|
v.SetDefault("storage.driver", "memory")
|
||||||
|
v.SetDefault("storage.redis.enabled", false)
|
||||||
|
v.SetDefault("storage.persistence.enabled", false)
|
||||||
|
v.SetDefault("storage.persistence.driver", "postgres")
|
||||||
|
|
||||||
|
// redis
|
||||||
|
v.SetDefault("redis.addr", "localhost:6379")
|
||||||
|
v.SetDefault("redis.password", "")
|
||||||
|
v.SetDefault("redis.db", 0)
|
||||||
|
|
||||||
|
// auth
|
||||||
|
v.SetDefault("auth.access_ttl", 15)
|
||||||
|
v.SetDefault("auth.refresh_ttl", 10080)
|
||||||
|
|
||||||
|
// log
|
||||||
|
v.SetDefault("log.level", "info")
|
||||||
|
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 显式绑定敏感信息环境变量。
|
||||||
|
// 只绑定不应出现在 config.yaml 中的敏感字段,非敏感配置通过 config.yaml 管理。
|
||||||
|
func bindEnvVars(v *viper.Viper) {
|
||||||
|
// app.env 特殊处理:环境变量 APP_ENV 覆盖 config.yaml 中的 app.env
|
||||||
|
v.BindEnv("app.env", "APP_ENV")
|
||||||
|
|
||||||
|
// AI API Key
|
||||||
|
v.BindEnv("ai.stt.api_key", "CAMTALK_AI_STT_API_KEY")
|
||||||
|
v.BindEnv("ai.llm.api_key", "CAMTALK_AI_LLM_API_KEY")
|
||||||
|
v.BindEnv("ai.tts.api_key", "CAMTALK_AI_TTS_API_KEY")
|
||||||
|
|
||||||
|
// JWT
|
||||||
|
v.BindEnv("auth.jwt_secret", "CAMTALK_AUTH_JWT_SECRET")
|
||||||
|
|
||||||
|
// 数据库
|
||||||
|
v.BindEnv("storage.dsn", "CAMTALK_STORAGE_DSN")
|
||||||
|
v.BindEnv("storage.persistence.dsn", "CAMTALK_STORAGE_DSN")
|
||||||
|
v.BindEnv("storage.redis.enabled", "CAMTALK_STORAGE_REDIS_ENABLED")
|
||||||
|
v.BindEnv("storage.persistence.enabled", "CAMTALK_STORAGE_PERSISTENCE_ENABLED")
|
||||||
|
v.BindEnv("storage.persistence.driver", "CAMTALK_STORAGE_PERSISTENCE_DRIVER")
|
||||||
|
|
||||||
|
// Redis(密码可能包含特殊字符,通过环境变量设置更安全)
|
||||||
|
v.BindEnv("redis.addr", "CAMTALK_REDIS_ADDR")
|
||||||
|
v.BindEnv("redis.password", "CAMTALK_REDIS_PASSWORD")
|
||||||
|
}
|
||||||
|
|||||||
172
backend/internal/eino/adapter.go
Normal file
172
backend/internal/eino/adapter.go
Normal file
@@ -0,0 +1,172 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"io"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/orchestrator"
|
||||||
|
"github.com/hhs/camtalk/internal/session"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// EinoOrchestrator 实现 orchestrator.Orchestrator 接口。
|
||||||
|
// 将 Eino Graph 包装为现有接口,WS Handler 几乎不用改。
|
||||||
|
type EinoOrchestrator struct {
|
||||||
|
graph *PipelineGraph
|
||||||
|
sessionMgr session.Manager
|
||||||
|
model string
|
||||||
|
callbacks compose.Option // 运行时 Callback option
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewEinoOrchestrator 创建 Eino 编排器适配器。
|
||||||
|
func NewEinoOrchestrator(graph *PipelineGraph, sessionMgr session.Manager, model string) *EinoOrchestrator {
|
||||||
|
return &EinoOrchestrator{
|
||||||
|
graph: graph,
|
||||||
|
sessionMgr: sessionMgr,
|
||||||
|
model: model,
|
||||||
|
callbacks: compose.WithCallbacks(BuildCallbackHandler()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProcessQuery 实现 orchestrator.Orchestrator 接口。
|
||||||
|
func (e *EinoOrchestrator) ProcessQuery(
|
||||||
|
ctx context.Context,
|
||||||
|
sessionID string,
|
||||||
|
req models.WsQuery,
|
||||||
|
sender orchestrator.Sender,
|
||||||
|
) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
startTime := time.Now()
|
||||||
|
|
||||||
|
// 1. 设置活跃请求
|
||||||
|
if err := e.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer e.sessionMgr.ClearActiveRequest(ctx, sessionID)
|
||||||
|
|
||||||
|
// 2. 获取会话配置
|
||||||
|
sess, err := e.sessionMgr.Get(ctx, sessionID)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("get session failed", "error", err)
|
||||||
|
sender.SendError(models.WsError{
|
||||||
|
Type: "error",
|
||||||
|
RequestID: req.RequestID,
|
||||||
|
Code: "SESSION_NOT_FOUND",
|
||||||
|
Message: "会话不存在",
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. 解码音频和图片
|
||||||
|
var audioData []byte
|
||||||
|
if req.Text == "" && req.Audio != "" {
|
||||||
|
audioData, err = base64.StdEncoding.DecodeString(req.Audio)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("audio decode failed", "error", err)
|
||||||
|
sender.SendError(models.WsError{
|
||||||
|
Type: "error",
|
||||||
|
RequestID: req.RequestID,
|
||||||
|
Code: "INVALID_MESSAGE",
|
||||||
|
Message: "音频数据解码失败",
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var imageData []byte
|
||||||
|
if req.Image != "" {
|
||||||
|
imageData, err = base64.StdEncoding.DecodeString(req.Image)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("image decode failed", "error", err)
|
||||||
|
sender.SendError(models.WsError{
|
||||||
|
Type: "error",
|
||||||
|
RequestID: req.RequestID,
|
||||||
|
Code: "INVALID_MESSAGE",
|
||||||
|
Message: "图片数据解码失败",
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 4. 构建 Graph 输入
|
||||||
|
input := buildPipelineInput(req, sessionID, sess, audioData, imageData)
|
||||||
|
|
||||||
|
// 5. 注入 context 值(供 Callback 和 Lambda 节点使用)
|
||||||
|
ctx = WithSender(ctx, sender)
|
||||||
|
ctx = WithRequestID(ctx, req.RequestID)
|
||||||
|
ctx = trace.WithSessionID(ctx, sessionID)
|
||||||
|
ctx = WithStartTime(ctx, startTime)
|
||||||
|
|
||||||
|
// 创建 State 并从 input 复制元数据
|
||||||
|
state := genLocalState(ctx)
|
||||||
|
state.SessionID = input.SessionID
|
||||||
|
state.RequestID = input.RequestID
|
||||||
|
state.ImageData = input.ImageData
|
||||||
|
state.Scenario = input.Scenario
|
||||||
|
state.Language = input.Language
|
||||||
|
state.DetailLevel = sess.Config.DetailLevel
|
||||||
|
state.TTSEnabled = input.TTSEnabled
|
||||||
|
state.UserID = input.UserID
|
||||||
|
ctx = WithPipelineState(ctx, state)
|
||||||
|
|
||||||
|
// 6. 调用 Graph(Stream 模式 + 运行时 Callback)
|
||||||
|
streamReader, err := e.graph.Runnable.Stream(ctx, input, e.callbacks)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("graph stream start failed", "error", err)
|
||||||
|
sender.SendError(models.WsError{
|
||||||
|
Type: "error",
|
||||||
|
RequestID: req.RequestID,
|
||||||
|
Code: "INTERNAL_ERROR",
|
||||||
|
Message: "编排器启动失败",
|
||||||
|
})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 7. 消费 StreamReader(触发整条链路执行,side effects 推送消息到客户端)
|
||||||
|
var output PipelineOutput
|
||||||
|
for {
|
||||||
|
o, err := streamReader.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
log.Errorw("graph stream consume error", "error", err)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
output = o
|
||||||
|
}
|
||||||
|
|
||||||
|
// 8. 追加用户消息到历史(使用 STT 结果,兼容文本输入和语音输入)
|
||||||
|
userText := output.TranscribedText
|
||||||
|
if userText == "" {
|
||||||
|
userText = req.Text // fallback 到原始文本输入
|
||||||
|
}
|
||||||
|
if userText != "" {
|
||||||
|
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
||||||
|
Role: "user",
|
||||||
|
Content: userText,
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("append user message failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 9. 追加助手消息到历史
|
||||||
|
if output.FullResponse != "" {
|
||||||
|
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: output.FullResponse,
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("append assistant message failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
latency := time.Since(startTime).Milliseconds()
|
||||||
|
log.Infow("eino pipeline completed", "latency_ms", latency)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
131
backend/internal/eino/callback.go
Normal file
131
backend/internal/eino/callback.go
Normal file
@@ -0,0 +1,131 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/callbacks"
|
||||||
|
"github.com/cloudwego/eino/components/model"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
callbacksHelper "github.com/cloudwego/eino/utils/callbacks"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/orchestrator"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// context key 类型,避免与其他包冲突。
|
||||||
|
type ctxKeySender struct{}
|
||||||
|
type ctxKeyState struct{}
|
||||||
|
|
||||||
|
// WithSender 将 Sender 注入 context。
|
||||||
|
func WithSender(ctx context.Context, sender orchestrator.Sender) context.Context {
|
||||||
|
return context.WithValue(ctx, ctxKeySender{}, sender)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithRequestID 将 requestID 注入 context(使用 trace 包)。
|
||||||
|
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||||
|
return trace.WithRequestID(ctx, requestID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithPipelineState 将 PipelineState 注入 context。
|
||||||
|
func WithPipelineState(ctx context.Context, state *PipelineState) context.Context {
|
||||||
|
return context.WithValue(ctx, ctxKeyState{}, state)
|
||||||
|
}
|
||||||
|
|
||||||
|
// senderFromCtx 从 context 获取 Sender。
|
||||||
|
func senderFromCtx(ctx context.Context) orchestrator.Sender {
|
||||||
|
s, _ := ctx.Value(ctxKeySender{}).(orchestrator.Sender)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// requestIDFromCtx 从 context 获取 requestID(使用 trace 包)。
|
||||||
|
func requestIDFromCtx(ctx context.Context) string {
|
||||||
|
return trace.GetRequestID(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// stateFromCtx 从 context 获取 PipelineState。
|
||||||
|
func stateFromCtx(ctx context.Context) *PipelineState {
|
||||||
|
s, _ := ctx.Value(ctxKeyState{}).(*PipelineState)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildCallbackHandler 构建 Eino Callback Handler。
|
||||||
|
//
|
||||||
|
// 核心职责:ChatModel 节点通过 OnEndWithStreamOutput 逐 token 推送 llm_chunk 到客户端,
|
||||||
|
// 同时累积完整文本到 PipelineState。
|
||||||
|
//
|
||||||
|
// 其他节点的消息推送(stt_result、tts_audio、llm_done)由各 Lambda 内部直接调用 Sender。
|
||||||
|
func BuildCallbackHandler() callbacks.Handler {
|
||||||
|
return callbacksHelper.NewHandlerHelper().
|
||||||
|
ChatModel(&callbacksHelper.ModelCallbackHandler{
|
||||||
|
OnEndWithStreamOutput: func(ctx context.Context, info *callbacks.RunInfo, output *schema.StreamReader[*model.CallbackOutput]) context.Context {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
sender := senderFromCtx(ctx)
|
||||||
|
requestID := requestIDFromCtx(ctx)
|
||||||
|
state := stateFromCtx(ctx)
|
||||||
|
|
||||||
|
if sender == nil || requestID == "" {
|
||||||
|
log.Warnw("ModelCallback: missing sender or request_id in context",
|
||||||
|
"node", info.Name)
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
// 异步消费流,避免阻塞框架的下游处理。
|
||||||
|
// 框架对流做了内部拷贝,此 goroutine 读取独立副本。
|
||||||
|
go func() {
|
||||||
|
defer output.Close()
|
||||||
|
|
||||||
|
for {
|
||||||
|
chunk, err := output.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Errorw("ModelCallback: stream recv error",
|
||||||
|
"node", info.Name, "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if chunk == nil || chunk.Message == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
delta := chunk.Message.Content
|
||||||
|
if delta == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// 推送 llm_chunk 到客户端
|
||||||
|
if err := sender.SendLLMChunk(models.WsLLMChunk{
|
||||||
|
Type: "llm_chunk",
|
||||||
|
RequestID: requestID,
|
||||||
|
Delta: delta,
|
||||||
|
Role: "assistant",
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("ModelCallback: send llm_chunk failed", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 累积完整文本到 State
|
||||||
|
if state != nil {
|
||||||
|
state.AppendText(delta)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 记录 token 用量(流的最后一帧携带)
|
||||||
|
if chunk.TokenUsage != nil && state != nil {
|
||||||
|
state.mu.Lock()
|
||||||
|
state.TokenUsage = &TokenUsage{
|
||||||
|
Prompt: chunk.TokenUsage.PromptTokens,
|
||||||
|
Completion: chunk.TokenUsage.CompletionTokens,
|
||||||
|
Total: chunk.TokenUsage.TotalTokens,
|
||||||
|
}
|
||||||
|
state.mu.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return ctx
|
||||||
|
},
|
||||||
|
}).
|
||||||
|
Handler()
|
||||||
|
}
|
||||||
119
backend/internal/eino/graph.go
Normal file
119
backend/internal/eino/graph.go
Normal file
@@ -0,0 +1,119 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
openaiImpl "github.com/cloudwego/eino-ext/components/model/openai"
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/stt"
|
||||||
|
"github.com/hhs/camtalk/internal/ai/tts"
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/session"
|
||||||
|
"github.com/hhs/camtalk/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
nodeSTT = "stt"
|
||||||
|
nodeHistory = "history"
|
||||||
|
nodeLLM = "llm"
|
||||||
|
nodeMessageToString = "msg2str"
|
||||||
|
nodeSplitter = "splitter"
|
||||||
|
nodeTTS = "tts"
|
||||||
|
nodeDone = "done"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PipelineGraph 封装编译后的 Eino Graph。
|
||||||
|
type PipelineGraph struct {
|
||||||
|
Runnable compose.Runnable[PipelineInput, PipelineOutput]
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPipelineGraph 构建 CamTalk AI 编排 Graph。
|
||||||
|
//
|
||||||
|
// 拓扑:START → STT → History → ChatModel → Splitter → TTS → Done → END
|
||||||
|
//
|
||||||
|
// Graph 使用 Stream 模式调用,ChatModel 实现真正的 token 级流式输出。
|
||||||
|
// LLM token 通过 Callback 的 OnEndWithStreamOutput 实时推送到客户端。
|
||||||
|
func NewPipelineGraph(
|
||||||
|
ctx context.Context,
|
||||||
|
cfg *config.Config,
|
||||||
|
sttService stt.Service,
|
||||||
|
ttsService tts.Service,
|
||||||
|
sessionMgr session.Manager,
|
||||||
|
scenarioRepo store.UserScenarioRepository,
|
||||||
|
) (*PipelineGraph, error) {
|
||||||
|
log := logger.Log
|
||||||
|
|
||||||
|
// 1. 创建 eino-ext ChatModel(对接 DashScope OpenAI 兼容接口)
|
||||||
|
chatModel, err := openaiImpl.NewChatModel(ctx, &openaiImpl.ChatModelConfig{
|
||||||
|
APIKey: cfg.AI.LLM.APIKey,
|
||||||
|
Model: cfg.AI.LLM.Model,
|
||||||
|
BaseURL: cfg.AI.LLM.Endpoint,
|
||||||
|
Timeout: time.Duration(cfg.AI.LLM.Timeout) * time.Second,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
log.Infow("Eino ChatModel 初始化成功",
|
||||||
|
"model", cfg.AI.LLM.Model,
|
||||||
|
"endpoint", cfg.AI.LLM.Endpoint)
|
||||||
|
|
||||||
|
// 2. 构建 Graph(值类型,非指针)
|
||||||
|
g := compose.NewGraph[PipelineInput, PipelineOutput](
|
||||||
|
compose.WithGenLocalState(genLocalState),
|
||||||
|
)
|
||||||
|
|
||||||
|
// 3. 添加节点
|
||||||
|
maxHistory := cfg.Session.MaxHistory
|
||||||
|
|
||||||
|
_ = g.AddLambdaNode(nodeSTT, NewSTTLambda(sttService))
|
||||||
|
_ = g.AddLambdaNode(nodeHistory, NewHistoryLambda(sessionMgr.GetHistory, scenarioRepo, maxHistory))
|
||||||
|
_ = g.AddChatModelNode(nodeLLM, chatModel)
|
||||||
|
_ = g.AddLambdaNode(nodeMessageToString, NewMessageToStringLambda())
|
||||||
|
_ = g.AddLambdaNode(nodeSplitter, NewSplitterLambda())
|
||||||
|
_ = g.AddLambdaNode(nodeTTS, NewTTSLambda(
|
||||||
|
ttsService,
|
||||||
|
cfg.AI.TTS.Voice,
|
||||||
|
cfg.AI.TTS.Speed,
|
||||||
|
cfg.AI.TTS.OutputFormat,
|
||||||
|
cfg.AI.TTS.SampleRate,
|
||||||
|
))
|
||||||
|
_ = g.AddLambdaNode(nodeDone, NewDoneLambda(cfg.AI.LLM.Model))
|
||||||
|
|
||||||
|
// 4. 连接边
|
||||||
|
_ = g.AddEdge(compose.START, nodeSTT)
|
||||||
|
_ = g.AddEdge(nodeSTT, nodeHistory)
|
||||||
|
_ = g.AddEdge(nodeHistory, nodeLLM)
|
||||||
|
_ = g.AddEdge(nodeLLM, nodeMessageToString)
|
||||||
|
_ = g.AddEdge(nodeMessageToString, nodeSplitter)
|
||||||
|
_ = g.AddEdge(nodeSplitter, nodeTTS)
|
||||||
|
_ = g.AddEdge(nodeTTS, nodeDone)
|
||||||
|
_ = g.AddEdge(nodeDone, compose.END)
|
||||||
|
|
||||||
|
// 5. 编译(回调在运行时通过 Stream option 传入)
|
||||||
|
runnable, err := g.Compile(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infow("Eino Graph 编译成功", "nodes", 7)
|
||||||
|
return &PipelineGraph{Runnable: runnable}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildPipelineInput 从 WebSocket 请求和会话配置构建 Graph 输入。
|
||||||
|
func buildPipelineInput(req models.WsQuery, sessionID string, sess *models.Session, audioData, imageData []byte) PipelineInput {
|
||||||
|
return PipelineInput{
|
||||||
|
AudioData: audioData,
|
||||||
|
ImageData: imageData,
|
||||||
|
Text: req.Text,
|
||||||
|
SessionID: sessionID,
|
||||||
|
RequestID: req.RequestID,
|
||||||
|
Language: sess.Config.Language,
|
||||||
|
Scenario: sess.Config.Scenario,
|
||||||
|
TTSEnabled: sess.Config.TTSEnabled,
|
||||||
|
UserID: sess.UserID,
|
||||||
|
}
|
||||||
|
}
|
||||||
236
backend/internal/eino/graph_test.go
Normal file
236
backend/internal/eino/graph_test.go
Normal file
@@ -0,0 +1,236 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/mock"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/stt"
|
||||||
|
"github.com/hhs/camtalk/internal/ai/tts"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/orchestrator"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- Mock STT Service ---
|
||||||
|
|
||||||
|
type mockSTTService struct {
|
||||||
|
mock.Mock
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSTTService) Recognize(ctx context.Context, audio []byte, opts stt.Options) (string, error) {
|
||||||
|
args := m.Called(ctx, audio, opts)
|
||||||
|
return args.String(0), args.Error(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Mock TTS Service ---
|
||||||
|
|
||||||
|
type mockTTSService struct {
|
||||||
|
mock.Mock
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockTTSService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts tts.Options) (<-chan tts.Chunk, error) {
|
||||||
|
args := m.Called(ctx, textStream, opts)
|
||||||
|
return args.Get(0).(<-chan tts.Chunk), args.Error(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Mock Sender ---
|
||||||
|
|
||||||
|
type mockSender struct {
|
||||||
|
mock.Mock
|
||||||
|
STTResults []models.WsSTTResult
|
||||||
|
LLMChunks []models.WsLLMChunk
|
||||||
|
LLMDones []models.WsLLMDone
|
||||||
|
TTSAudios []models.WsTTSAudio
|
||||||
|
Errors []models.WsError
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSender) SendSTTResult(result models.WsSTTResult) error {
|
||||||
|
m.STTResults = append(m.STTResults, result)
|
||||||
|
return m.Called(result).Error(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSender) SendLLMChunk(chunk models.WsLLMChunk) error {
|
||||||
|
m.LLMChunks = append(m.LLMChunks, chunk)
|
||||||
|
return m.Called(chunk).Error(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSender) SendLLMDone(done models.WsLLMDone) error {
|
||||||
|
m.LLMDones = append(m.LLMDones, done)
|
||||||
|
return m.Called(done).Error(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSender) SendTTSAudio(audio models.WsTTSAudio) error {
|
||||||
|
m.TTSAudios = append(m.TTSAudios, audio)
|
||||||
|
return m.Called(audio).Error(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSender) SendError(err models.WsError) error {
|
||||||
|
m.Errors = append(m.Errors, err)
|
||||||
|
return m.Called(err).Error(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Tests ---
|
||||||
|
|
||||||
|
func TestDetectImageMimeType(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
data []byte
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{"JPEG", []byte{0xFF, 0xD8, 0xFF, 0xE0}, "image/jpeg"},
|
||||||
|
{"PNG", []byte{0x89, 0x50, 0x4E, 0x47}, "image/png"},
|
||||||
|
{"GIF", []byte{0x47, 0x49, 0x46, 0x38}, "image/gif"},
|
||||||
|
{"WebP", []byte{0x52, 0x49, 0x46, 0x46}, "image/webp"},
|
||||||
|
{"Unknown", []byte{0x00, 0x00, 0x00}, "image/jpeg"},
|
||||||
|
{"Short", []byte{0xFF}, "image/jpeg"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := detectImageMimeType(tt.data)
|
||||||
|
assert.Equal(t, tt.expected, result)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildPipelineInput(t *testing.T) {
|
||||||
|
req := models.WsQuery{
|
||||||
|
Text: "你好",
|
||||||
|
RequestID: "req-1",
|
||||||
|
}
|
||||||
|
sess := &models.Session{
|
||||||
|
Config: models.SessionConfig{
|
||||||
|
Language: "zh-CN",
|
||||||
|
Scenario: "free_chat",
|
||||||
|
TTSEnabled: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
input := buildPipelineInput(req, "sess-1", sess, nil, nil)
|
||||||
|
require.Equal(t, "你好", input.Text)
|
||||||
|
require.Equal(t, "sess-1", input.SessionID)
|
||||||
|
require.Equal(t, "req-1", input.RequestID)
|
||||||
|
require.Equal(t, "zh-CN", input.Language)
|
||||||
|
require.Equal(t, "free_chat", input.Scenario)
|
||||||
|
require.True(t, input.TTSEnabled)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildPipelineInput_WithAudioData(t *testing.T) {
|
||||||
|
req := models.WsQuery{
|
||||||
|
Audio: "base64audio",
|
||||||
|
RequestID: "req-2",
|
||||||
|
}
|
||||||
|
sess := &models.Session{
|
||||||
|
Config: models.SessionConfig{
|
||||||
|
Language: "en",
|
||||||
|
Scenario: "free_chat",
|
||||||
|
TTSEnabled: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
audioData := []byte("fake-audio-bytes")
|
||||||
|
imageData := []byte("fake-image-bytes")
|
||||||
|
|
||||||
|
input := buildPipelineInput(req, "sess-2", sess, audioData, imageData)
|
||||||
|
require.Equal(t, audioData, input.AudioData)
|
||||||
|
require.Equal(t, imageData, input.ImageData)
|
||||||
|
require.False(t, input.TTSEnabled)
|
||||||
|
require.Equal(t, "en", input.Language)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPipelineState_AppendAndGet(t *testing.T) {
|
||||||
|
state := genLocalState(context.Background())
|
||||||
|
|
||||||
|
state.AppendText("Hello ")
|
||||||
|
state.AppendText("World")
|
||||||
|
|
||||||
|
require.Equal(t, "Hello World", state.GetFullResponse())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPipelineState_ConcurrentAccess(t *testing.T) {
|
||||||
|
state := genLocalState(context.Background())
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
state.AppendText("a")
|
||||||
|
}
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
_ = state.GetFullResponse()
|
||||||
|
}
|
||||||
|
|
||||||
|
<-done
|
||||||
|
require.Equal(t, 100, len(state.GetFullResponse()))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestContextInjection(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
sender := &mockSender{}
|
||||||
|
ctx = WithSender(ctx, sender)
|
||||||
|
ctx = WithRequestID(ctx, "req-123")
|
||||||
|
ctx = trace.WithSessionID(ctx, "sess-456")
|
||||||
|
ctx = WithStartTime(ctx, time.Now())
|
||||||
|
ctx = WithPipelineState(ctx, genLocalState(ctx))
|
||||||
|
|
||||||
|
require.NotNil(t, senderFromCtx(ctx))
|
||||||
|
require.Equal(t, "req-123", requestIDFromCtx(ctx))
|
||||||
|
require.NotNil(t, stateFromCtx(ctx))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLatencyFromCtx(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// No start time set
|
||||||
|
require.Equal(t, int64(0), latencyFromCtx(ctx))
|
||||||
|
|
||||||
|
// With start time
|
||||||
|
start := time.Now().Add(-100 * time.Millisecond)
|
||||||
|
ctx = WithStartTime(ctx, start)
|
||||||
|
latency := latencyFromCtx(ctx)
|
||||||
|
require.Greater(t, latency, int64(0))
|
||||||
|
require.Less(t, latency, int64(1000)) // should be < 1 second
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEinoOrchestrator_ImplementsInterface(t *testing.T) {
|
||||||
|
// Compile-time check that EinoOrchestrator implements orchestrator.Orchestrator
|
||||||
|
var _ orchestrator.Orchestrator = (*EinoOrchestrator)(nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewSTTLambda_ReturnsNonNil(t *testing.T) {
|
||||||
|
mockSTT := &mockSTTService{}
|
||||||
|
lambda := NewSTTLambda(mockSTT)
|
||||||
|
require.NotNil(t, lambda)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewHistoryLambda_ReturnsNonNil(t *testing.T) {
|
||||||
|
fetcher := func(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
lambda := NewHistoryLambda(fetcher, nil, 10)
|
||||||
|
require.NotNil(t, lambda)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewSplitterLambda_ReturnsNonNil(t *testing.T) {
|
||||||
|
lambda := NewSplitterLambda()
|
||||||
|
require.NotNil(t, lambda)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewTTSLambda_ReturnsNonNil(t *testing.T) {
|
||||||
|
mockTTS := &mockTTSService{}
|
||||||
|
lambda := NewTTSLambda(mockTTS, "alloy", 1.0, "mp3", 24000)
|
||||||
|
require.NotNil(t, lambda)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewDoneLambda_ReturnsNonNil(t *testing.T) {
|
||||||
|
lambda := NewDoneLambda("test-model")
|
||||||
|
require.NotNil(t, lambda)
|
||||||
|
}
|
||||||
86
backend/internal/eino/nodes_done.go
Normal file
86
backend/internal/eino/nodes_done.go
Normal file
@@ -0,0 +1,86 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ctxKeyStartTime 请求开始时间的 context key。
|
||||||
|
type ctxKeyStartTime struct{}
|
||||||
|
|
||||||
|
// WithStartTime 将请求开始时间注入 context。
|
||||||
|
func WithStartTime(ctx context.Context, t time.Time) context.Context {
|
||||||
|
return context.WithValue(ctx, ctxKeyStartTime{}, t)
|
||||||
|
}
|
||||||
|
|
||||||
|
// latencyFromCtx 从 context 获取开始时间并计算延迟(毫秒)。
|
||||||
|
func latencyFromCtx(ctx context.Context) int64 {
|
||||||
|
if startTime, ok := ctx.Value(ctxKeyStartTime{}).(time.Time); ok {
|
||||||
|
return time.Since(startTime).Milliseconds()
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDoneLambda 创建 Done Lambda 节点。
|
||||||
|
// 输入: struct{}(TTS 完成信号)→ 输出: *PipelineOutput
|
||||||
|
//
|
||||||
|
// 从 PipelineState 读取完整回复和 token 用量,发送 llm_done 到客户端。
|
||||||
|
// 历史消息追加由适配器负责(避免重复写入)。
|
||||||
|
func NewDoneLambda(defaultModel string) *compose.Lambda {
|
||||||
|
return compose.InvokableLambda(func(ctx context.Context, _ struct{}) (PipelineOutput, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
sender := senderFromCtx(ctx)
|
||||||
|
state := stateFromCtx(ctx)
|
||||||
|
|
||||||
|
if state == nil {
|
||||||
|
return PipelineOutput{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
state.mu.Lock()
|
||||||
|
fullResponse := state.FullResponse.String()
|
||||||
|
transcribedText := state.TranscribedText
|
||||||
|
tokenUsage := state.TokenUsage
|
||||||
|
requestID := state.RequestID
|
||||||
|
modelName := defaultModel
|
||||||
|
state.mu.Unlock()
|
||||||
|
|
||||||
|
// 发送 llm_done
|
||||||
|
if sender != nil && requestID != "" {
|
||||||
|
done := models.WsLLMDone{
|
||||||
|
Type: "llm_done",
|
||||||
|
RequestID: requestID,
|
||||||
|
FullText: fullResponse,
|
||||||
|
Model: modelName,
|
||||||
|
LatencyMs: latencyFromCtx(ctx),
|
||||||
|
}
|
||||||
|
if tokenUsage != nil {
|
||||||
|
done.TokensUsed = struct {
|
||||||
|
Prompt int `json:"prompt"`
|
||||||
|
Completion int `json:"completion"`
|
||||||
|
Total int `json:"total"`
|
||||||
|
}{
|
||||||
|
Prompt: tokenUsage.Prompt,
|
||||||
|
Completion: tokenUsage.Completion,
|
||||||
|
Total: tokenUsage.Total,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := sender.SendLLMDone(done); err != nil {
|
||||||
|
log.Errorw("send llm_done failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infow("query processing completed", "response_length", len(fullResponse))
|
||||||
|
|
||||||
|
return PipelineOutput{
|
||||||
|
TranscribedText: transcribedText,
|
||||||
|
FullResponse: fullResponse,
|
||||||
|
Model: modelName,
|
||||||
|
TokenUsage: tokenUsage,
|
||||||
|
}, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
151
backend/internal/eino/nodes_history.go
Normal file
151
backend/internal/eino/nodes_history.go
Normal file
@@ -0,0 +1,151 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/llm"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/store"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewHistoryLambda 创建历史组装 Lambda 节点。
|
||||||
|
// 输入: *STTOutput → 输出: []*schema.Message
|
||||||
|
//
|
||||||
|
// 从 PipelineState 读取请求元数据(SessionID、Scenario、ImageData 等),
|
||||||
|
// 构建系统提示词,组装历史消息和当前用户输入(含多模态图片)。
|
||||||
|
func NewHistoryLambda(
|
||||||
|
historyFetcher func(ctx context.Context, sessionID string, limit int) ([]models.Message, error),
|
||||||
|
scenarioRepo store.UserScenarioRepository,
|
||||||
|
maxHistory int,
|
||||||
|
) *compose.Lambda {
|
||||||
|
return compose.InvokableLambda(func(ctx context.Context, sttOut STTOutput) ([]*schema.Message, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
// 从 State 读取请求元数据
|
||||||
|
state := stateFromCtx(ctx)
|
||||||
|
if state == nil {
|
||||||
|
return []*schema.Message{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
state.mu.Lock()
|
||||||
|
sessionID := state.SessionID
|
||||||
|
requestID := state.RequestID
|
||||||
|
imageData := state.ImageData
|
||||||
|
scenario := state.Scenario
|
||||||
|
detailLevel := state.DetailLevel
|
||||||
|
language := sttOut.Language
|
||||||
|
userID := state.UserID
|
||||||
|
state.mu.Unlock()
|
||||||
|
|
||||||
|
// 加载用户自建情景(如果有 userID 和 scenarioRepo)
|
||||||
|
var customScenarios map[string]string
|
||||||
|
var customGreetings map[string]string
|
||||||
|
if userID != "" && scenarioRepo != nil {
|
||||||
|
scenarios, err := scenarioRepo.FindByUserID(ctx, userID)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnw("load user scenarios failed", "user_id", userID, "error", err)
|
||||||
|
} else if len(scenarios) > 0 {
|
||||||
|
customScenarios = make(map[string]string, len(scenarios))
|
||||||
|
customGreetings = make(map[string]string, len(scenarios))
|
||||||
|
for _, s := range scenarios {
|
||||||
|
customScenarios[s.ID] = s.Prompt
|
||||||
|
if s.Greeting != "" {
|
||||||
|
customGreetings[s.ID] = s.Greeting
|
||||||
|
}
|
||||||
|
}
|
||||||
|
log.Debugw("loaded user scenarios", "user_id", userID, "count", len(scenarios))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 构建系统提示词(支持用户自建情景)
|
||||||
|
scenarioPrompt := llm.GetScenarioPrompt(scenario, language, customScenarios)
|
||||||
|
systemPrompt := llm.BuildSystemPrompt(language, detailLevel, scenarioPrompt)
|
||||||
|
|
||||||
|
// 构建 system message(仅文本,多模态内容只能放在 user 角色)
|
||||||
|
systemMsg := &schema.Message{
|
||||||
|
Role: schema.System,
|
||||||
|
Content: systemPrompt,
|
||||||
|
}
|
||||||
|
|
||||||
|
messages := []*schema.Message{systemMsg}
|
||||||
|
|
||||||
|
// 获取并追加历史消息
|
||||||
|
if historyFetcher != nil && sessionID != "" {
|
||||||
|
history, err := historyFetcher(ctx, sessionID, maxHistory)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnw("fetch history failed, continuing", "error", err, "request_id", requestID)
|
||||||
|
} else {
|
||||||
|
for _, msg := range history {
|
||||||
|
messages = append(messages, &schema.Message{
|
||||||
|
Role: schema.RoleType(msg.Role),
|
||||||
|
Content: msg.Content,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 追加当前用户输入(含图片,多模态内容只能放在 user 角色)
|
||||||
|
// 注意:不能同时设置 Content 和 UserInputMultiContent,需要统一放到 MultiContent 中
|
||||||
|
if len(imageData) > 0 {
|
||||||
|
base64Str := base64.StdEncoding.EncodeToString(imageData)
|
||||||
|
mimeType := detectImageMimeType(imageData)
|
||||||
|
parts := []schema.MessageInputPart{
|
||||||
|
{
|
||||||
|
Type: schema.ChatMessagePartTypeText,
|
||||||
|
Text: sttOut.Text,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: schema.ChatMessagePartTypeImageURL,
|
||||||
|
Image: &schema.MessageInputImage{
|
||||||
|
MessagePartCommon: schema.MessagePartCommon{
|
||||||
|
Base64Data: &base64Str,
|
||||||
|
MIMEType: mimeType,
|
||||||
|
},
|
||||||
|
Detail: schema.ImageURLDetailAuto,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
messages = append(messages, &schema.Message{
|
||||||
|
Role: schema.User,
|
||||||
|
UserInputMultiContent: parts,
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
messages = append(messages, &schema.Message{
|
||||||
|
Role: schema.User,
|
||||||
|
Content: sttOut.Text,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("history assembled",
|
||||||
|
"message_count", len(messages),
|
||||||
|
"has_image", len(imageData) > 0,
|
||||||
|
"scenario", scenario)
|
||||||
|
|
||||||
|
return messages, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// detectImageMimeType 简单检测图片 MIME 类型。
|
||||||
|
func detectImageMimeType(data []byte) string {
|
||||||
|
if len(data) < 4 {
|
||||||
|
return "image/jpeg"
|
||||||
|
}
|
||||||
|
if data[0] == 0xFF && data[1] == 0xD8 && data[2] == 0xFF {
|
||||||
|
return "image/jpeg"
|
||||||
|
}
|
||||||
|
if data[0] == 0x89 && data[1] == 0x50 && data[2] == 0x4E && data[3] == 0x47 {
|
||||||
|
return "image/png"
|
||||||
|
}
|
||||||
|
if data[0] == 0x47 && data[1] == 0x49 && data[2] == 0x46 {
|
||||||
|
return "image/gif"
|
||||||
|
}
|
||||||
|
if data[0] == 0x52 && data[1] == 0x49 && data[2] == 0x46 && data[3] == 0x46 {
|
||||||
|
return "image/webp"
|
||||||
|
}
|
||||||
|
return "image/jpeg"
|
||||||
|
}
|
||||||
102
backend/internal/eino/nodes_splitter.go
Normal file
102
backend/internal/eino/nodes_splitter.go
Normal file
@@ -0,0 +1,102 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
)
|
||||||
|
|
||||||
|
// sentenceDelimiters 句子分隔符集合。
|
||||||
|
var sentenceDelimiters = map[rune]bool{
|
||||||
|
'。': true,
|
||||||
|
'!': true,
|
||||||
|
'?': true,
|
||||||
|
'\n': true,
|
||||||
|
'.': true,
|
||||||
|
'!': true,
|
||||||
|
'?': true,
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMessageToStringLambda 创建 Message → String 转换 Lambda 节点。
|
||||||
|
// 输入: *schema.Message → 输出: string
|
||||||
|
//
|
||||||
|
// 提取 Message.Content 文本,供 Splitter 节点消费。
|
||||||
|
func NewMessageToStringLambda() *compose.Lambda {
|
||||||
|
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[*schema.Message]) (*schema.StreamReader[string], error) {
|
||||||
|
sr, sw := schema.Pipe[string](8)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer sw.Close()
|
||||||
|
defer input.Close()
|
||||||
|
|
||||||
|
for {
|
||||||
|
msg, err := input.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sw.Send("", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if msg != nil && msg.Content != "" {
|
||||||
|
sw.Send(msg.Content, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return sr, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSplitterLambda 创建句子分割 Transform Lambda 节点。
|
||||||
|
// 输入: StreamReader[string](LLM token 流)→ 输出: StreamReader[string](完整句子流)
|
||||||
|
//
|
||||||
|
// 逐字符累积,按句子分隔符切分。每切出一个完整句子就输出一次,
|
||||||
|
// 供下游 TTS 节点实时合成。
|
||||||
|
func NewSplitterLambda() *compose.Lambda {
|
||||||
|
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[string], error) {
|
||||||
|
sr, sw := schema.Pipe[string](8)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer sw.Close()
|
||||||
|
defer input.Close()
|
||||||
|
|
||||||
|
var buffer strings.Builder
|
||||||
|
|
||||||
|
for {
|
||||||
|
chunk, err := input.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
// 流结束,flush 剩余缓冲
|
||||||
|
if buffer.Len() > 0 {
|
||||||
|
text := strings.TrimSpace(buffer.String())
|
||||||
|
if text != "" {
|
||||||
|
sw.Send(text, nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sw.Send("", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 逐字符累积,按句子分隔符切分
|
||||||
|
for _, r := range chunk {
|
||||||
|
buffer.WriteRune(r)
|
||||||
|
if sentenceDelimiters[r] {
|
||||||
|
text := strings.TrimSpace(buffer.String())
|
||||||
|
if text != "" {
|
||||||
|
sw.Send(text, nil)
|
||||||
|
}
|
||||||
|
buffer.Reset()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return sr, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
135
backend/internal/eino/nodes_stt.go
Normal file
135
backend/internal/eino/nodes_stt.go
Normal file
@@ -0,0 +1,135 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/stt"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
"github.com/hhs/camtalk/internal/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewSTTLambda 创建 STT Lambda 节点。
|
||||||
|
// 输入: PipelineInput → 输出: STTOutput
|
||||||
|
//
|
||||||
|
// 文本输入模式:跳过 STT,直接返回用户输入文本。
|
||||||
|
// 语音模式:调用 sttService.Recognize() 进行语音识别。
|
||||||
|
// 识别结果通过 Sender 发送 stt_result 到客户端。
|
||||||
|
func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||||
|
return compose.InvokableLambda(func(ctx context.Context, input PipelineInput) (STTOutput, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
sender := senderFromCtx(ctx)
|
||||||
|
requestID := requestIDFromCtx(ctx)
|
||||||
|
|
||||||
|
// 将输入元数据写入 State,供下游节点(History、Done)读取
|
||||||
|
if state := stateFromCtx(ctx); state != nil {
|
||||||
|
state.mu.Lock()
|
||||||
|
state.SessionID = input.SessionID
|
||||||
|
state.RequestID = input.RequestID
|
||||||
|
state.ImageData = input.ImageData
|
||||||
|
state.Scenario = input.Scenario
|
||||||
|
state.DetailLevel = "low"
|
||||||
|
state.Language = input.Language
|
||||||
|
state.TTSEnabled = input.TTSEnabled
|
||||||
|
state.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// 文本输入模式:跳过 STT
|
||||||
|
if input.Text != "" {
|
||||||
|
log.Debugw("text input mode, skipping stt",
|
||||||
|
"text_len", len(input.Text),
|
||||||
|
"text_preview", util.Truncate(input.Text, 50))
|
||||||
|
|
||||||
|
// 发送 stt_result 保持前端消息流一致性
|
||||||
|
if sender != nil {
|
||||||
|
if err := sender.SendSTTResult(models.WsSTTResult{
|
||||||
|
Type: "stt_result",
|
||||||
|
RequestID: requestID,
|
||||||
|
Text: input.Text,
|
||||||
|
IsFinal: true,
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("send stt_result failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 写入 State
|
||||||
|
if state := stateFromCtx(ctx); state != nil {
|
||||||
|
state.mu.Lock()
|
||||||
|
state.TranscribedText = input.Text
|
||||||
|
state.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
return STTOutput{
|
||||||
|
Text: input.Text,
|
||||||
|
Language: input.Language,
|
||||||
|
IsSkipped: true,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 语音模式:解码音频
|
||||||
|
if len(input.AudioData) == 0 {
|
||||||
|
return STTOutput{}, fmt.Errorf("stt: no audio data provided")
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("stt recognition started", "audio_bytes", len(input.AudioData))
|
||||||
|
|
||||||
|
// 调用 STT 服务
|
||||||
|
text, err := sttService.Recognize(ctx, input.AudioData, stt.Options{
|
||||||
|
Encoding: "pcm_s16le",
|
||||||
|
SampleRate: 16000,
|
||||||
|
Language: input.Language,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("stt recognition failed", "error", err)
|
||||||
|
if sender != nil {
|
||||||
|
sender.SendError(models.WsError{
|
||||||
|
Type: "error",
|
||||||
|
RequestID: requestID,
|
||||||
|
Code: "STT_ERROR",
|
||||||
|
Message: "语音识别失败: " + err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return STTOutput{}, fmt.Errorf("stt: recognize: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// STT 返回空文本
|
||||||
|
if strings.TrimSpace(text) == "" {
|
||||||
|
log.Infow("stt returned empty text")
|
||||||
|
text = "(未识别到语音)"
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("stt recognition completed",
|
||||||
|
"text_len", len(text),
|
||||||
|
"text_preview", util.Truncate(text, 50))
|
||||||
|
|
||||||
|
// 发送 stt_result
|
||||||
|
if sender != nil {
|
||||||
|
if err := sender.SendSTTResult(models.WsSTTResult{
|
||||||
|
Type: "stt_result",
|
||||||
|
RequestID: requestID,
|
||||||
|
Text: text,
|
||||||
|
IsFinal: true,
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("send stt_result failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 写入 State
|
||||||
|
if state := stateFromCtx(ctx); state != nil {
|
||||||
|
state.mu.Lock()
|
||||||
|
state.TranscribedText = text
|
||||||
|
state.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
return STTOutput{
|
||||||
|
Text: text,
|
||||||
|
Language: input.Language,
|
||||||
|
IsSkipped: false,
|
||||||
|
}, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
116
backend/internal/eino/nodes_tts.go
Normal file
116
backend/internal/eino/nodes_tts.go
Normal file
@@ -0,0 +1,116 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/tts"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewTTSLambda 创建 TTS Transform Lambda 节点。
|
||||||
|
// 输入: StreamReader[string](句子流)→ 输出: StreamReader[struct{}](结果流)
|
||||||
|
//
|
||||||
|
// 流式消费每个句子,调用 ttsService.SynthesizeStream() 合成,
|
||||||
|
// 逐 chunk 推送 tts_audio 到客户端。TTS 失败静默跳过。
|
||||||
|
func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, ttsOutputFmt string, ttsSampleRate int) *compose.Lambda {
|
||||||
|
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[struct{}], error) {
|
||||||
|
sr, sw := schema.Pipe[struct{}](8)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer sw.Close()
|
||||||
|
defer input.Close()
|
||||||
|
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
sender := senderFromCtx(ctx)
|
||||||
|
requestID := requestIDFromCtx(ctx)
|
||||||
|
|
||||||
|
if sender == nil || requestID == "" {
|
||||||
|
// 消费并丢弃流
|
||||||
|
for {
|
||||||
|
_, err := input.Recv()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 收集句子,按批次合成 TTS
|
||||||
|
var sentences []string
|
||||||
|
for {
|
||||||
|
sentence, err := input.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
log.Errorw("TTS: stream recv error", "error", err)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if sentence != "" {
|
||||||
|
sentences = append(sentences, sentence)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(sentences) == 0 {
|
||||||
|
sw.Send(struct{}{}, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infow("开始 TTS 合成", "sentence_count", len(sentences))
|
||||||
|
|
||||||
|
// 将句子数组转为 channel
|
||||||
|
sentenceCh := make(chan string, len(sentences))
|
||||||
|
for _, s := range sentences {
|
||||||
|
sentenceCh <- s
|
||||||
|
}
|
||||||
|
close(sentenceCh)
|
||||||
|
|
||||||
|
// 调用 TTS 服务
|
||||||
|
ttsStream, err := ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{
|
||||||
|
Voice: ttsVoice,
|
||||||
|
Speed: ttsSpeed,
|
||||||
|
OutputFmt: ttsOutputFmt,
|
||||||
|
SampleRate: ttsSampleRate,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("TTS 合成启动失败(已跳过)", "error", err)
|
||||||
|
sw.Send(struct{}{}, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 消费 TTS 音频流,推送到客户端
|
||||||
|
for chunk := range ttsStream {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
log.Debugw("tts stream interrupted")
|
||||||
|
sw.Send(struct{}{}, ctx.Err())
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
audioBase64 := base64.StdEncoding.EncodeToString(chunk.Audio)
|
||||||
|
|
||||||
|
if err := sender.SendTTSAudio(models.WsTTSAudio{
|
||||||
|
Type: "tts_audio",
|
||||||
|
RequestID: requestID,
|
||||||
|
Audio: audioBase64,
|
||||||
|
MimeType: "audio/mp3",
|
||||||
|
IsLast: chunk.IsLast,
|
||||||
|
Final: chunk.Final,
|
||||||
|
}); err != nil {
|
||||||
|
log.Errorw("发送 tts_audio 失败", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infow("TTS 合成完成")
|
||||||
|
sw.Send(struct{}{}, nil)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return sr, nil
|
||||||
|
})
|
||||||
|
}
|
||||||
46
backend/internal/eino/state.go
Normal file
46
backend/internal/eino/state.go
Normal file
@@ -0,0 +1,46 @@
|
|||||||
|
package eino
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PipelineState Graph 全局状态,用于跨节点收集数据。
|
||||||
|
// 通过 compose.WithGenLocalState 注册,各节点通过 compose.ProcessState 读写。
|
||||||
|
type PipelineState struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
FullResponse strings.Builder // LLM 完整回复(由 Callback 累积)
|
||||||
|
TranscribedText string // STT 识别文本
|
||||||
|
Model string // 实际使用的模型名
|
||||||
|
TokenUsage *TokenUsage // token 用量
|
||||||
|
|
||||||
|
// 从 PipelineInput 复制的元数据,供下游节点(History、Done)读取
|
||||||
|
SessionID string
|
||||||
|
RequestID string
|
||||||
|
ImageData []byte
|
||||||
|
Scenario string
|
||||||
|
DetailLevel string
|
||||||
|
Language string
|
||||||
|
TTSEnabled bool
|
||||||
|
UserID string // 新增:用户 ID,用于加载自建情景
|
||||||
|
}
|
||||||
|
|
||||||
|
// genLocalState 创建每请求的 PipelineState 实例。
|
||||||
|
func genLocalState(ctx context.Context) *PipelineState {
|
||||||
|
return &PipelineState{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// AppendText 追加文本到 FullResponse(线程安全)。
|
||||||
|
func (s *PipelineState) AppendText(text string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.FullResponse.WriteString(text)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetFullResponse 获取完整回复文本(线程安全)。
|
||||||
|
func (s *PipelineState) GetFullResponse() string {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return s.FullResponse.String()
|
||||||
|
}
|
||||||
38
backend/internal/eino/types.go
Normal file
38
backend/internal/eino/types.go
Normal file
@@ -0,0 +1,38 @@
|
|||||||
|
// Package eino 基于 CloudWeGo Eino 框架的 AI 编排层。
|
||||||
|
// 使用 Eino Graph 替代手写 goroutine 管道,实现声明式 STT → LLM → TTS 编排。
|
||||||
|
package eino
|
||||||
|
|
||||||
|
// PipelineInput Graph 统一输入。
|
||||||
|
type PipelineInput struct {
|
||||||
|
AudioData []byte // base64 解码后的音频(可选)
|
||||||
|
ImageData []byte // base64 解码后的图像(可选)
|
||||||
|
Text string // 直接文本输入(可选,跳过 STT)
|
||||||
|
SessionID string
|
||||||
|
RequestID string
|
||||||
|
Language string // zh / en
|
||||||
|
Scenario string // free_chat, interviewer, etc.
|
||||||
|
TTSEnabled bool
|
||||||
|
UserID string // 用户 ID,用于加载自建情景
|
||||||
|
}
|
||||||
|
|
||||||
|
// PipelineOutput Graph 统一输出。
|
||||||
|
type PipelineOutput struct {
|
||||||
|
TranscribedText string // STT 结果
|
||||||
|
FullResponse string // LLM 完整回复
|
||||||
|
Model string // 实际使用的模型名
|
||||||
|
TokenUsage *TokenUsage // token 用量
|
||||||
|
}
|
||||||
|
|
||||||
|
// STTOutput STT 节点输出。
|
||||||
|
type STTOutput struct {
|
||||||
|
Text string
|
||||||
|
Language string
|
||||||
|
IsSkipped bool // 文本输入模式跳过了 STT
|
||||||
|
}
|
||||||
|
|
||||||
|
// TokenUsage token 用量统计。
|
||||||
|
type TokenUsage struct {
|
||||||
|
Prompt int
|
||||||
|
Completion int
|
||||||
|
Total int
|
||||||
|
}
|
||||||
43
backend/internal/models/user_scenario.go
Normal file
43
backend/internal/models/user_scenario.go
Normal file
@@ -0,0 +1,43 @@
|
|||||||
|
package models
|
||||||
|
|
||||||
|
import "time"
|
||||||
|
|
||||||
|
// UserScenario 用户自建情景。
|
||||||
|
type UserScenario struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
UserID string `json:"user_id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Icon string `json:"icon"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Prompt string `json:"prompt"`
|
||||||
|
Greeting string `json:"greeting,omitempty"`
|
||||||
|
Language string `json:"language"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateUserScenarioRequest 创建用户情景请求。
|
||||||
|
type CreateUserScenarioRequest struct {
|
||||||
|
Name string `json:"name" binding:"required,min=2,max=50"`
|
||||||
|
Icon string `json:"icon,omitempty"`
|
||||||
|
Description string `json:"description,omitempty" binding:"omitempty,max=100"`
|
||||||
|
Prompt string `json:"prompt" binding:"required,min=10,max=2000"`
|
||||||
|
Greeting string `json:"greeting,omitempty" binding:"omitempty,max=500"`
|
||||||
|
Language string `json:"language,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateUserScenarioRequest 更新用户情景请求。
|
||||||
|
type UpdateUserScenarioRequest struct {
|
||||||
|
Name *string `json:"name,omitempty" binding:"omitempty,min=2,max=50"`
|
||||||
|
Icon *string `json:"icon,omitempty"`
|
||||||
|
Description *string `json:"description,omitempty" binding:"omitempty,max=100"`
|
||||||
|
Prompt *string `json:"prompt,omitempty" binding:"omitempty,min=10,max=2000"`
|
||||||
|
Greeting *string `json:"greeting,omitempty" binding:"omitempty,max=500"`
|
||||||
|
Language *string `json:"language,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UserScenarioListResponse 用户情景列表响应。
|
||||||
|
type UserScenarioListResponse struct {
|
||||||
|
Scenarios []*UserScenario `json:"scenarios"`
|
||||||
|
Total int `json:"total"`
|
||||||
|
}
|
||||||
@@ -14,13 +14,11 @@ type Orchestrator interface {
|
|||||||
// ctx 用于整体超时和中断控制。
|
// ctx 用于整体超时和中断控制。
|
||||||
// sessionID 用于会话管理和历史获取。
|
// sessionID 用于会话管理和历史获取。
|
||||||
// req 包含图像和音频数据。
|
// req 包含图像和音频数据。
|
||||||
// history 是最近的对话历史。
|
|
||||||
// sender 用于向客户端推送消息。
|
// sender 用于向客户端推送消息。
|
||||||
ProcessQuery(
|
ProcessQuery(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
sessionID string,
|
sessionID string,
|
||||||
req models.WsQuery,
|
req models.WsQuery,
|
||||||
history []models.Message,
|
|
||||||
sender Sender,
|
sender Sender,
|
||||||
) error
|
) error
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,403 +0,0 @@
|
|||||||
package orchestrator
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/base64"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
"unicode/utf8"
|
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/ai/llm"
|
|
||||||
"github.com/hhs/camtalk/internal/ai/stt"
|
|
||||||
"github.com/hhs/camtalk/internal/ai/tts"
|
|
||||||
"github.com/hhs/camtalk/internal/config"
|
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
|
||||||
"github.com/hhs/camtalk/internal/session"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Pipeline 实现 Orchestrator 接口,管理 STT → LLM → TTS 流式管道。
|
|
||||||
type Pipeline struct {
|
|
||||||
sttService stt.Service
|
|
||||||
llmService llm.Service
|
|
||||||
ttsService tts.Service
|
|
||||||
sessionMgr session.Manager
|
|
||||||
model string // LLM 模型名,用于 llm_done 上报
|
|
||||||
ttsVoice string // TTS 音色
|
|
||||||
ttsSpeed float64 // TTS 语速
|
|
||||||
ttsOutputFmt string // TTS 输出格式
|
|
||||||
ttsSampleRate int // TTS 输出采样率
|
|
||||||
}
|
|
||||||
|
|
||||||
// New 创建 Pipeline 实例。
|
|
||||||
func New(
|
|
||||||
sttService stt.Service,
|
|
||||||
llmService llm.Service,
|
|
||||||
ttsService tts.Service,
|
|
||||||
sessionMgr session.Manager,
|
|
||||||
cfg *config.Config,
|
|
||||||
) *Pipeline {
|
|
||||||
return &Pipeline{
|
|
||||||
sttService: sttService,
|
|
||||||
llmService: llmService,
|
|
||||||
ttsService: ttsService,
|
|
||||||
sessionMgr: sessionMgr,
|
|
||||||
model: cfg.AI.LLM.Model,
|
|
||||||
ttsVoice: cfg.AI.TTS.Voice,
|
|
||||||
ttsSpeed: cfg.AI.TTS.Speed,
|
|
||||||
ttsOutputFmt: cfg.AI.TTS.OutputFormat,
|
|
||||||
ttsSampleRate: cfg.AI.TTS.SampleRate,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ProcessQuery 实现 Orchestrator 接口。
|
|
||||||
func (p *Pipeline) ProcessQuery(
|
|
||||||
ctx context.Context,
|
|
||||||
sessionID string,
|
|
||||||
req models.WsQuery,
|
|
||||||
history []models.Message,
|
|
||||||
sender Sender,
|
|
||||||
) error {
|
|
||||||
log := logger.Log
|
|
||||||
startTime := time.Now()
|
|
||||||
|
|
||||||
// 解码音频数据(文本输入模式可跳过)
|
|
||||||
var audio []byte
|
|
||||||
if req.Text == "" && req.Audio != "" {
|
|
||||||
var err error
|
|
||||||
audio, err = base64.StdEncoding.DecodeString(req.Audio)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("音频解码失败", "error", err)
|
|
||||||
sender.SendError(models.WsError{
|
|
||||||
Type: "error",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Code: "INVALID_MESSAGE",
|
|
||||||
Message: "音频数据解码失败",
|
|
||||||
})
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 解码图片数据(可选)
|
|
||||||
var image []byte
|
|
||||||
if req.Image != "" {
|
|
||||||
var err error
|
|
||||||
image, err = base64.StdEncoding.DecodeString(req.Image)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("图片解码失败", "error", err)
|
|
||||||
sender.SendError(models.WsError{
|
|
||||||
Type: "error",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Code: "INVALID_MESSAGE",
|
|
||||||
Message: "图片数据解码失败",
|
|
||||||
})
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 设置活跃请求
|
|
||||||
if err := p.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil {
|
|
||||||
log.Errorw("设置活跃请求失败", "error", err)
|
|
||||||
}
|
|
||||||
defer p.sessionMgr.ClearActiveRequest(ctx, sessionID)
|
|
||||||
|
|
||||||
// 获取会话配置
|
|
||||||
sess, err := p.sessionMgr.Get(ctx, sessionID)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("获取会话失败", "error", err)
|
|
||||||
sender.SendError(models.WsError{
|
|
||||||
Type: "error",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Code: "SESSION_NOT_FOUND",
|
|
||||||
Message: "会话不存在",
|
|
||||||
})
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Step 1: 获取用户文本(语音识别或直接使用输入文本)
|
|
||||||
var userText string
|
|
||||||
if req.Text != "" {
|
|
||||||
// 文本输入模式:跳过 STT,直接使用用户输入的文本
|
|
||||||
log.Infow("使用文本输入", "request_id", req.RequestID, "text", req.Text)
|
|
||||||
userText = req.Text
|
|
||||||
|
|
||||||
// 发送 stt_result 以保持前端消息流一致性
|
|
||||||
if err := sender.SendSTTResult(models.WsSTTResult{
|
|
||||||
Type: "stt_result",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Text: userText,
|
|
||||||
IsFinal: true,
|
|
||||||
}); err != nil {
|
|
||||||
log.Errorw("发送 STT 结果失败", "error", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
// 语音模式:执行 STT 语音识别
|
|
||||||
log.Infow("开始语音识别", "request_id", req.RequestID, "audio_bytes", len(audio))
|
|
||||||
sttResult, err := p.sttService.Recognize(ctx, audio, stt.Options{
|
|
||||||
Encoding: "pcm_s16le",
|
|
||||||
SampleRate: 16000,
|
|
||||||
Language: sess.Config.Language,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("语音识别失败", "error", err, "audio_bytes", len(audio))
|
|
||||||
sender.SendError(models.WsError{
|
|
||||||
Type: "error",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Code: "STT_ERROR",
|
|
||||||
Message: "语音识别失败: " + err.Error(),
|
|
||||||
})
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
userText = sttResult
|
|
||||||
|
|
||||||
// STT 返回空文本:未识别到语音,发送结果后直接返回(不调 LLM)
|
|
||||||
if strings.TrimSpace(userText) == "" {
|
|
||||||
log.Infow("语音识别结果为空", "request_id", req.RequestID)
|
|
||||||
userText = "(未识别到语音)"
|
|
||||||
if err := sender.SendSTTResult(models.WsSTTResult{
|
|
||||||
Type: "stt_result",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Text: userText,
|
|
||||||
IsFinal: true,
|
|
||||||
}); err != nil {
|
|
||||||
log.Errorw("发送 STT 结果失败", "error", err)
|
|
||||||
}
|
|
||||||
// 发送空的 llm_done 以结束本轮处理
|
|
||||||
latency := time.Since(startTime).Milliseconds()
|
|
||||||
_ = sender.SendLLMDone(models.WsLLMDone{
|
|
||||||
Type: "llm_done",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
FullText: "",
|
|
||||||
Model: p.model,
|
|
||||||
LatencyMs: latency,
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// 发送 STT 结果
|
|
||||||
if err := sender.SendSTTResult(models.WsSTTResult{
|
|
||||||
Type: "stt_result",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Text: userText,
|
|
||||||
IsFinal: true,
|
|
||||||
}); err != nil {
|
|
||||||
log.Errorw("发送 STT 结果失败", "error", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 追加用户消息到历史
|
|
||||||
p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
|
||||||
Role: "user",
|
|
||||||
Content: userText,
|
|
||||||
})
|
|
||||||
|
|
||||||
// Step 2+3: LLM 流式推理 + TTS 并行合成
|
|
||||||
log.Infow("开始 LLM 推理", "request_id", req.RequestID, "scenario", sess.Config.Scenario)
|
|
||||||
llmReq := llm.Request{
|
|
||||||
Image: image,
|
|
||||||
Text: userText,
|
|
||||||
History: history,
|
|
||||||
Language: sess.Config.Language,
|
|
||||||
SystemPrompt: llm.GetScenarioPrompt(sess.Config.Scenario, sess.Config.Language),
|
|
||||||
}
|
|
||||||
|
|
||||||
llmStream, err := p.llmService.ChatStream(ctx, llmReq)
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("LLM 流式推理启动失败", "error", err)
|
|
||||||
sender.SendError(models.WsError{
|
|
||||||
Type: "error",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
Code: "LLM_ERROR",
|
|
||||||
Message: "LLM 推理失败",
|
|
||||||
})
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// 创建句子切分器
|
|
||||||
sentenceCh := make(chan string, 4)
|
|
||||||
splitter := NewSplitter(sentenceCh)
|
|
||||||
|
|
||||||
// 并行:LLM 消费 + TTS 合成
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
var fullText string
|
|
||||||
var ttsErr error
|
|
||||||
|
|
||||||
// goroutine 1: 消费 LLM token + 句子切分
|
|
||||||
var tokenUsage *llm.TokenUsage
|
|
||||||
wg.Add(1)
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
defer close(sentenceCh)
|
|
||||||
fullText, tokenUsage = p.consumeLLMStream(ctx, llmStream, req.RequestID, sender, splitter)
|
|
||||||
}()
|
|
||||||
|
|
||||||
// goroutine 2: TTS 合成(如果启用)
|
|
||||||
if sess.Config.TTSEnabled {
|
|
||||||
wg.Add(1)
|
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
|
||||||
log.Infow("开始 TTS 合成", "request_id", req.RequestID)
|
|
||||||
ttsErr = p.synthesizeTTS(ctx, sentenceCh, req.RequestID, sender)
|
|
||||||
}()
|
|
||||||
} else {
|
|
||||||
// 如果 TTS 未启用,需要消费 sentenceCh 防止阻塞
|
|
||||||
go func() {
|
|
||||||
for range sentenceCh {
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
}
|
|
||||||
|
|
||||||
// 等待所有 goroutine 完成
|
|
||||||
wg.Wait()
|
|
||||||
|
|
||||||
// TTS 失败静默跳过
|
|
||||||
if ttsErr != nil {
|
|
||||||
log.Warnw("TTS 合成失败(已跳过)", "error", ttsErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 追加助手消息到历史
|
|
||||||
p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
|
||||||
Role: "assistant",
|
|
||||||
Content: fullText,
|
|
||||||
})
|
|
||||||
|
|
||||||
// 发送 llm_done
|
|
||||||
latency := time.Since(startTime).Milliseconds()
|
|
||||||
done := models.WsLLMDone{
|
|
||||||
Type: "llm_done",
|
|
||||||
RequestID: req.RequestID,
|
|
||||||
FullText: fullText,
|
|
||||||
Model: p.model,
|
|
||||||
LatencyMs: latency,
|
|
||||||
}
|
|
||||||
if tokenUsage != nil {
|
|
||||||
done.TokensUsed = struct {
|
|
||||||
Prompt int `json:"prompt"`
|
|
||||||
Completion int `json:"completion"`
|
|
||||||
Total int `json:"total"`
|
|
||||||
}{
|
|
||||||
Prompt: tokenUsage.Prompt,
|
|
||||||
Completion: tokenUsage.Completion,
|
|
||||||
Total: tokenUsage.Total,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := sender.SendLLMDone(done); err != nil {
|
|
||||||
log.Errorw("发送 llm_done 失败", "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Infow("查询处理完成",
|
|
||||||
"request_id", req.RequestID,
|
|
||||||
"latency_ms", latency,
|
|
||||||
"text_length", utf8.RuneCountInString(fullText),
|
|
||||||
)
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// consumeLLMStream 消费 LLM 流式输出,发送 llm_chunk 并进行句子切分。
|
|
||||||
// 返回完整文本和 token 用量。
|
|
||||||
func (p *Pipeline) consumeLLMStream(
|
|
||||||
ctx context.Context,
|
|
||||||
stream <-chan llm.Chunk,
|
|
||||||
requestID string,
|
|
||||||
sender Sender,
|
|
||||||
splitter *Splitter,
|
|
||||||
) (string, *llm.TokenUsage) {
|
|
||||||
log := logger.Log
|
|
||||||
var fullText strings.Builder
|
|
||||||
var tokenUsage *llm.TokenUsage
|
|
||||||
|
|
||||||
for chunk := range stream {
|
|
||||||
// 检查上下文是否已取消
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
log.Infow("LLM 流被中断", "request_id", requestID)
|
|
||||||
return fullText.String(), tokenUsage
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
if chunk.Done {
|
|
||||||
// 流结束,记录 token 用量
|
|
||||||
if chunk.TokensUsed != nil {
|
|
||||||
tokenUsage = chunk.TokensUsed
|
|
||||||
log.Infow("LLM 用量统计",
|
|
||||||
"request_id", requestID,
|
|
||||||
"prompt_tokens", tokenUsage.Prompt,
|
|
||||||
"completion_tokens", tokenUsage.Completion,
|
|
||||||
"total_tokens", tokenUsage.Total,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
|
||||||
|
|
||||||
// 累积全文
|
|
||||||
fullText.WriteString(chunk.Delta)
|
|
||||||
|
|
||||||
// 发送 llm_chunk
|
|
||||||
if err := sender.SendLLMChunk(models.WsLLMChunk{
|
|
||||||
Type: "llm_chunk",
|
|
||||||
RequestID: requestID,
|
|
||||||
Delta: chunk.Delta,
|
|
||||||
Role: "assistant",
|
|
||||||
}); err != nil {
|
|
||||||
log.Errorw("发送 llm_chunk 失败", "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 句子切分
|
|
||||||
splitter.Feed(chunk.Delta)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 刷新切分器中的剩余文本
|
|
||||||
splitter.Flush()
|
|
||||||
|
|
||||||
return fullText.String(), tokenUsage
|
|
||||||
}
|
|
||||||
|
|
||||||
// synthesizeTTS 从句子 channel 读取文本,进行 TTS 合成并发送音频。
|
|
||||||
func (p *Pipeline) synthesizeTTS(
|
|
||||||
ctx context.Context,
|
|
||||||
sentenceCh <-chan string,
|
|
||||||
requestID string,
|
|
||||||
sender Sender,
|
|
||||||
) error {
|
|
||||||
log := logger.Log
|
|
||||||
|
|
||||||
ttsStream, err := p.ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{
|
|
||||||
Voice: p.ttsVoice,
|
|
||||||
Speed: p.ttsSpeed,
|
|
||||||
OutputFmt: p.ttsOutputFmt,
|
|
||||||
SampleRate: p.ttsSampleRate,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
log.Errorw("TTS 合成启动失败", "error", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// 消费 TTS 音频流
|
|
||||||
for chunk := range ttsStream {
|
|
||||||
// 检查上下文是否已取消
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
log.Infow("TTS 流被中断", "request_id", requestID)
|
|
||||||
return ctx.Err()
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
// Base64 编码音频数据
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(chunk.Audio)
|
|
||||||
|
|
||||||
if err := sender.SendTTSAudio(models.WsTTSAudio{
|
|
||||||
Type: "tts_audio",
|
|
||||||
RequestID: requestID,
|
|
||||||
Audio: audioBase64,
|
|
||||||
MimeType: "audio/mp3",
|
|
||||||
IsLast: chunk.IsLast,
|
|
||||||
Final: chunk.Final,
|
|
||||||
}); err != nil {
|
|
||||||
log.Errorw("发送 tts_audio 失败", "error", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,713 +0,0 @@
|
|||||||
package orchestrator
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/base64"
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/mock"
|
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/ai/llm"
|
|
||||||
"github.com/hhs/camtalk/internal/ai/stt"
|
|
||||||
"github.com/hhs/camtalk/internal/ai/tts"
|
|
||||||
"github.com/hhs/camtalk/internal/config"
|
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
|
||||||
"github.com/hhs/camtalk/internal/session"
|
|
||||||
)
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
logger.Init("debug", "console")
|
|
||||||
}
|
|
||||||
|
|
||||||
// MockSTTService mock STT 服务
|
|
||||||
type MockSTTService struct {
|
|
||||||
mock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSTTService) Recognize(ctx context.Context, audio []byte, opts stt.Options) (string, error) {
|
|
||||||
args := m.Called(ctx, audio, opts)
|
|
||||||
return args.String(0), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// MockLLMService mock LLM 服务
|
|
||||||
type MockLLMService struct {
|
|
||||||
mock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockLLMService) ChatStream(ctx context.Context, req llm.Request) (<-chan llm.Chunk, error) {
|
|
||||||
args := m.Called(ctx, req)
|
|
||||||
if args.Get(0) == nil {
|
|
||||||
return nil, args.Error(1)
|
|
||||||
}
|
|
||||||
return args.Get(0).(<-chan llm.Chunk), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// MockTTSService mock TTS 服务
|
|
||||||
type MockTTSService struct {
|
|
||||||
mock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockTTSService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts tts.Options) (<-chan tts.Chunk, error) {
|
|
||||||
args := m.Called(ctx, textStream, opts)
|
|
||||||
if args.Get(0) == nil {
|
|
||||||
return nil, args.Error(1)
|
|
||||||
}
|
|
||||||
return args.Get(0).(<-chan tts.Chunk), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// MockSessionManager mock 会话管理器
|
|
||||||
type MockSessionManager struct {
|
|
||||||
mock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
|
||||||
args := m.Called(ctx, userID, config)
|
|
||||||
return args.String(0), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
|
||||||
args := m.Called(ctx, sessionID, title)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) ListByUser(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) {
|
|
||||||
args := m.Called(ctx, userID, page, size)
|
|
||||||
return args.Get(0).([]session.ConversationSummary), args.Int(1), args.Error(2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
|
||||||
args := m.Called(ctx, sessionID)
|
|
||||||
if args.Get(0) == nil {
|
|
||||||
return nil, args.Error(1)
|
|
||||||
}
|
|
||||||
return args.Get(0).(*models.Session), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
|
||||||
args := m.Called(ctx, sessionID, patch)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
|
||||||
args := m.Called(ctx, sessionID, limit)
|
|
||||||
return args.Get(0).([]models.Message), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
|
|
||||||
args := m.Called(ctx, sessionID, msg)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
|
|
||||||
args := m.Called(ctx, sessionID, requestID)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
|
|
||||||
args := m.Called(ctx, sessionID)
|
|
||||||
return args.String(0), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
|
|
||||||
args := m.Called(ctx, sessionID)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) Touch(ctx context.Context, sessionID string) error {
|
|
||||||
args := m.Called(ctx, sessionID)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) Destroy(ctx context.Context, sessionID string) error {
|
|
||||||
args := m.Called(ctx, sessionID)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSessionManager) ActiveCount() int {
|
|
||||||
args := m.Called()
|
|
||||||
return args.Int(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// MockSender mock WebSocket 发送器
|
|
||||||
type MockSender struct {
|
|
||||||
mock.Mock
|
|
||||||
STTResults []models.WsSTTResult
|
|
||||||
LLMChunks []models.WsLLMChunk
|
|
||||||
LLMDones []models.WsLLMDone
|
|
||||||
TTSAudios []models.WsTTSAudio
|
|
||||||
Errors []models.WsError
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewMockSender() *MockSender {
|
|
||||||
return &MockSender{
|
|
||||||
STTResults: make([]models.WsSTTResult, 0),
|
|
||||||
LLMChunks: make([]models.WsLLMChunk, 0),
|
|
||||||
LLMDones: make([]models.WsLLMDone, 0),
|
|
||||||
TTSAudios: make([]models.WsTTSAudio, 0),
|
|
||||||
Errors: make([]models.WsError, 0),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSender) SendSTTResult(result models.WsSTTResult) error {
|
|
||||||
m.STTResults = append(m.STTResults, result)
|
|
||||||
args := m.Called(result)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSender) SendLLMChunk(chunk models.WsLLMChunk) error {
|
|
||||||
m.LLMChunks = append(m.LLMChunks, chunk)
|
|
||||||
args := m.Called(chunk)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSender) SendLLMDone(done models.WsLLMDone) error {
|
|
||||||
m.LLMDones = append(m.LLMDones, done)
|
|
||||||
args := m.Called(done)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSender) SendTTSAudio(audio models.WsTTSAudio) error {
|
|
||||||
m.TTSAudios = append(m.TTSAudios, audio)
|
|
||||||
args := m.Called(audio)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *MockSender) SendError(err models.WsError) error {
|
|
||||||
m.Errors = append(m.Errors, err)
|
|
||||||
args := m.Called(err)
|
|
||||||
return args.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 辅助函数:创建 LLM 流式响应
|
|
||||||
func createLLMStream(chunks []llm.Chunk) <-chan llm.Chunk {
|
|
||||||
ch := make(chan llm.Chunk, len(chunks))
|
|
||||||
for _, chunk := range chunks {
|
|
||||||
ch <- chunk
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch
|
|
||||||
}
|
|
||||||
|
|
||||||
// 辅助函数:创建 TTS 流式响应
|
|
||||||
func createTTSStream(chunks []tts.Chunk) <-chan tts.Chunk {
|
|
||||||
ch := make(chan tts.Chunk, len(chunks))
|
|
||||||
for _, chunk := range chunks {
|
|
||||||
ch <- chunk
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_Success 测试完整流程
|
|
||||||
func TestProcessQuery_Success(t *testing.T) {
|
|
||||||
// 准备测试数据
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
imageData := []byte("test image")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
imageBase64 := base64.StdEncoding.EncodeToString(imageData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Image: imageBase64,
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
TTSEnabled: true,
|
|
||||||
DetailLevel: "low",
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 创建 mock
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
// 设置 mock 期望
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
|
||||||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, stt.Options{
|
|
||||||
Encoding: "pcm_s16le",
|
|
||||||
SampleRate: 16000,
|
|
||||||
Language: "zh-CN",
|
|
||||||
}).Return("你好,世界", nil)
|
|
||||||
|
|
||||||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
llmChunks := []llm.Chunk{
|
|
||||||
{Delta: "你好"},
|
|
||||||
{Delta: ",世界!"},
|
|
||||||
{Done: true, TokensUsed: &llm.TokenUsage{Prompt: 10, Completion: 5, Total: 15}},
|
|
||||||
}
|
|
||||||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
|
||||||
|
|
||||||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
|
||||||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
ttsChunks := []tts.Chunk{
|
|
||||||
{Audio: []byte("audio1"), IsLast: false},
|
|
||||||
{Audio: []byte("audio2"), IsLast: true},
|
|
||||||
}
|
|
||||||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).Return(createTTSStream(ttsChunks), nil)
|
|
||||||
|
|
||||||
mockSender.On("SendTTSAudio", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
// 创建 Pipeline
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
// 执行
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
// 验证
|
|
||||||
assert.NoError(t, err)
|
|
||||||
assert.Len(t, mockSender.STTResults, 1)
|
|
||||||
assert.Equal(t, "你好,世界", mockSender.STTResults[0].Text)
|
|
||||||
assert.Len(t, mockSender.LLMChunks, 2)
|
|
||||||
assert.Len(t, mockSender.LLMDones, 1)
|
|
||||||
assert.Len(t, mockSender.TTSAudios, 2)
|
|
||||||
|
|
||||||
mockSTT.AssertExpectations(t)
|
|
||||||
mockLLM.AssertExpectations(t)
|
|
||||||
mockTTS.AssertExpectations(t)
|
|
||||||
mockSession.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_STTError 测试 STT 失败降级
|
|
||||||
func TestProcessQuery_STTError(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(&models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}, nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).
|
|
||||||
Return("", errors.New("STT service unavailable"))
|
|
||||||
|
|
||||||
mockSender.On("SendError", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
|
||||||
assert.Len(t, mockSender.Errors, 1)
|
|
||||||
assert.Equal(t, "STT_ERROR", mockSender.Errors[0].Code)
|
|
||||||
|
|
||||||
mockSTT.AssertExpectations(t)
|
|
||||||
mockLLM.AssertNotCalled(t, "ChatStream")
|
|
||||||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_LLMError 测试 LLM 失败降级
|
|
||||||
func TestProcessQuery_LLMError(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
TTSEnabled: true,
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
|
||||||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
|
||||||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).
|
|
||||||
Return(nil, errors.New("LLM service unavailable"))
|
|
||||||
|
|
||||||
mockSender.On("SendError", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
|
||||||
assert.Len(t, mockSender.Errors, 1)
|
|
||||||
assert.Equal(t, "LLM_ERROR", mockSender.Errors[0].Code)
|
|
||||||
|
|
||||||
mockSTT.AssertExpectations(t)
|
|
||||||
mockLLM.AssertExpectations(t)
|
|
||||||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_TTSError 测试 TTS 失败静默跳过
|
|
||||||
func TestProcessQuery_TTSError(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
TTSEnabled: true,
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
|
||||||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
|
||||||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
llmChunks := []llm.Chunk{
|
|
||||||
{Delta: "你好"},
|
|
||||||
{Done: true},
|
|
||||||
}
|
|
||||||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
|
||||||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
|
||||||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).
|
|
||||||
Return(nil, errors.New("TTS service unavailable"))
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
// TTS 失败应该静默跳过,不返回错误
|
|
||||||
assert.NoError(t, err)
|
|
||||||
assert.Len(t, mockSender.LLMDones, 1)
|
|
||||||
assert.Len(t, mockSender.TTSAudios, 0)
|
|
||||||
|
|
||||||
mockSTT.AssertExpectations(t)
|
|
||||||
mockLLM.AssertExpectations(t)
|
|
||||||
mockTTS.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_ContextCancelled 测试上下文取消(Interrupt)
|
|
||||||
func TestProcessQuery_ContextCancelled(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
TTSEnabled: true,
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
|
||||||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
|
||||||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
// 创建一个会延迟的 LLM 流,以便我们可以取消上下文
|
|
||||||
llmCh := make(chan llm.Chunk)
|
|
||||||
go func() {
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
llmCh <- llm.Chunk{Delta: "你"}
|
|
||||||
time.Sleep(100 * time.Millisecond)
|
|
||||||
llmCh <- llm.Chunk{Delta: "好"}
|
|
||||||
close(llmCh)
|
|
||||||
}()
|
|
||||||
|
|
||||||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return((<-chan llm.Chunk)(llmCh), nil)
|
|
||||||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
|
||||||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
// 创建一个会延迟的 TTS 流
|
|
||||||
ttsCh := make(chan tts.Chunk)
|
|
||||||
go func() {
|
|
||||||
time.Sleep(200 * time.Millisecond)
|
|
||||||
close(ttsCh)
|
|
||||||
}()
|
|
||||||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).Return((<-chan tts.Chunk)(ttsCh), nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
// 创建可取消的上下文
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
|
|
||||||
// 在 50ms 后取消
|
|
||||||
go func() {
|
|
||||||
time.Sleep(50 * time.Millisecond)
|
|
||||||
cancel()
|
|
||||||
}()
|
|
||||||
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
// 上下文取消后,流程应该正常完成(中断流但不返回错误)
|
|
||||||
assert.NoError(t, err)
|
|
||||||
|
|
||||||
mockSTT.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_DisabledTTS 测试 TTS 未启用的情况
|
|
||||||
func TestProcessQuery_DisabledTTS(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
session := &models.Session{
|
|
||||||
ID: "session-123",
|
|
||||||
Config: models.SessionConfig{
|
|
||||||
TTSEnabled: false, // TTS 未启用
|
|
||||||
Language: "zh-CN",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
|
||||||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
|
||||||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
llmChunks := []llm.Chunk{
|
|
||||||
{Delta: "你好"},
|
|
||||||
{Done: true},
|
|
||||||
}
|
|
||||||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
|
||||||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
|
||||||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
assert.NoError(t, err)
|
|
||||||
assert.Len(t, mockSender.LLMDones, 1)
|
|
||||||
assert.Len(t, mockSender.TTSAudios, 0)
|
|
||||||
|
|
||||||
// TTS 不应该被调用
|
|
||||||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSplitter 测试句子切分器
|
|
||||||
func TestSplitter(t *testing.T) {
|
|
||||||
ch := make(chan string, 10)
|
|
||||||
splitter := NewSplitter(ch)
|
|
||||||
|
|
||||||
// 输入包含多个句子的文本
|
|
||||||
splitter.Feed("你好。")
|
|
||||||
splitter.Feed("世界!")
|
|
||||||
splitter.Feed("这是")
|
|
||||||
splitter.Feed("一个测试。")
|
|
||||||
splitter.Flush()
|
|
||||||
|
|
||||||
// 应该有 3 个句子
|
|
||||||
assert.Equal(t, 3, len(ch))
|
|
||||||
assert.Equal(t, "你好。", <-ch)
|
|
||||||
assert.Equal(t, "世界!", <-ch)
|
|
||||||
assert.Equal(t, "这是一个测试。", <-ch)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSplitter_NoDelimiter 测试没有分隔符的情况
|
|
||||||
func TestSplitter_NoDelimiter(t *testing.T) {
|
|
||||||
ch := make(chan string, 10)
|
|
||||||
splitter := NewSplitter(ch)
|
|
||||||
|
|
||||||
splitter.Feed("没有分隔符的文本")
|
|
||||||
splitter.Flush()
|
|
||||||
|
|
||||||
// 应该有 1 个句子(Flush 会发送剩余内容)
|
|
||||||
assert.Equal(t, 1, len(ch))
|
|
||||||
assert.Equal(t, "没有分隔符的文本", <-ch)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSplitter_Empty 测试空输入
|
|
||||||
func TestSplitter_Empty(t *testing.T) {
|
|
||||||
ch := make(chan string, 10)
|
|
||||||
splitter := NewSplitter(ch)
|
|
||||||
|
|
||||||
splitter.Flush()
|
|
||||||
|
|
||||||
// 应该没有句子
|
|
||||||
assert.Equal(t, 0, len(ch))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_InvalidAudio 测试无效音频数据
|
|
||||||
func TestProcessQuery_InvalidAudio(t *testing.T) {
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: "invalid-base64!!!",
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSender.On("SendError", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
|
||||||
assert.Len(t, mockSender.Errors, 1)
|
|
||||||
assert.Equal(t, "INVALID_MESSAGE", mockSender.Errors[0].Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProcessQuery_SessionNotFound 测试会话不存在
|
|
||||||
func TestProcessQuery_SessionNotFound(t *testing.T) {
|
|
||||||
audioData := []byte("test audio")
|
|
||||||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
|
||||||
|
|
||||||
req := models.WsQuery{
|
|
||||||
Type: "query",
|
|
||||||
RequestID: "req-123",
|
|
||||||
Audio: audioBase64,
|
|
||||||
}
|
|
||||||
|
|
||||||
mockSTT := new(MockSTTService)
|
|
||||||
mockLLM := new(MockLLMService)
|
|
||||||
mockTTS := new(MockTTSService)
|
|
||||||
mockSession := new(MockSessionManager)
|
|
||||||
mockSender := NewMockSender()
|
|
||||||
|
|
||||||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
|
||||||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
|
||||||
mockSession.On("Get", mock.Anything, "session-123").Return(nil, errors.New("session not found"))
|
|
||||||
|
|
||||||
mockSender.On("SendError", mock.Anything).Return(nil)
|
|
||||||
|
|
||||||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, &config.Config{
|
|
||||||
AI: config.AIConfig{
|
|
||||||
LLM: config.LLMConfig{Model: "gpt-4o"},
|
|
||||||
TTS: config.TTSConfig{Voice: "alloy", Speed: 1.0, OutputFormat: "mp3", SampleRate: 24000},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
|
||||||
|
|
||||||
assert.Error(t, err)
|
|
||||||
assert.Len(t, mockSender.Errors, 1)
|
|
||||||
assert.Equal(t, "SESSION_NOT_FOUND", mockSender.Errors[0].Code)
|
|
||||||
}
|
|
||||||
@@ -1,55 +0,0 @@
|
|||||||
package orchestrator
|
|
||||||
|
|
||||||
import "strings"
|
|
||||||
|
|
||||||
// sentenceDelimiters 句子分隔符集合。
|
|
||||||
var sentenceDelimiters = map[rune]bool{
|
|
||||||
'。': true,
|
|
||||||
'!': true,
|
|
||||||
'?': true,
|
|
||||||
'\n': true,
|
|
||||||
'.': true,
|
|
||||||
'!': true,
|
|
||||||
'?': true,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Splitter 句子切分器。
|
|
||||||
// 将流式文本按句子边界切分,发送到 channel 供 TTS 合成。
|
|
||||||
type Splitter struct {
|
|
||||||
ch chan<- string
|
|
||||||
buffer strings.Builder
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewSplitter 创建句子切分器。
|
|
||||||
// ch 用于接收切分后的句子文本。
|
|
||||||
func NewSplitter(ch chan<- string) *Splitter {
|
|
||||||
return &Splitter{
|
|
||||||
ch: ch,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Feed 输入增量文本,遇到句子分隔符时发送完整句子。
|
|
||||||
func (s *Splitter) Feed(delta string) {
|
|
||||||
for _, r := range delta {
|
|
||||||
s.buffer.WriteRune(r)
|
|
||||||
if sentenceDelimiters[r] {
|
|
||||||
s.flushBuffer()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Flush 刷新缓冲区中的剩余文本(即使没有句子分隔符)。
|
|
||||||
func (s *Splitter) Flush() {
|
|
||||||
if s.buffer.Len() > 0 {
|
|
||||||
s.flushBuffer()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// flushBuffer 将缓冲区内容发送到 channel 并清空。
|
|
||||||
func (s *Splitter) flushBuffer() {
|
|
||||||
text := strings.TrimSpace(s.buffer.String())
|
|
||||||
if text != "" {
|
|
||||||
s.ch <- text
|
|
||||||
}
|
|
||||||
s.buffer.Reset()
|
|
||||||
}
|
|
||||||
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()
|
||||||
|
}
|
||||||
51
backend/internal/ratelimit/middleware.go
Normal file
51
backend/internal/ratelimit/middleware.go
Normal file
@@ -0,0 +1,51 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
log.Warnw("rate limited",
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
"path", c.Request.URL.Path,
|
||||||
|
"limit_key", key,
|
||||||
|
"retry_after_sec", int(retryAfter.Seconds()+0.5))
|
||||||
|
|
||||||
|
// 设置 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"))
|
||||||
|
}
|
||||||
132
backend/internal/ratelimit/redis_bucket.go
Normal file
132
backend/internal/ratelimit/redis_bucket.go
Normal file
@@ -0,0 +1,132 @@
|
|||||||
|
package ratelimit
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/config"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
"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) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
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 {
|
||||||
|
log.Errorw("rate limit check failed", "key", key, "error", err)
|
||||||
|
// 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
|
||||||
|
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
|
||||||
|
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)
|
||||||
|
}
|
||||||
@@ -18,6 +18,7 @@ type ConversationSummary struct {
|
|||||||
Title string `json:"title"`
|
Title string `json:"title"`
|
||||||
LastMessage string `json:"last_message"`
|
LastMessage string `json:"last_message"`
|
||||||
MessageCount int `json:"message_count"`
|
MessageCount int `json:"message_count"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
UpdatedAt time.Time `json:"updated_at"`
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -142,11 +142,11 @@ func (m *MemoryManager) Create(ctx context.Context, userID string, config models
|
|||||||
}
|
}
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Write-Through:异步写 PG
|
// Write-Through:异步写 PG(使用 Background context,避免 HTTP 请求结束后 context 被取消)
|
||||||
if m.sessRepo != nil {
|
if m.sessRepo != nil {
|
||||||
go func() {
|
go func() {
|
||||||
cfgJSON, _ := json.Marshal(config)
|
cfgJSON, _ := json.Marshal(config)
|
||||||
if err := m.sessRepo.Save(ctx, store.SessionRecord{
|
if err := m.sessRepo.Save(context.Background(), store.SessionRecord{
|
||||||
ID: id, UserID: userID, Title: models.DefaultSessionTitle,
|
ID: id, UserID: userID, Title: models.DefaultSessionTitle,
|
||||||
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
|
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
@@ -204,11 +204,11 @@ func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patc
|
|||||||
cfg := entry.session.Config
|
cfg := entry.session.Config
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Write-Through:异步更新 PG
|
// Write-Through:异步更新 PG(使用 Background context)
|
||||||
if m.sessRepo != nil {
|
if m.sessRepo != nil {
|
||||||
go func() {
|
go func() {
|
||||||
cfgJSON, _ := json.Marshal(cfg)
|
cfgJSON, _ := json.Marshal(cfg)
|
||||||
if err := m.sessRepo.UpdateConfig(ctx, sessionID, cfgJSON); err != nil {
|
if err := m.sessRepo.UpdateConfig(context.Background(), sessionID, cfgJSON); err != nil {
|
||||||
logger.Log.Warnw("update session config in DB failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("update session config in DB failed", "session", sessionID, "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -233,10 +233,10 @@ func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title
|
|||||||
entry.lastActive = time.Now()
|
entry.lastActive = time.Now()
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Write-Through:异步更新 PG
|
// Write-Through:异步更新 PG(使用 Background context)
|
||||||
if m.sessRepo != nil {
|
if m.sessRepo != nil {
|
||||||
go func() {
|
go func() {
|
||||||
if err := m.sessRepo.UpdateTitle(ctx, sessionID, title); err != nil {
|
if err := m.sessRepo.UpdateTitle(context.Background(), sessionID, title); err != nil {
|
||||||
logger.Log.Warnw("update session title in DB failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("update session title in DB failed", "session", sessionID, "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -271,6 +271,7 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
|
|||||||
list = append(list, ConversationSummary{
|
list = append(list, ConversationSummary{
|
||||||
ID: rec.ID,
|
ID: rec.ID,
|
||||||
Title: rec.Title,
|
Title: rec.Title,
|
||||||
|
CreatedAt: rec.CreatedAt,
|
||||||
UpdatedAt: rec.UpdatedAt,
|
UpdatedAt: rec.UpdatedAt,
|
||||||
})
|
})
|
||||||
sessionIDs = append(sessionIDs, rec.ID)
|
sessionIDs = append(sessionIDs, rec.ID)
|
||||||
@@ -320,6 +321,7 @@ func (m *MemoryManager) listByUserFromMemory(ctx context.Context, userID string,
|
|||||||
summary := ConversationSummary{
|
summary := ConversationSummary{
|
||||||
ID: entry.session.ID,
|
ID: entry.session.ID,
|
||||||
Title: entry.session.Title,
|
Title: entry.session.Title,
|
||||||
|
CreatedAt: entry.session.CreatedAt,
|
||||||
UpdatedAt: entry.lastActive,
|
UpdatedAt: entry.lastActive,
|
||||||
}
|
}
|
||||||
summary.MessageCount = len(entry.history)
|
summary.MessageCount = len(entry.history)
|
||||||
@@ -394,8 +396,10 @@ func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg m
|
|||||||
entry.history = append(entry.history, msg)
|
entry.history = append(entry.history, msg)
|
||||||
|
|
||||||
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
|
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
|
||||||
|
titleUpdated := false
|
||||||
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
|
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
|
||||||
entry.session.Title = generateTitle(msg.Content)
|
entry.session.Title = generateTitle(msg.Content)
|
||||||
|
titleUpdated = true
|
||||||
}
|
}
|
||||||
|
|
||||||
// 超过上限时裁剪,保留最新的 maxHistory 条
|
// 超过上限时裁剪,保留最新的 maxHistory 条
|
||||||
@@ -406,13 +410,31 @@ func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg m
|
|||||||
now := time.Now()
|
now := time.Now()
|
||||||
entry.lastActive = now
|
entry.lastActive = now
|
||||||
entry.session.UpdatedAt = now
|
entry.session.UpdatedAt = now
|
||||||
|
|
||||||
|
// 复制标题(释放锁后安全使用)
|
||||||
|
persistTitle := entry.session.Title
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Write-Through:异步写冷存储,不阻塞调用方
|
// Write-Through:消息同步写入 PostgreSQL(保证调用顺序 = 插入顺序,
|
||||||
|
// 避免用户消息和 AI 消息的异步 goroutine 执行顺序不确定导致排序错乱)
|
||||||
if m.msgRepo != nil {
|
if m.msgRepo != nil {
|
||||||
|
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
|
||||||
|
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write-Through:异步更新会话元数据(标题 + updated_at)到 PostgreSQL
|
||||||
|
if m.sessRepo != nil {
|
||||||
go func() {
|
go func() {
|
||||||
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
|
if titleUpdated {
|
||||||
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
|
if err := m.sessRepo.UpdateTitle(context.Background(), sessionID, persistTitle); err != nil {
|
||||||
|
logger.Log.Warnw("persist session title failed", "session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// 即使标题没变,也要刷新 updated_at(保证列表排序正确)
|
||||||
|
if err := m.sessRepo.Touch(context.Background(), sessionID); err != nil {
|
||||||
|
logger.Log.Warnw("touch session in DB failed", "session", sessionID, "error", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
@@ -540,10 +562,10 @@ func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
|
|||||||
delete(m.sessions, sessionID)
|
delete(m.sessions, sessionID)
|
||||||
m.mu.Unlock()
|
m.mu.Unlock()
|
||||||
|
|
||||||
// Write-Through:异步删除 PG
|
// Write-Through:异步删除 PG(使用 Background context)
|
||||||
if m.sessRepo != nil {
|
if m.sessRepo != nil {
|
||||||
go func() {
|
go func() {
|
||||||
if err := m.sessRepo.Delete(ctx, sessionID); err != nil {
|
if err := m.sessRepo.Delete(context.Background(), sessionID); err != nil {
|
||||||
logger.Log.Warnw("delete session from DB failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("delete session from DB failed", "session", sessionID, "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -10,8 +10,9 @@ import (
|
|||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"github.com/redis/go-redis/v9"
|
"github.com/redis/go-redis/v9"
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
"github.com/hhs/camtalk/internal/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
// RedisManager 基于 Redis 的 SessionManager 实现。
|
// RedisManager 基于 Redis 的 SessionManager 实现。
|
||||||
@@ -36,13 +37,23 @@ func NewRedisManager(rdb *redis.Client, ttl time.Duration, maxHistory int) *Redi
|
|||||||
return &RedisManager{rdb: rdb, ttl: ttl, maxHistory: maxHistory}
|
return &RedisManager{rdb: rdb, ttl: ttl, maxHistory: maxHistory}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Ping 检查 Redis 连接是否正常。
|
||||||
|
func (m *RedisManager) Ping(ctx context.Context) error {
|
||||||
|
return m.rdb.Ping(ctx).Err()
|
||||||
|
}
|
||||||
|
|
||||||
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
|
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
|
||||||
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
|
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
|
||||||
func userSessKey(id string) string { return fmt.Sprintf("user:%s:sessions", id) }
|
func userSessKey(id string) string { return fmt.Sprintf("user:%s:sessions", id) }
|
||||||
|
|
||||||
// Create 创建新会话。userID 为空表示匿名会话。
|
// Create 创建新会话。userID 为空表示匿名会话。
|
||||||
func (m *RedisManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
func (m *RedisManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||||
id := uuidNew()
|
return m.CreateWithID(ctx, uuidNew(), userID, config)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateWithID 使用指定 ID 创建新会话。
|
||||||
|
// 供 TieredManager 调用,确保 L1/L2 使用相同的 session ID。
|
||||||
|
func (m *RedisManager) CreateWithID(ctx context.Context, id string, userID string, config models.SessionConfig) (string, error) {
|
||||||
now := time.Now().UTC()
|
now := time.Now().UTC()
|
||||||
|
|
||||||
pipe := m.rdb.Pipeline()
|
pipe := m.rdb.Pipeline()
|
||||||
@@ -77,7 +88,8 @@ func (m *RedisManager) Create(ctx context.Context, userID string, config models.
|
|||||||
return "", fmt.Errorf("redis create session: %w", err)
|
return "", fmt.Errorf("redis create session: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.Log.Debugw("redis session created", "session", id, "user_id", userID)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Debugw("redis session created", "session_id", id, "user_id", userID)
|
||||||
return id, nil
|
return id, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -86,8 +98,11 @@ const placeholderHistoryMark = "__placeholder__"
|
|||||||
|
|
||||||
// Get 获取会话。
|
// Get 获取会话。
|
||||||
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
|
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
|
||||||
return nil, fmt.Errorf("redis get session: %w", err)
|
return nil, fmt.Errorf("redis get session: %w", err)
|
||||||
}
|
}
|
||||||
if len(vals) == 0 {
|
if len(vals) == 0 {
|
||||||
@@ -105,6 +120,7 @@ func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Sessi
|
|||||||
sess.Config.DetailLevel = vals["config.detail_level"]
|
sess.Config.DetailLevel = vals["config.detail_level"]
|
||||||
sess.Config.Language = vals["config.language"]
|
sess.Config.Language = vals["config.language"]
|
||||||
|
|
||||||
|
log.Debugw("redis session retrieved", "session_id", sessionID)
|
||||||
return sess, nil
|
return sess, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -140,7 +156,9 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
|
|||||||
|
|
||||||
// 刷新 TTL
|
// 刷新 TTL
|
||||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||||
logger.Log.Debugw("redis session config updated", "session", sessionID)
|
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Debugw("redis session config updated", "session_id", sessionID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -160,7 +178,9 @@ func (m *RedisManager) UpdateTitle(ctx context.Context, sessionID string, title
|
|||||||
}
|
}
|
||||||
|
|
||||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||||
logger.Log.Debugw("redis session title updated", "session", sessionID, "title", title)
|
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Debugw("redis session title updated", "session_id", sessionID, "title", title)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -275,7 +295,11 @@ func (m *RedisManager) GetHistory(ctx context.Context, sessionID string, limit i
|
|||||||
}
|
}
|
||||||
var msg models.Message
|
var msg models.Message
|
||||||
if err := json.Unmarshal([]byte(raw), &msg); err != nil {
|
if err := json.Unmarshal([]byte(raw), &msg); err != nil {
|
||||||
logger.Log.Warnw("invalid history entry", "session", sessionID, "raw", raw)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("invalid history entry",
|
||||||
|
"session_id", sessionID,
|
||||||
|
"raw_len", len(raw),
|
||||||
|
"raw_preview", util.Truncate(raw, 100))
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
msgs = append(msgs, msg)
|
msgs = append(msgs, msg)
|
||||||
@@ -426,7 +450,8 @@ func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
|
|||||||
m.rdb.SRem(ctx, userSessKey(userID), sessionID)
|
m.rdb.SRem(ctx, userSessKey(userID), sessionID)
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.Log.Debugw("redis session destroyed", "session", sessionID)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Debugw("redis session destroyed", "session_id", sessionID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
361
backend/internal/session/tiered.go
Normal file
361
backend/internal/session/tiered.go
Normal file
@@ -0,0 +1,361 @@
|
|||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TieredManager 三级存储 SessionManager 实现。
|
||||||
|
//
|
||||||
|
// L1(内存)→ L2(Redis)→ L3(PostgreSQL)
|
||||||
|
//
|
||||||
|
// 读:L1 miss → L2 miss → L3,回填到 L1+L2
|
||||||
|
// 写:L1 → L2(同步)→ L3(异步)
|
||||||
|
// 降级:Redis 不可用时,回退到 L1+L3 模式
|
||||||
|
type TieredManager struct {
|
||||||
|
l1 *MemoryManager // L1: 内存缓存
|
||||||
|
l2 *RedisManager // L2: Redis(可选)
|
||||||
|
sessRepo store.SessionRepository // L3: PostgreSQL 会话持久化(可选)
|
||||||
|
msgRepo store.MessageRepository // L3: PostgreSQL 消息持久化(可选)
|
||||||
|
|
||||||
|
redisOK atomic.Bool // Redis 健康状态
|
||||||
|
stopCh chan struct{} // 停止信号
|
||||||
|
}
|
||||||
|
|
||||||
|
// TieredOption TieredManager 的函数式选项。
|
||||||
|
type TieredOption func(*TieredManager)
|
||||||
|
|
||||||
|
// WithTieredSessionRepository 注入 L3 会话持久化仓库。
|
||||||
|
func WithTieredSessionRepository(repo store.SessionRepository) TieredOption {
|
||||||
|
return func(m *TieredManager) {
|
||||||
|
m.sessRepo = repo
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithTieredMessageRepository 注入 L3 消息持久化仓库。
|
||||||
|
func WithTieredMessageRepository(repo store.MessageRepository) TieredOption {
|
||||||
|
return func(m *TieredManager) {
|
||||||
|
m.msgRepo = repo
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTieredManager 创建三级存储 SessionManager。
|
||||||
|
// l2 为 nil 时降级为 L1+L3 模式。
|
||||||
|
func NewTieredManager(
|
||||||
|
ttl time.Duration,
|
||||||
|
maxHistory int,
|
||||||
|
l2 *RedisManager,
|
||||||
|
opts ...TieredOption,
|
||||||
|
) *TieredManager {
|
||||||
|
m := &TieredManager{
|
||||||
|
l2: l2,
|
||||||
|
stopCh: make(chan struct{}),
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(m)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 初始化 L1(内存),注入 L3 仓库实现 Write-Through
|
||||||
|
var l1Opts []Option
|
||||||
|
if m.sessRepo != nil {
|
||||||
|
l1Opts = append(l1Opts, WithSessionRepository(m.sessRepo))
|
||||||
|
}
|
||||||
|
if m.msgRepo != nil {
|
||||||
|
l1Opts = append(l1Opts, WithMessageRepository(m.msgRepo))
|
||||||
|
}
|
||||||
|
m.l1 = NewMemoryManager(ttl, maxHistory, l1Opts...)
|
||||||
|
|
||||||
|
// 初始化 Redis 健康状态
|
||||||
|
if l2 != nil {
|
||||||
|
m.redisOK.Store(true)
|
||||||
|
go m.healthCheck()
|
||||||
|
}
|
||||||
|
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// healthCheck 定期检查 Redis 健康状态。
|
||||||
|
func (m *TieredManager) healthCheck() {
|
||||||
|
ticker := time.NewTicker(30 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ticker.C:
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
|
err := m.l2.Ping(ctx)
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
wasOK := m.redisOK.Load()
|
||||||
|
isOK := err == nil
|
||||||
|
m.redisOK.Store(isOK)
|
||||||
|
|
||||||
|
if wasOK && !isOK {
|
||||||
|
logger.Log.Warn("Redis connection lost, degrading to L1+L3 mode")
|
||||||
|
} else if !wasOK && isOK {
|
||||||
|
logger.Log.Info("Redis connection restored, resuming L1+L2+L3 mode")
|
||||||
|
}
|
||||||
|
case <-m.stopCh:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isRedisOK 检查 Redis 是否可用。
|
||||||
|
func (m *TieredManager) isRedisOK() bool {
|
||||||
|
return m.l2 != nil && m.redisOK.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop 停止 TieredManager(清理后台 goroutine)。
|
||||||
|
func (m *TieredManager) Stop() {
|
||||||
|
close(m.stopCh)
|
||||||
|
m.l1.Stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create 创建新会话。
|
||||||
|
// 写入:L1 → L2(同步)→ L3(异步)
|
||||||
|
func (m *TieredManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||||
|
// L1: 内存
|
||||||
|
id, err := m.l1.Create(ctx, userID, config)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步),使用 L1 生成的 ID 保证一致性
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if _, err := m.l2.CreateWithID(ctx, id, userID, config); err != nil {
|
||||||
|
logger.Log.Warnw("Redis Create failed, continuing without L2",
|
||||||
|
"session", id, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||||
|
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get 获取会话。
|
||||||
|
// 读取:L1 → L2(回填 L1)→ L3(回填 L1+L2)
|
||||||
|
func (m *TieredManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||||
|
// L1: 内存
|
||||||
|
sess, err := m.l1.Get(ctx, sessionID)
|
||||||
|
if err == nil {
|
||||||
|
return sess, nil
|
||||||
|
}
|
||||||
|
if err != ErrSessionNotFound {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis
|
||||||
|
if m.isRedisOK() {
|
||||||
|
sess, err = m.l2.Get(ctx, sessionID)
|
||||||
|
if err == nil {
|
||||||
|
// 回填 L1
|
||||||
|
history, _ := m.l2.GetHistory(ctx, sessionID, 0)
|
||||||
|
m.l1.LoadSession(sess, history)
|
||||||
|
return sess, nil
|
||||||
|
}
|
||||||
|
if err != ErrSessionNotFound {
|
||||||
|
logger.Log.Warnw("Redis Get failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// L3: PostgreSQL(由 L1 的 Cache-Aside 处理)
|
||||||
|
// L1.Get 已经实现了从 PostgreSQL 恢复的逻辑
|
||||||
|
return nil, ErrSessionNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateConfig 更新会话配置。
|
||||||
|
// 写入:L1 → L2(同步)→ L3(异步)
|
||||||
|
func (m *TieredManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
||||||
|
// L1: 内存
|
||||||
|
if err := m.l1.UpdateConfig(ctx, sessionID, patch); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.UpdateConfig(ctx, sessionID, patch); err != nil {
|
||||||
|
logger.Log.Warnw("Redis UpdateConfig failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateTitle 更新会话标题。
|
||||||
|
// 写入:L1 → L2(同步)→ L3(异步)
|
||||||
|
func (m *TieredManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
||||||
|
// L1: 内存
|
||||||
|
if err := m.l1.UpdateTitle(ctx, sessionID, title); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.UpdateTitle(ctx, sessionID, title); err != nil {
|
||||||
|
logger.Log.Warnw("Redis UpdateTitle failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListByUser 查询用户的会话列表。
|
||||||
|
// 读取:L1 + L2 + L3 合并去重
|
||||||
|
func (m *TieredManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
||||||
|
// 优先使用 L1(已集成 L3 回退逻辑)
|
||||||
|
return m.l1.ListByUser(ctx, userID, page, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetHistory 获取对话历史。
|
||||||
|
// 读取:L1 → L2(回填 L1)→ L3(回填 L1+L2)
|
||||||
|
func (m *TieredManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||||
|
// L1: 内存
|
||||||
|
msgs, err := m.l1.GetHistory(ctx, sessionID, limit)
|
||||||
|
if err == nil && len(msgs) > 0 {
|
||||||
|
return msgs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis
|
||||||
|
if m.isRedisOK() {
|
||||||
|
msgs, err = m.l2.GetHistory(ctx, sessionID, limit)
|
||||||
|
if err == nil && len(msgs) > 0 {
|
||||||
|
// 回填 L1(通过 Get 触发)
|
||||||
|
m.l1.Get(ctx, sessionID)
|
||||||
|
return msgs, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// L3: PostgreSQL(由 L1 的 Cache-Aside 处理)
|
||||||
|
return nil, ErrSessionNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
// AppendMessage 追加消息。
|
||||||
|
// 写入:L1 → L2(同步)→ L3(异步)
|
||||||
|
func (m *TieredManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
|
||||||
|
// L1: 内存(Write-Through 到 L3)
|
||||||
|
if err := m.l1.AppendMessage(ctx, sessionID, msg); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.AppendMessage(ctx, sessionID, msg); err != nil {
|
||||||
|
logger.Log.Warnw("Redis AppendMessage failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetActiveRequest 设置当前活跃请求。
|
||||||
|
func (m *TieredManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
|
||||||
|
// L1: 内存
|
||||||
|
if err := m.l1.SetActiveRequest(ctx, sessionID, requestID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.SetActiveRequest(ctx, sessionID, requestID); err != nil {
|
||||||
|
logger.Log.Warnw("Redis SetActiveRequest failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||||
|
func (m *TieredManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
|
||||||
|
// L1: 内存
|
||||||
|
id, err := m.l1.GetActiveRequestID(ctx, sessionID)
|
||||||
|
if err == nil && id != "" {
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis
|
||||||
|
if m.isRedisOK() {
|
||||||
|
id, err = m.l2.GetActiveRequestID(ctx, sessionID)
|
||||||
|
if err == nil && id != "" {
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearActiveRequest 清除当前活跃请求。
|
||||||
|
func (m *TieredManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
|
||||||
|
// L1: 内存
|
||||||
|
if err := m.l1.ClearActiveRequest(ctx, sessionID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.ClearActiveRequest(ctx, sessionID); err != nil {
|
||||||
|
logger.Log.Warnw("Redis ClearActiveRequest failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Touch 刷新会话活跃时间。
|
||||||
|
func (m *TieredManager) Touch(ctx context.Context, sessionID string) error {
|
||||||
|
// L1: 内存
|
||||||
|
if err := m.l1.Touch(ctx, sessionID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.Touch(ctx, sessionID); err != nil {
|
||||||
|
logger.Log.Warnw("Redis Touch failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Destroy 销毁会话。
|
||||||
|
// 写入:L1 → L2 → L3
|
||||||
|
func (m *TieredManager) Destroy(ctx context.Context, sessionID string) error {
|
||||||
|
// L1: 内存(Write-Through 到 L3)
|
||||||
|
if err := m.l1.Destroy(ctx, sessionID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// L2: Redis(同步)
|
||||||
|
if m.isRedisOK() {
|
||||||
|
if err := m.l2.Destroy(ctx, sessionID); err != nil {
|
||||||
|
logger.Log.Warnw("Redis Destroy failed",
|
||||||
|
"session", sessionID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ActiveCount 返回活跃会话数量。
|
||||||
|
func (m *TieredManager) ActiveCount() int {
|
||||||
|
return m.l1.ActiveCount()
|
||||||
|
}
|
||||||
172
backend/internal/store/cached_user.go
Normal file
172
backend/internal/store/cached_user.go
Normal file
@@ -0,0 +1,172 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Redis key 前缀。
|
||||||
|
const (
|
||||||
|
refreshTokenPrefix = "auth:refresh:" // auth:refresh:{token_hash} → user_id
|
||||||
|
userRefreshPrefix = "auth:user_refresh:" // auth:user_refresh:{user_id} → Set of token_hash
|
||||||
|
)
|
||||||
|
|
||||||
|
// CachedUserRepository 装饰器,为 UserRepository 的 refresh token 操作增加 Redis 缓存。
|
||||||
|
// 读路径:Redis miss → DB → 回填 Redis。
|
||||||
|
// 写路径:同步双写 Redis + DB。
|
||||||
|
// 删路径:同步双删 Redis + DB。
|
||||||
|
// Redis 操作失败时降级到纯 DB,不阻断主流程。
|
||||||
|
type CachedUserRepository struct {
|
||||||
|
inner UserRepository
|
||||||
|
rdb *redis.Client
|
||||||
|
backfillTTL time.Duration // DB 回填 Redis 时使用的默认 TTL
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCachedUserRepository 创建带 Redis 缓存的 UserRepository 装饰器。
|
||||||
|
// backfillTTL: 从 DB 回填 Redis 时使用的 TTL(因 DB 接口不返回 expiresAt)。
|
||||||
|
func NewCachedUserRepository(inner UserRepository, rdb *redis.Client, backfillTTL time.Duration) *CachedUserRepository {
|
||||||
|
if backfillTTL <= 0 {
|
||||||
|
backfillTTL = 24 * time.Hour
|
||||||
|
}
|
||||||
|
return &CachedUserRepository{
|
||||||
|
inner: inner,
|
||||||
|
rdb: rdb,
|
||||||
|
backfillTTL: backfillTTL,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// refreshTokenKey 生成 refresh token 的 Redis key。
|
||||||
|
func refreshTokenKey(tokenHash string) string {
|
||||||
|
return refreshTokenPrefix + tokenHash
|
||||||
|
}
|
||||||
|
|
||||||
|
// userRefreshKey 生成用户 refresh token 集合的 Redis key。
|
||||||
|
func userRefreshKey(userID string) string {
|
||||||
|
return userRefreshPrefix + userID
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- 委托方法(不做缓存) ---
|
||||||
|
|
||||||
|
func (r *CachedUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
|
||||||
|
return r.inner.Create(ctx, username, passwordHash)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *CachedUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
|
||||||
|
return r.inner.FindByUsername(ctx, username)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *CachedUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
|
||||||
|
return r.inner.FindByID(ctx, id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- 缓存方法 ---
|
||||||
|
|
||||||
|
// SaveRefreshToken Write-Through:先写 DB,再写 Redis。
|
||||||
|
func (r *CachedUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||||
|
// 先写 DB
|
||||||
|
if err := r.inner.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 写 Redis(SET + SADD),设置 TTL 为 token 剩余有效期
|
||||||
|
ttl := time.Until(expiresAt)
|
||||||
|
if ttl <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
key := refreshTokenKey(tokenHash)
|
||||||
|
pipe := r.rdb.Pipeline()
|
||||||
|
pipe.Set(ctx, key, userID, ttl)
|
||||||
|
pipe.SAdd(ctx, userRefreshKey(userID), tokenHash)
|
||||||
|
if _, err := pipe.Exec(ctx); err != nil {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("redis cache write failed for refresh token", "error", err)
|
||||||
|
// 降级:DB 已写入成功,Redis 失败不影响正确性
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FindRefreshToken Read-Through:先查 Redis,miss 时查 DB 并回填。
|
||||||
|
func (r *CachedUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||||
|
key := refreshTokenKey(tokenHash)
|
||||||
|
|
||||||
|
// 查 Redis
|
||||||
|
userID, err := r.rdb.Get(ctx, key).Result()
|
||||||
|
if err == nil {
|
||||||
|
return userID, nil
|
||||||
|
}
|
||||||
|
// redis.Nil 表示 key 不存在,其他错误记录日志后降级到 DB
|
||||||
|
if err != redis.Nil {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("redis cache read failed for refresh token", "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 降级到 DB
|
||||||
|
userID, err = r.inner.FindRefreshToken(ctx, tokenHash)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 回填 Redis(SET + SADD),TTL 使用保守默认值
|
||||||
|
go func() {
|
||||||
|
bgCtx := context.Background()
|
||||||
|
pipe := r.rdb.Pipeline()
|
||||||
|
pipe.Set(bgCtx, key, userID, r.backfillTTL)
|
||||||
|
pipe.SAdd(bgCtx, userRefreshKey(userID), tokenHash)
|
||||||
|
_, _ = pipe.Exec(bgCtx)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return userID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteRefreshToken 双删:先删 DB,再删 Redis。
|
||||||
|
func (r *CachedUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
|
||||||
|
// 先从 Redis 获取 user_id(用于从集合中移除)
|
||||||
|
userID, _ := r.rdb.Get(ctx, refreshTokenKey(tokenHash)).Result()
|
||||||
|
|
||||||
|
// 删 DB
|
||||||
|
if err := r.inner.DeleteRefreshToken(ctx, tokenHash); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 删 Redis
|
||||||
|
key := refreshTokenKey(tokenHash)
|
||||||
|
pipe := r.rdb.Pipeline()
|
||||||
|
pipe.Del(ctx, key)
|
||||||
|
if userID != "" {
|
||||||
|
pipe.SRem(ctx, userRefreshKey(userID), tokenHash)
|
||||||
|
}
|
||||||
|
if _, err := pipe.Exec(ctx); err != nil {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("redis cache delete failed for refresh token", "error", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteUserRefreshTokens 批量清理:先从 Redis 获取集合,逐个删缓存,再删 DB。
|
||||||
|
func (r *CachedUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
|
||||||
|
userKey := userRefreshKey(userID)
|
||||||
|
|
||||||
|
// 从 Redis 获取该用户所有 token hash
|
||||||
|
hashes, _ := r.rdb.SMembers(ctx, userKey).Result()
|
||||||
|
|
||||||
|
// 批量删除 Redis 缓存
|
||||||
|
if len(hashes) > 0 {
|
||||||
|
keys := make([]string, 0, len(hashes)+1)
|
||||||
|
for _, h := range hashes {
|
||||||
|
keys = append(keys, refreshTokenKey(h))
|
||||||
|
}
|
||||||
|
keys = append(keys, userKey)
|
||||||
|
if err := r.rdb.Del(ctx, keys...).Err(); err != nil {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("redis cache batch delete failed for user refresh tokens", "error", err, "user_id", userID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 删 DB(无论 Redis 是否成功都执行)
|
||||||
|
return r.inner.DeleteUserRefreshTokens(ctx, userID)
|
||||||
|
}
|
||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"github.com/jackc/pgx/v5/pgxpool"
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。
|
// PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。
|
||||||
@@ -21,14 +22,24 @@ func NewPgMessageRepository(pool *pgxpool.Pool) *PgMessageRepository {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error {
|
func (r *PgMessageRepository) SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
_, err := r.pool.Exec(ctx,
|
_, err := r.pool.Exec(ctx,
|
||||||
`INSERT INTO messages (session_id, role, content, tokens_used) VALUES ($1, $2, $3, $4)`,
|
`INSERT INTO messages (session_id, role, content, tokens_used) VALUES ($1, $2, $3, $4)`,
|
||||||
sessionID, msg.Role, msg.Content, tokensUsed,
|
sessionID, msg.Role, msg.Content, tokensUsed,
|
||||||
)
|
)
|
||||||
return err
|
if err != nil {
|
||||||
|
log.Errorw("save message failed", "session_id", sessionID, "role", msg.Role, "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("message saved", "session_id", sessionID, "role", msg.Role, "tokens_used", tokensUsed)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) {
|
func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
if limit <= 0 {
|
if limit <= 0 {
|
||||||
limit = 50
|
limit = 50
|
||||||
}
|
}
|
||||||
@@ -56,6 +67,7 @@ func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string,
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("get messages failed", "session_id", sessionID, "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -64,31 +76,39 @@ func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string,
|
|||||||
rows[i], rows[j] = rows[j], rows[i]
|
rows[i], rows[j] = rows[j], rows[i]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("messages retrieved", "session_id", sessionID, "count", len(rows))
|
||||||
return rows, nil
|
return rows, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) {
|
func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
pgxRows, err := r.pool.Query(ctx, query, args...)
|
pgxRows, err := r.pool.Query(ctx, query, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("query messages failed", "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
defer pgxRows.Close()
|
defer pgxRows.Close()
|
||||||
|
|
||||||
var messages []StoredMessage
|
messages := make([]StoredMessage, 0)
|
||||||
for pgxRows.Next() {
|
for pgxRows.Next() {
|
||||||
var m StoredMessage
|
var m StoredMessage
|
||||||
if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil {
|
if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil {
|
||||||
|
log.Errorw("scan message row failed", "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
messages = append(messages, m)
|
messages = append(messages, m)
|
||||||
}
|
}
|
||||||
if err := pgxRows.Err(); err != nil {
|
if err := pgxRows.Err(); err != nil {
|
||||||
|
log.Errorw("iterate message rows failed", "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return messages, nil
|
return messages, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
|
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var m StoredMessage
|
var m StoredMessage
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT id, session_id, role, content, tokens_used, created_at
|
`SELECT id, session_id, role, content, tokens_used, created_at
|
||||||
@@ -102,24 +122,34 @@ func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID stri
|
|||||||
return nil, ErrMessageNotFound
|
return nil, ErrMessageNotFound
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("get last message failed", "session_id", sessionID, "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("last message retrieved", "session_id", sessionID, "message_id", m.ID)
|
||||||
return &m, nil
|
return &m, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
|
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var count int
|
var count int
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
|
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
|
||||||
sessionID,
|
sessionID,
|
||||||
).Scan(&count)
|
).Scan(&count)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("get message count failed", "session_id", sessionID, "error", err)
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("message count retrieved", "session_id", sessionID, "count", count)
|
||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
|
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
if len(sessionIDs) == 0 {
|
if len(sessionIDs) == 0 {
|
||||||
return map[string]SessionMessageStats{}, nil
|
return map[string]SessionMessageStats{}, nil
|
||||||
}
|
}
|
||||||
@@ -143,6 +173,7 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
|
|||||||
sessionIDs,
|
sessionIDs,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
@@ -152,12 +183,16 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
|
|||||||
var sid string
|
var sid string
|
||||||
var stats SessionMessageStats
|
var stats SessionMessageStats
|
||||||
if err := rows.Scan(&sid, &stats.MessageCount, &stats.LastMessage); err != nil {
|
if err := rows.Scan(&sid, &stats.MessageCount, &stats.LastMessage); err != nil {
|
||||||
|
log.Errorw("scan message stats row failed", "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
result[sid] = stats
|
result[sid] = stats
|
||||||
}
|
}
|
||||||
if err := rows.Err(); err != nil {
|
if err := rows.Err(); err != nil {
|
||||||
|
log.Errorw("iterate message stats rows failed", "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session message stats retrieved", "session_count", len(sessionIDs), "result_count", len(result))
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
"github.com/jackc/pgx/v5/pgxpool"
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
|
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
|
||||||
@@ -19,6 +21,8 @@ func NewPgSessionRepository(pool *pgxpool.Pool) *PgSessionRepository {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
_, err := r.pool.Exec(ctx,
|
_, err := r.pool.Exec(ctx,
|
||||||
`INSERT INTO sessions (id, user_id, title, config, created_at, updated_at)
|
`INSERT INTO sessions (id, user_id, title, config, created_at, updated_at)
|
||||||
VALUES ($1, $2, $3, $4, $5, $6)
|
VALUES ($1, $2, $3, $4, $5, $6)
|
||||||
@@ -28,10 +32,18 @@ func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
|||||||
updated_at = EXCLUDED.updated_at`,
|
updated_at = EXCLUDED.updated_at`,
|
||||||
s.ID, s.UserID, s.Title, s.Config, s.CreatedAt, s.UpdatedAt,
|
s.ID, s.UserID, s.Title, s.Config, s.CreatedAt, s.UpdatedAt,
|
||||||
)
|
)
|
||||||
return err
|
if err != nil {
|
||||||
|
log.Errorw("save session failed", "session_id", s.ID, "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("session saved", "session_id", s.ID, "user_id", s.UserID)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*SessionRecord, error) {
|
func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*SessionRecord, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var s SessionRecord
|
var s SessionRecord
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT id, user_id, title, config, created_at, updated_at
|
`SELECT id, user_id, title, config, created_at, updated_at
|
||||||
@@ -41,12 +53,17 @@ func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*Session
|
|||||||
return nil, ErrSessionNotFound
|
return nil, ErrSessionNotFound
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("find session failed", "session_id", id, "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session found", "session_id", id)
|
||||||
return &s, nil
|
return &s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error) {
|
func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
if page <= 0 {
|
if page <= 0 {
|
||||||
page = 1
|
page = 1
|
||||||
}
|
}
|
||||||
@@ -60,6 +77,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
|||||||
if err := r.pool.QueryRow(ctx,
|
if err := r.pool.QueryRow(ctx,
|
||||||
`SELECT COUNT(*) FROM sessions WHERE user_id = $1`, userID,
|
`SELECT COUNT(*) FROM sessions WHERE user_id = $1`, userID,
|
||||||
).Scan(&total); err != nil {
|
).Scan(&total); err != nil {
|
||||||
|
log.Errorw("count user sessions failed", "user_id", userID, "error", err)
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -73,6 +91,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
|||||||
userID, size, offset,
|
userID, size, offset,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("find user sessions failed", "user_id", userID, "error", err)
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
@@ -81,66 +100,90 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
|||||||
for rows.Next() {
|
for rows.Next() {
|
||||||
var s SessionRecord
|
var s SessionRecord
|
||||||
if err := rows.Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt); err != nil {
|
if err := rows.Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt); err != nil {
|
||||||
|
log.Errorw("scan session row failed", "user_id", userID, "error", err)
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
list = append(list, s)
|
list = append(list, s)
|
||||||
}
|
}
|
||||||
if err := rows.Err(); err != nil {
|
if err := rows.Err(); err != nil {
|
||||||
|
log.Errorw("iterate session rows failed", "user_id", userID, "error", err)
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("user sessions found", "user_id", userID, "count", len(list), "total", total)
|
||||||
return list, total, nil
|
return list, total, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) UpdateTitle(ctx context.Context, id string, title string) error {
|
func (r *PgSessionRepository) UpdateTitle(ctx context.Context, id string, title string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
tag, err := r.pool.Exec(ctx,
|
tag, err := r.pool.Exec(ctx,
|
||||||
`UPDATE sessions SET title = $2, updated_at = NOW() WHERE id = $1`,
|
`UPDATE sessions SET title = $2, updated_at = NOW() WHERE id = $1`,
|
||||||
id, title,
|
id, title,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("update session title failed", "session_id", id, "error", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if tag.RowsAffected() == 0 {
|
if tag.RowsAffected() == 0 {
|
||||||
return ErrSessionNotFound
|
return ErrSessionNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session title updated", "session_id", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) UpdateConfig(ctx context.Context, id string, configJSON []byte) error {
|
func (r *PgSessionRepository) UpdateConfig(ctx context.Context, id string, configJSON []byte) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
tag, err := r.pool.Exec(ctx,
|
tag, err := r.pool.Exec(ctx,
|
||||||
`UPDATE sessions SET config = $2, updated_at = NOW() WHERE id = $1`,
|
`UPDATE sessions SET config = $2, updated_at = NOW() WHERE id = $1`,
|
||||||
id, configJSON,
|
id, configJSON,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("update session config failed", "session_id", id, "error", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if tag.RowsAffected() == 0 {
|
if tag.RowsAffected() == 0 {
|
||||||
return ErrSessionNotFound
|
return ErrSessionNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session config updated", "session_id", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) Touch(ctx context.Context, id string) error {
|
func (r *PgSessionRepository) Touch(ctx context.Context, id string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
tag, err := r.pool.Exec(ctx,
|
tag, err := r.pool.Exec(ctx,
|
||||||
`UPDATE sessions SET updated_at = NOW() WHERE id = $1`, id,
|
`UPDATE sessions SET updated_at = NOW() WHERE id = $1`, id,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("touch session failed", "session_id", id, "error", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if tag.RowsAffected() == 0 {
|
if tag.RowsAffected() == 0 {
|
||||||
return ErrSessionNotFound
|
return ErrSessionNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session touched", "session_id", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgSessionRepository) Delete(ctx context.Context, id string) error {
|
func (r *PgSessionRepository) Delete(ctx context.Context, id string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
tag, err := r.pool.Exec(ctx,
|
tag, err := r.pool.Exec(ctx,
|
||||||
`DELETE FROM sessions WHERE id = $1`, id,
|
`DELETE FROM sessions WHERE id = $1`, id,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("delete session failed", "session_id", id, "error", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if tag.RowsAffected() == 0 {
|
if tag.RowsAffected() == 0 {
|
||||||
return ErrSessionNotFound
|
return ErrSessionNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("session deleted", "session_id", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
"github.com/jackc/pgx/v5/pgxpool"
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PgUserRepository 基于 PostgreSQL 的 UserRepository 实现。
|
// PgUserRepository 基于 PostgreSQL 的 UserRepository 实现。
|
||||||
@@ -20,18 +22,25 @@ func NewPgUserRepository(pool *pgxpool.Pool) *PgUserRepository {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
|
func (r *PgUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var id string
|
var id string
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`INSERT INTO users (username, password_hash) VALUES ($1, $2) RETURNING id`,
|
`INSERT INTO users (username, password_hash) VALUES ($1, $2) RETURNING id`,
|
||||||
username, passwordHash,
|
username, passwordHash,
|
||||||
).Scan(&id)
|
).Scan(&id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("create user failed", "username", username, "error", err)
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("user created", "user_id", id, "username", username)
|
||||||
return id, nil
|
return id, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
|
func (r *PgUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var u User
|
var u User
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE username = $1`,
|
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE username = $1`,
|
||||||
@@ -41,12 +50,17 @@ func (r *PgUserRepository) FindByUsername(ctx context.Context, username string)
|
|||||||
return nil, ErrUserNotFound
|
return nil, ErrUserNotFound
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("find user by username failed", "username", username, "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("user found by username", "user_id", u.ID, "username", username)
|
||||||
return &u, nil
|
return &u, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
|
func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var u User
|
var u User
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE id = $1`,
|
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE id = $1`,
|
||||||
@@ -56,20 +70,33 @@ func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, erro
|
|||||||
return nil, ErrUserNotFound
|
return nil, ErrUserNotFound
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("find user by id failed", "user_id", id, "error", err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("user found by id", "user_id", id)
|
||||||
return &u, nil
|
return &u, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
func (r *PgUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
_, err := r.pool.Exec(ctx,
|
_, err := r.pool.Exec(ctx,
|
||||||
`INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`,
|
`INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`,
|
||||||
userID, tokenHash, expiresAt,
|
userID, tokenHash, expiresAt,
|
||||||
)
|
)
|
||||||
return err
|
if err != nil {
|
||||||
|
log.Errorw("save refresh token failed", "user_id", userID, "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("refresh token saved", "user_id", userID)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
var userID string
|
var userID string
|
||||||
err := r.pool.QueryRow(ctx,
|
err := r.pool.QueryRow(ctx,
|
||||||
`SELECT user_id FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW()`,
|
`SELECT user_id FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW()`,
|
||||||
@@ -79,23 +106,42 @@ func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash strin
|
|||||||
return "", ErrRefreshTokenNotFound
|
return "", ErrRefreshTokenNotFound
|
||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
log.Errorw("find refresh token failed", "error", err)
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Debugw("refresh token found", "user_id", userID)
|
||||||
return userID, nil
|
return userID, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
|
func (r *PgUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
_, err := r.pool.Exec(ctx,
|
_, err := r.pool.Exec(ctx,
|
||||||
`DELETE FROM refresh_tokens WHERE token_hash = $1`,
|
`DELETE FROM refresh_tokens WHERE token_hash = $1`,
|
||||||
tokenHash,
|
tokenHash,
|
||||||
)
|
)
|
||||||
return err
|
if err != nil {
|
||||||
|
log.Errorw("delete refresh token failed", "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("refresh token deleted")
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *PgUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
|
func (r *PgUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
_, err := r.pool.Exec(ctx,
|
_, err := r.pool.Exec(ctx,
|
||||||
`DELETE FROM refresh_tokens WHERE user_id = $1`,
|
`DELETE FROM refresh_tokens WHERE user_id = $1`,
|
||||||
userID,
|
userID,
|
||||||
)
|
)
|
||||||
return err
|
if err != nil {
|
||||||
|
log.Errorw("delete user refresh tokens failed", "user_id", userID, "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user refresh tokens deleted", "user_id", userID)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
276
backend/internal/store/user_scenario_repository.go
Normal file
276
backend/internal/store/user_scenario_repository.go
Normal file
@@ -0,0 +1,276 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UserScenarioRepository 用户自建情景仓储接口。
|
||||||
|
type UserScenarioRepository interface {
|
||||||
|
Create(ctx context.Context, scenario *models.UserScenario) error
|
||||||
|
FindByID(ctx context.Context, id string) (*models.UserScenario, error)
|
||||||
|
FindByIDAndUserID(ctx context.Context, id, userID string) (*models.UserScenario, error)
|
||||||
|
FindByUserID(ctx context.Context, userID string) ([]*models.UserScenario, error)
|
||||||
|
Update(ctx context.Context, scenario *models.UserScenario) error
|
||||||
|
Delete(ctx context.Context, id string) error
|
||||||
|
CountByUserID(ctx context.Context, userID string) (int, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PostgresUserScenarioRepo PostgreSQL 实现。
|
||||||
|
type PostgresUserScenarioRepo struct {
|
||||||
|
pool *pgxpool.Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPostgresUserScenarioRepo 创建 PostgreSQL 用户情景仓储。
|
||||||
|
func NewPostgresUserScenarioRepo(pool *pgxpool.Pool) UserScenarioRepository {
|
||||||
|
return &PostgresUserScenarioRepo{pool: pool}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create 创建用户情景。
|
||||||
|
func (r *PostgresUserScenarioRepo) Create(ctx context.Context, scenario *models.UserScenario) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `
|
||||||
|
INSERT INTO user_scenarios (id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at)
|
||||||
|
VALUES ($1, $2, $3, $4, NULLIF($5, ''), $6, NULLIF($7, ''), $8, $9, $10)
|
||||||
|
RETURNING id, created_at, updated_at
|
||||||
|
`
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
scenario.CreatedAt = now
|
||||||
|
scenario.UpdatedAt = now
|
||||||
|
|
||||||
|
if scenario.ID == "" {
|
||||||
|
scenario.ID = uuid.New().String()
|
||||||
|
}
|
||||||
|
if scenario.Icon == "" {
|
||||||
|
scenario.Icon = "✨"
|
||||||
|
}
|
||||||
|
if scenario.Language == "" {
|
||||||
|
scenario.Language = "zh-CN"
|
||||||
|
}
|
||||||
|
|
||||||
|
err := r.pool.QueryRow(ctx, query,
|
||||||
|
scenario.ID,
|
||||||
|
scenario.UserID,
|
||||||
|
scenario.Name,
|
||||||
|
scenario.Icon,
|
||||||
|
scenario.Description,
|
||||||
|
scenario.Prompt,
|
||||||
|
scenario.Greeting,
|
||||||
|
scenario.Language,
|
||||||
|
scenario.CreatedAt,
|
||||||
|
scenario.UpdatedAt,
|
||||||
|
).Scan(&scenario.ID, &scenario.CreatedAt, &scenario.UpdatedAt)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("create user scenario failed", "user_id", scenario.UserID, "name", scenario.Name, "error", err)
|
||||||
|
return fmt.Errorf("create user scenario: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenario created", "scenario_id", scenario.ID, "user_id", scenario.UserID, "name", scenario.Name)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FindByID 根据 ID 查找情景。
|
||||||
|
func (r *PostgresUserScenarioRepo) FindByID(ctx context.Context, id string) (*models.UserScenario, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `
|
||||||
|
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||||
|
FROM user_scenarios
|
||||||
|
WHERE id = $1
|
||||||
|
`
|
||||||
|
|
||||||
|
var scenario models.UserScenario
|
||||||
|
err := r.pool.QueryRow(ctx, query, id).Scan(
|
||||||
|
&scenario.ID,
|
||||||
|
&scenario.UserID,
|
||||||
|
&scenario.Name,
|
||||||
|
&scenario.Icon,
|
||||||
|
&scenario.Description,
|
||||||
|
&scenario.Prompt,
|
||||||
|
&scenario.Greeting,
|
||||||
|
&scenario.Language,
|
||||||
|
&scenario.CreatedAt,
|
||||||
|
&scenario.UpdatedAt,
|
||||||
|
)
|
||||||
|
|
||||||
|
if err == pgx.ErrNoRows {
|
||||||
|
return nil, fmt.Errorf("user scenario not found: %s", id)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("find user scenario failed", "scenario_id", id, "error", err)
|
||||||
|
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenario found", "scenario_id", id)
|
||||||
|
return &scenario, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FindByIDAndUserID 根据 ID 和用户 ID 查找情景(权限校验)。
|
||||||
|
func (r *PostgresUserScenarioRepo) FindByIDAndUserID(ctx context.Context, id, userID string) (*models.UserScenario, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `
|
||||||
|
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||||
|
FROM user_scenarios
|
||||||
|
WHERE id = $1 AND user_id = $2
|
||||||
|
`
|
||||||
|
|
||||||
|
var scenario models.UserScenario
|
||||||
|
err := r.pool.QueryRow(ctx, query, id, userID).Scan(
|
||||||
|
&scenario.ID,
|
||||||
|
&scenario.UserID,
|
||||||
|
&scenario.Name,
|
||||||
|
&scenario.Icon,
|
||||||
|
&scenario.Description,
|
||||||
|
&scenario.Prompt,
|
||||||
|
&scenario.Greeting,
|
||||||
|
&scenario.Language,
|
||||||
|
&scenario.CreatedAt,
|
||||||
|
&scenario.UpdatedAt,
|
||||||
|
)
|
||||||
|
|
||||||
|
if err == pgx.ErrNoRows {
|
||||||
|
return nil, fmt.Errorf("user scenario not found or no permission")
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("find user scenario by id and user failed", "scenario_id", id, "user_id", userID, "error", err)
|
||||||
|
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenario found by id and user", "scenario_id", id, "user_id", userID)
|
||||||
|
return &scenario, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FindByUserID 查找用户的所有情景。
|
||||||
|
func (r *PostgresUserScenarioRepo) FindByUserID(ctx context.Context, userID string) ([]*models.UserScenario, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `
|
||||||
|
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||||
|
FROM user_scenarios
|
||||||
|
WHERE user_id = $1
|
||||||
|
ORDER BY created_at DESC
|
||||||
|
`
|
||||||
|
|
||||||
|
rows, err := r.pool.Query(ctx, query, userID)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("find user scenarios failed", "user_id", userID, "error", err)
|
||||||
|
return nil, fmt.Errorf("find user scenarios: %w", err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
var scenarios []*models.UserScenario
|
||||||
|
for rows.Next() {
|
||||||
|
var s models.UserScenario
|
||||||
|
err := rows.Scan(
|
||||||
|
&s.ID,
|
||||||
|
&s.UserID,
|
||||||
|
&s.Name,
|
||||||
|
&s.Icon,
|
||||||
|
&s.Description,
|
||||||
|
&s.Prompt,
|
||||||
|
&s.Greeting,
|
||||||
|
&s.Language,
|
||||||
|
&s.CreatedAt,
|
||||||
|
&s.UpdatedAt,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("scan user scenario row failed", "user_id", userID, "error", err)
|
||||||
|
return nil, fmt.Errorf("scan user scenario: %w", err)
|
||||||
|
}
|
||||||
|
scenarios = append(scenarios, &s)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = rows.Err(); err != nil {
|
||||||
|
log.Errorw("iterate user scenarios failed", "user_id", userID, "error", err)
|
||||||
|
return nil, fmt.Errorf("iterate user scenarios: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenarios found", "user_id", userID, "count", len(scenarios))
|
||||||
|
return scenarios, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update 更新用户情景。
|
||||||
|
func (r *PostgresUserScenarioRepo) Update(ctx context.Context, scenario *models.UserScenario) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `
|
||||||
|
UPDATE user_scenarios
|
||||||
|
SET name = $1, icon = $2, description = $3, prompt = $4, greeting = $5, language = $6, updated_at = $7
|
||||||
|
WHERE id = $8 AND user_id = $9
|
||||||
|
RETURNING updated_at
|
||||||
|
`
|
||||||
|
|
||||||
|
scenario.UpdatedAt = time.Now()
|
||||||
|
|
||||||
|
err := r.pool.QueryRow(ctx, query,
|
||||||
|
scenario.Name,
|
||||||
|
scenario.Icon,
|
||||||
|
scenario.Description,
|
||||||
|
scenario.Prompt,
|
||||||
|
scenario.Greeting,
|
||||||
|
scenario.Language,
|
||||||
|
scenario.UpdatedAt,
|
||||||
|
scenario.ID,
|
||||||
|
scenario.UserID,
|
||||||
|
).Scan(&scenario.UpdatedAt)
|
||||||
|
|
||||||
|
if err == pgx.ErrNoRows {
|
||||||
|
return fmt.Errorf("user scenario not found or no permission")
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("update user scenario failed", "scenario_id", scenario.ID, "user_id", scenario.UserID, "error", err)
|
||||||
|
return fmt.Errorf("update user scenario: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenario updated", "scenario_id", scenario.ID, "user_id", scenario.UserID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete 删除用户情景。
|
||||||
|
func (r *PostgresUserScenarioRepo) Delete(ctx context.Context, id string) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `DELETE FROM user_scenarios WHERE id = $1`
|
||||||
|
|
||||||
|
result, err := r.pool.Exec(ctx, query, id)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("delete user scenario failed", "scenario_id", id, "error", err)
|
||||||
|
return fmt.Errorf("delete user scenario: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if result.RowsAffected() == 0 {
|
||||||
|
return fmt.Errorf("user scenario not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenario deleted", "scenario_id", id)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CountByUserID 统计用户的情景数量。
|
||||||
|
func (r *PostgresUserScenarioRepo) CountByUserID(ctx context.Context, userID string) (int, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
query := `SELECT COUNT(*) FROM user_scenarios WHERE user_id = $1`
|
||||||
|
|
||||||
|
var count int
|
||||||
|
err := r.pool.QueryRow(ctx, query, userID).Scan(&count)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("count user scenarios failed", "user_id", userID, "error", err)
|
||||||
|
return 0, fmt.Errorf("count user scenarios: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("user scenarios counted", "user_id", userID, "count", count)
|
||||||
|
return count, nil
|
||||||
|
}
|
||||||
46
backend/internal/trace/context.go
Normal file
46
backend/internal/trace/context.go
Normal file
@@ -0,0 +1,46 @@
|
|||||||
|
package trace
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
type traceIDKey struct{}
|
||||||
|
type requestIDKey struct{}
|
||||||
|
type sessionIDKey struct{}
|
||||||
|
|
||||||
|
// WithTraceID 将 trace ID 注入 context(连接级/会话级标识)
|
||||||
|
func WithTraceID(ctx context.Context, traceID string) context.Context {
|
||||||
|
return context.WithValue(ctx, traceIDKey{}, traceID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetTraceID 从 context 提取 trace ID
|
||||||
|
func GetTraceID(ctx context.Context) string {
|
||||||
|
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithRequestID 将 request ID 注入 context(单次请求/查询标识)
|
||||||
|
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||||
|
return context.WithValue(ctx, requestIDKey{}, requestID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRequestID 从 context 提取 request ID
|
||||||
|
func GetRequestID(ctx context.Context) string {
|
||||||
|
if v, ok := ctx.Value(requestIDKey{}).(string); ok {
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithSessionID 将 session ID 注入 context(会话存储标识)
|
||||||
|
func WithSessionID(ctx context.Context, sessionID string) context.Context {
|
||||||
|
return context.WithValue(ctx, sessionIDKey{}, sessionID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetSessionID 从 context 提取 session ID
|
||||||
|
func GetSessionID(ctx context.Context) string {
|
||||||
|
if v, ok := ctx.Value(sessionIDKey{}).(string); ok {
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
42
backend/internal/trace/eino_test.go
Normal file
42
backend/internal/trace/eino_test.go
Normal file
@@ -0,0 +1,42 @@
|
|||||||
|
package trace_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEinoContextPropagation(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
testTraceID := "01J5TEST123456789"
|
||||||
|
ctx = trace.WithTraceID(ctx, testTraceID)
|
||||||
|
|
||||||
|
var capturedTraceID string
|
||||||
|
|
||||||
|
g := compose.NewGraph[string, string]()
|
||||||
|
g.AddLambdaNode("test_node", compose.InvokableLambda(
|
||||||
|
func(ctx context.Context, input string) (string, error) {
|
||||||
|
capturedTraceID = trace.GetTraceID(ctx)
|
||||||
|
return "ok", nil
|
||||||
|
},
|
||||||
|
))
|
||||||
|
g.AddEdge(compose.START, "test_node")
|
||||||
|
g.AddEdge("test_node", compose.END)
|
||||||
|
|
||||||
|
runnable, err := g.Compile(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("compile failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = runnable.Invoke(ctx, "test_input")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("invoke failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if capturedTraceID != testTraceID {
|
||||||
|
t.Errorf("trace_id lost in Eino propagation: got %q, want %q",
|
||||||
|
capturedTraceID, testTraceID)
|
||||||
|
}
|
||||||
|
}
|
||||||
63
backend/internal/trace/gin_logger.go
Normal file
63
backend/internal/trace/gin_logger.go
Normal file
@@ -0,0 +1,63 @@
|
|||||||
|
package trace
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GinLogger 记录每个 HTTP 请求的 method/path/status/latency
|
||||||
|
func GinLogger() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
start := time.Now()
|
||||||
|
path := c.Request.URL.Path
|
||||||
|
query := c.Request.URL.RawQuery
|
||||||
|
|
||||||
|
c.Next()
|
||||||
|
|
||||||
|
latency := time.Since(start).Milliseconds()
|
||||||
|
status := c.Writer.Status()
|
||||||
|
log := FromContext(c.Request.Context())
|
||||||
|
|
||||||
|
fields := []interface{}{
|
||||||
|
"method", c.Request.Method,
|
||||||
|
"path", path,
|
||||||
|
"status", status,
|
||||||
|
"latency_ms", latency,
|
||||||
|
"client_ip", c.ClientIP(),
|
||||||
|
}
|
||||||
|
if query != "" {
|
||||||
|
fields = append(fields, "query", query)
|
||||||
|
}
|
||||||
|
if errStr := c.Errors.String(); errStr != "" {
|
||||||
|
fields = append(fields, "errors", errStr)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case status >= 500:
|
||||||
|
log.Errorw("request completed", fields...)
|
||||||
|
case status >= 400:
|
||||||
|
log.Warnw("request completed", fields...)
|
||||||
|
default:
|
||||||
|
log.Infow("request completed", fields...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GinRecovery 自定义 panic 恢复中间件,使用 zap 记录
|
||||||
|
func GinRecovery() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
defer func() {
|
||||||
|
if err := recover(); err != nil {
|
||||||
|
log := FromContext(c.Request.Context())
|
||||||
|
log.Errorw("panic recovered",
|
||||||
|
"error", err,
|
||||||
|
"path", c.Request.URL.Path,
|
||||||
|
"method", c.Request.Method,
|
||||||
|
"client_ip", c.ClientIP())
|
||||||
|
c.AbortWithStatus(500)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
22
backend/internal/trace/id.go
Normal file
22
backend/internal/trace/id.go
Normal file
@@ -0,0 +1,22 @@
|
|||||||
|
package trace
|
||||||
|
|
||||||
|
import (
|
||||||
|
cryptorand "crypto/rand"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/oklog/ulid/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
var entropyPool = sync.Pool{
|
||||||
|
New: func() interface{} {
|
||||||
|
return ulid.Monotonic(cryptorand.Reader, 0)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateTraceID 生成并发安全的 ULID trace ID
|
||||||
|
func GenerateTraceID() string {
|
||||||
|
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
|
||||||
|
defer entropyPool.Put(entropy)
|
||||||
|
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
|
||||||
|
}
|
||||||
25
backend/internal/trace/logger.go
Normal file
25
backend/internal/trace/logger.go
Normal file
@@ -0,0 +1,25 @@
|
|||||||
|
package trace
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
|
||||||
|
func FromContext(ctx context.Context) *zap.SugaredLogger {
|
||||||
|
log := logger.Log
|
||||||
|
|
||||||
|
if traceID := GetTraceID(ctx); traceID != "" {
|
||||||
|
log = log.With("trace_id", traceID)
|
||||||
|
}
|
||||||
|
if requestID := GetRequestID(ctx); requestID != "" {
|
||||||
|
log = log.With("request_id", requestID)
|
||||||
|
}
|
||||||
|
if sessionID := GetSessionID(ctx); sessionID != "" {
|
||||||
|
log = log.With("session_id", sessionID)
|
||||||
|
}
|
||||||
|
|
||||||
|
return log
|
||||||
|
}
|
||||||
17
backend/internal/trace/middleware.go
Normal file
17
backend/internal/trace/middleware.go
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
package trace
|
||||||
|
|
||||||
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
|
||||||
|
func TraceMiddleware() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
traceID := GenerateTraceID()
|
||||||
|
ctx := WithTraceID(c.Request.Context(), traceID)
|
||||||
|
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
|
||||||
|
|
||||||
|
c.Request = c.Request.WithContext(ctx)
|
||||||
|
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
|
||||||
|
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
9
backend/internal/util/string.go
Normal file
9
backend/internal/util/string.go
Normal file
@@ -0,0 +1,9 @@
|
|||||||
|
package util
|
||||||
|
|
||||||
|
// Truncate 截断字符串到指定长度,超出部分用 "..." 替换
|
||||||
|
func Truncate(s string, maxLen int) string {
|
||||||
|
if len(s) <= maxLen {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return s[:maxLen] + "..."
|
||||||
|
}
|
||||||
@@ -3,6 +3,7 @@ package ws
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -10,13 +11,16 @@ import (
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/llm"
|
||||||
"github.com/hhs/camtalk/internal/auth"
|
"github.com/hhs/camtalk/internal/auth"
|
||||||
"github.com/hhs/camtalk/internal/config"
|
"github.com/hhs/camtalk/internal/config"
|
||||||
"github.com/hhs/camtalk/internal/errors"
|
"github.com/hhs/camtalk/internal/errors"
|
||||||
"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"
|
||||||
|
"github.com/hhs/camtalk/internal/store"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// newUpgrader 根据配置创建 WebSocket upgrader。
|
// newUpgrader 根据配置创建 WebSocket upgrader。
|
||||||
@@ -40,12 +44,12 @@ func newUpgrader(cfg *config.Config) websocket.Upgrader {
|
|||||||
|
|
||||||
// Client 代表一个 WebSocket 客户端连接。
|
// Client 代表一个 WebSocket 客户端连接。
|
||||||
type Client struct {
|
type Client struct {
|
||||||
conn *websocket.Conn
|
conn *websocket.Conn
|
||||||
sessionID string
|
sessionID string
|
||||||
sessionMgr session.Manager
|
sessionMgr session.Manager
|
||||||
orchestrator orchestrator.Orchestrator
|
orchestrator orchestrator.Orchestrator
|
||||||
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
|
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
||||||
@@ -92,21 +96,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, scenarioRepo store.UserScenarioRepository) 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
|
||||||
|
|
||||||
maxHistory := cfg.Session.MaxHistory
|
|
||||||
|
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory, tokenMgr)
|
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr, limiter, scenarioRepo)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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, maxHistory int, tokenMgr *auth.TokenManager) {
|
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter, scenarioRepo store.UserScenarioRepository) {
|
||||||
|
|
||||||
// --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) ---
|
// --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) ---
|
||||||
token := c.Query("token")
|
token := c.Query("token")
|
||||||
@@ -132,9 +134,20 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 生成连接级 trace ID(整个 WebSocket 生命周期使用)
|
||||||
|
ctx := c.Request.Context()
|
||||||
|
traceID := trace.GetTraceID(ctx)
|
||||||
|
if traceID == "" {
|
||||||
|
// 如果 REST 中间件未生成(不应发生),fallback 生成
|
||||||
|
traceID = trace.GenerateTraceID()
|
||||||
|
ctx = trace.WithTraceID(ctx, traceID)
|
||||||
|
c.Request = c.Request.WithContext(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Errorw("websocket upgrade failed", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
@@ -143,13 +156,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
var sessionID string
|
var sessionID string
|
||||||
if conversationID != "" {
|
if conversationID != "" {
|
||||||
sessionID = conversationID
|
sessionID = conversationID
|
||||||
logger.Log.Infow("resuming conversation", "session", sessionID, "user_id", userID)
|
ctx = trace.WithSessionID(ctx, sessionID)
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Infow("resuming conversation", "user_id", userID)
|
||||||
} else {
|
} else {
|
||||||
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
|
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Log.Errorw("create session failed", "error", err)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Errorw("create session failed", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
ctx = trace.WithSessionID(ctx, sessionID)
|
||||||
}
|
}
|
||||||
|
|
||||||
client := &Client{
|
client := &Client{
|
||||||
@@ -166,7 +183,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
SessionID: sessionID,
|
SessionID: sessionID,
|
||||||
ServerVersion: version,
|
ServerVersion: version,
|
||||||
})
|
})
|
||||||
logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Infow("client connected", "user_id", userID, "username", username)
|
||||||
|
|
||||||
// 心跳检测
|
// 心跳检测
|
||||||
lastPong := time.Now()
|
lastPong := time.Now()
|
||||||
@@ -184,7 +202,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
select {
|
select {
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
if time.Since(lastPong) > heartbeatTimeout {
|
if time.Since(lastPong) > heartbeatTimeout {
|
||||||
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("heartbeat timeout")
|
||||||
conn.Close()
|
conn.Close()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -199,7 +218,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
_, message, err := conn.ReadMessage()
|
_, message, err := conn.ReadMessage()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
|
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
|
||||||
logger.Log.Warnw("ws read error", "error", err)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("ws read error", "error", err)
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -224,23 +244,36 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
|
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
|
||||||
|
// 注入 request ID 到 context
|
||||||
|
queryCtx := trace.WithRequestID(ctx, msg.RequestID)
|
||||||
|
log := trace.FromContext(queryCtx)
|
||||||
|
log.Infow("query received", "has_image", msg.Image != "", "has_audio", msg.Audio != "")
|
||||||
|
|
||||||
|
// 限流检查
|
||||||
|
if limiter != nil {
|
||||||
|
key := fmt.Sprintf("%s:query", userID)
|
||||||
|
allowed, retryAfter := limiter.Allow(context.Background(), key)
|
||||||
|
if !allowed {
|
||||||
|
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)
|
log.Warnw("touch session failed", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 标记活跃请求
|
// 标记活跃请求
|
||||||
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
||||||
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
|
log.Warnw("set active request failed", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 获取对话历史
|
|
||||||
history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, maxHistory)
|
|
||||||
|
|
||||||
// 创建可取消的 context
|
// 创建可取消的 context
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
processCtx, cancel := context.WithCancel(queryCtx)
|
||||||
client.mu.Lock()
|
client.mu.Lock()
|
||||||
client.cancelFuncs[msg.RequestID] = cancel
|
client.cancelFuncs[msg.RequestID] = cancel
|
||||||
client.mu.Unlock()
|
client.mu.Unlock()
|
||||||
@@ -260,8 +293,9 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, history, sender); err != nil {
|
if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil {
|
||||||
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
|
log := trace.FromContext(processCtx)
|
||||||
|
log.Errorw("process query failed", "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
@@ -282,10 +316,66 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
errors.SendWSError(client, errors.CodeInternalError, "", err)
|
errors.SendWSError(client, errors.CodeInternalError, "", err)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
logger.Log.Infow("config updated", "session", sessionID)
|
|
||||||
|
scenarioID := ""
|
||||||
|
if msg.Payload.Scenario != nil {
|
||||||
|
scenarioID = *msg.Payload.Scenario
|
||||||
|
}
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Infow("config updated", "scenario", scenarioID)
|
||||||
|
|
||||||
|
// 如果切换了情景(非自由对话),返回首句引导
|
||||||
|
if scenarioID != "" && scenarioID != "free_chat" {
|
||||||
|
sess, err := client.sessionMgr.Get(context.Background(), sessionID)
|
||||||
|
if err == nil && sess != nil {
|
||||||
|
// 加载用户自建情景
|
||||||
|
var customGreetings map[string]string
|
||||||
|
if sess.UserID != "" && scenarioRepo != nil {
|
||||||
|
scenarios, err := scenarioRepo.FindByUserID(context.Background(), sess.UserID)
|
||||||
|
if err == nil && len(scenarios) > 0 {
|
||||||
|
customGreetings = make(map[string]string, len(scenarios))
|
||||||
|
for _, s := range scenarios {
|
||||||
|
if s.Greeting != "" {
|
||||||
|
customGreetings[s.ID] = s.Greeting
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
greeting := llm.GetScenarioGreeting(scenarioID, sess.Config.Language, customGreetings)
|
||||||
|
if greeting != "" {
|
||||||
|
// 发送首句作为 AI 消息
|
||||||
|
_ = client.SendJSON(models.WsLLMChunk{
|
||||||
|
Type: "llm_chunk",
|
||||||
|
RequestID: "scenario_greeting",
|
||||||
|
Delta: greeting,
|
||||||
|
Role: "assistant",
|
||||||
|
})
|
||||||
|
|
||||||
|
doneMsg := models.WsLLMDone{
|
||||||
|
Type: "llm_done",
|
||||||
|
RequestID: "scenario_greeting",
|
||||||
|
FullText: greeting,
|
||||||
|
Model: "",
|
||||||
|
LatencyMs: 0,
|
||||||
|
}
|
||||||
|
doneMsg.TokensUsed.Prompt = 0
|
||||||
|
doneMsg.TokensUsed.Completion = 0
|
||||||
|
doneMsg.TokensUsed.Total = 0
|
||||||
|
_ = client.SendJSON(doneMsg)
|
||||||
|
|
||||||
|
// 追加首句到历史记录
|
||||||
|
_ = client.sessionMgr.AppendMessage(context.Background(), sessionID, models.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: greeting,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
case "interrupt":
|
case "interrupt":
|
||||||
logger.Log.Infow("interrupt received", "session", sessionID)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Infow("interrupt received")
|
||||||
|
|
||||||
// 获取活跃请求 ID 并取消
|
// 获取活跃请求 ID 并取消
|
||||||
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
||||||
@@ -313,12 +403,14 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
// 取消所有活跃请求
|
// 取消所有活跃请求
|
||||||
client.mu.Lock()
|
client.mu.Lock()
|
||||||
for reqID, cancel := range client.cancelFuncs {
|
for reqID, cancel := range client.cancelFuncs {
|
||||||
logger.Log.Infow("canceling active request on disconnect", "session", sessionID, "request", reqID)
|
log := trace.FromContext(ctx)
|
||||||
|
log.Infow("canceling active request on disconnect", "request", reqID)
|
||||||
cancel()
|
cancel()
|
||||||
}
|
}
|
||||||
client.cancelFuncs = make(map[string]context.CancelFunc)
|
client.cancelFuncs = make(map[string]context.CancelFunc)
|
||||||
client.mu.Unlock()
|
client.mu.Unlock()
|
||||||
|
|
||||||
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
||||||
logger.Log.Infow("client disconnected", "session", sessionID)
|
log = trace.FromContext(ctx)
|
||||||
|
log.Infow("client disconnected")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,11 +8,11 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"context"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"context"
|
|
||||||
|
|
||||||
"github.com/hhs/camtalk/internal/auth"
|
"github.com/hhs/camtalk/internal/auth"
|
||||||
"github.com/hhs/camtalk/internal/config"
|
"github.com/hhs/camtalk/internal/config"
|
||||||
@@ -48,7 +48,6 @@ func (m *MockOrchestrator) ProcessQuery(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
sessionID string,
|
sessionID string,
|
||||||
req models.WsQuery,
|
req models.WsQuery,
|
||||||
history []models.Message,
|
|
||||||
sender orchestrator.Sender,
|
sender orchestrator.Sender,
|
||||||
) error {
|
) error {
|
||||||
if m.Err != nil {
|
if m.Err != nil {
|
||||||
@@ -149,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, nil))
|
||||||
|
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
|
|
||||||
@@ -222,9 +221,9 @@ func TestWS_QueryFullFlow(t *testing.T) {
|
|||||||
imageB64 := base64.StdEncoding.EncodeToString([]byte("fake-image-data"))
|
imageB64 := base64.StdEncoding.EncodeToString([]byte("fake-image-data"))
|
||||||
|
|
||||||
mock := &MockOrchestrator{
|
mock := &MockOrchestrator{
|
||||||
STTResult: "你好,世界",
|
STTResult: "你好,世界",
|
||||||
LLMDeltas: []string{"你好", ",世界!"},
|
LLMDeltas: []string{"你好", ",世界!"},
|
||||||
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
|
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
|
||||||
}
|
}
|
||||||
|
|
||||||
srv, wsURL := setupTestServer(t, mock)
|
srv, wsURL := setupTestServer(t, mock)
|
||||||
@@ -333,7 +332,7 @@ func TestWS_UnknownMessageType(t *testing.T) {
|
|||||||
err := conn.WriteJSON(map[string]string{"type": "unknown_type"})
|
err := conn.WriteJSON(map[string]string{"type": "unknown_type"})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
errMsg := readJSON(t, conn)
|
errMsg := readJSON(t, conn)
|
||||||
assert.Equal(t, "error", errMsg["type"])
|
assert.Equal(t, "error", errMsg["type"])
|
||||||
assert.Equal(t, "INVALID_MESSAGE", errMsg["code"])
|
assert.Equal(t, "INVALID_MESSAGE", errMsg["code"])
|
||||||
assert.Contains(t, errMsg["message"], "unknown message type")
|
assert.Contains(t, errMsg["message"], "unknown message type")
|
||||||
@@ -592,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, nil))
|
||||||
|
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
return srv, tokenMgr, sessionMgr
|
return srv, tokenMgr, sessionMgr
|
||||||
@@ -643,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, nil))
|
||||||
srv := httptest.NewServer(r)
|
srv := httptest.NewServer(r)
|
||||||
defer srv.Close()
|
defer srv.Close()
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,14 @@ CREATE TABLE IF NOT EXISTS users (
|
|||||||
-- 用户名索引(用于登录查询)
|
-- 用户名索引(用于登录查询)
|
||||||
CREATE INDEX IF NOT EXISTS idx_users_username ON users(username);
|
CREATE INDEX IF NOT EXISTS idx_users_username ON users(username);
|
||||||
|
|
||||||
|
-- 表和列注释
|
||||||
|
COMMENT ON TABLE users IS '用户表,存储系统所有注册用户的基本信息';
|
||||||
|
COMMENT ON COLUMN users.id IS '用户唯一标识符 (UUID)';
|
||||||
|
COMMENT ON COLUMN users.username IS '用户名,最大 64 字符,全局唯一';
|
||||||
|
COMMENT ON COLUMN users.password_hash IS '密码哈希值,使用 bcrypt 算法(cost=10)';
|
||||||
|
COMMENT ON COLUMN users.created_at IS '用户注册时间';
|
||||||
|
COMMENT ON COLUMN users.updated_at IS '用户信息最后更新时间';
|
||||||
|
|
||||||
-- Refresh Token 表
|
-- Refresh Token 表
|
||||||
CREATE TABLE IF NOT EXISTS refresh_tokens (
|
CREATE TABLE IF NOT EXISTS refresh_tokens (
|
||||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
@@ -24,3 +32,11 @@ CREATE INDEX IF NOT EXISTS idx_refresh_tokens_token_hash ON refresh_tokens(token
|
|||||||
|
|
||||||
-- 用户 ID 索引(用于登出所有设备)
|
-- 用户 ID 索引(用于登出所有设备)
|
||||||
CREATE INDEX IF NOT EXISTS idx_refresh_tokens_user_id ON refresh_tokens(user_id);
|
CREATE INDEX IF NOT EXISTS idx_refresh_tokens_user_id ON refresh_tokens(user_id);
|
||||||
|
|
||||||
|
-- 表和列注释
|
||||||
|
COMMENT ON TABLE refresh_tokens IS 'Refresh Token 表,用于 JWT 双 token 机制的长期身份验证';
|
||||||
|
COMMENT ON COLUMN refresh_tokens.id IS 'Token 唯一标识符 (UUID)';
|
||||||
|
COMMENT ON COLUMN refresh_tokens.user_id IS '所属用户 ID,外键关联 users 表,用户删除时级联删除';
|
||||||
|
COMMENT ON COLUMN refresh_tokens.token_hash IS 'Token 哈希值,使用 SHA-256 算法,十六进制编码 (64 字符)';
|
||||||
|
COMMENT ON COLUMN refresh_tokens.expires_at IS 'Token 过期时间,默认有效期 7 天';
|
||||||
|
COMMENT ON COLUMN refresh_tokens.created_at IS 'Token 创建时间';
|
||||||
|
|||||||
@@ -2,10 +2,12 @@
|
|||||||
CREATE TABLE IF NOT EXISTS messages (
|
CREATE TABLE IF NOT EXISTS messages (
|
||||||
id BIGSERIAL PRIMARY KEY,
|
id BIGSERIAL PRIMARY KEY,
|
||||||
session_id UUID NOT NULL,
|
session_id UUID NOT NULL,
|
||||||
role VARCHAR(16) NOT NULL, -- "user" | "assistant" | "system"
|
role VARCHAR(10) NOT NULL,
|
||||||
content TEXT NOT NULL,
|
content TEXT NOT NULL,
|
||||||
tokens_used INTEGER NOT NULL DEFAULT 0,
|
tokens_used INTEGER NOT NULL DEFAULT 0,
|
||||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
|
||||||
|
CONSTRAINT check_tokens_non_negative CHECK (tokens_used >= 0)
|
||||||
);
|
);
|
||||||
|
|
||||||
-- 按会话查询消息(分页核心索引)
|
-- 按会话查询消息(分页核心索引)
|
||||||
@@ -15,3 +17,12 @@ CREATE INDEX IF NOT EXISTS idx_messages_session_id_created_at
|
|||||||
-- 按会话查询最后一条消息
|
-- 按会话查询最后一条消息
|
||||||
CREATE INDEX IF NOT EXISTS idx_messages_session_id_id_desc
|
CREATE INDEX IF NOT EXISTS idx_messages_session_id_id_desc
|
||||||
ON messages(session_id, id DESC);
|
ON messages(session_id, id DESC);
|
||||||
|
|
||||||
|
-- 表和列注释
|
||||||
|
COMMENT ON TABLE messages IS '消息表,存储所有会话的消息记录';
|
||||||
|
COMMENT ON COLUMN messages.id IS '消息唯一标识符,自增序列';
|
||||||
|
COMMENT ON COLUMN messages.session_id IS '所属会话 ID,关联 sessions 表';
|
||||||
|
COMMENT ON COLUMN messages.role IS '消息角色,可选值: ''user'' (用户), ''assistant'' (AI 助手), ''system'' (系统)';
|
||||||
|
COMMENT ON COLUMN messages.content IS '消息内容,无长度限制';
|
||||||
|
COMMENT ON COLUMN messages.tokens_used IS '消息消耗的 token 数量,用于计费统计';
|
||||||
|
COMMENT ON COLUMN messages.created_at IS '消息创建时间';
|
||||||
|
|||||||
@@ -1,11 +1,22 @@
|
|||||||
CREATE TABLE IF NOT EXISTS sessions (
|
CREATE TABLE IF NOT EXISTS sessions (
|
||||||
id UUID PRIMARY KEY,
|
id UUID PRIMARY KEY,
|
||||||
user_id UUID NOT NULL,
|
user_id UUID NOT NULL,
|
||||||
title VARCHAR(256) NOT NULL DEFAULT '新对话',
|
title VARCHAR(100) NOT NULL DEFAULT '新对话',
|
||||||
config JSONB NOT NULL DEFAULT '{}',
|
config JSONB NOT NULL DEFAULT '{}',
|
||||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||||
|
|
||||||
|
CONSTRAINT check_title_length CHECK (char_length(title) >= 1 AND char_length(title) <= 100)
|
||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions (user_id);
|
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions (user_id);
|
||||||
CREATE INDEX IF NOT EXISTS idx_sessions_user_updated ON sessions (user_id, updated_at DESC);
|
CREATE INDEX IF NOT EXISTS idx_sessions_user_updated ON sessions (user_id, updated_at DESC);
|
||||||
|
|
||||||
|
-- 表和列注释
|
||||||
|
COMMENT ON TABLE sessions IS '会话表,存储用户的对话会话信息';
|
||||||
|
COMMENT ON COLUMN sessions.id IS '会话唯一标识符 (UUID)';
|
||||||
|
COMMENT ON COLUMN sessions.user_id IS '所属用户 ID,关联 users 表';
|
||||||
|
COMMENT ON COLUMN sessions.title IS '会话标题,默认为"新对话",长度 1-100 字符';
|
||||||
|
COMMENT ON COLUMN sessions.config IS '会话配置 (JSONB),包含: tts_enabled (布尔), detail_level (''low''/''high''), language (语言代码), scenario (情景 ID)';
|
||||||
|
COMMENT ON COLUMN sessions.created_at IS '会话创建时间';
|
||||||
|
COMMENT ON COLUMN sessions.updated_at IS '会话最后更新时间';
|
||||||
|
|||||||
6
backend/migrations/004_user_scenarios.down.sql
Normal file
6
backend/migrations/004_user_scenarios.down.sql
Normal file
@@ -0,0 +1,6 @@
|
|||||||
|
-- 004_user_scenarios.down.sql
|
||||||
|
-- 回滚用户自建情景表
|
||||||
|
|
||||||
|
DROP INDEX IF EXISTS idx_user_scenarios_created_at;
|
||||||
|
DROP INDEX IF EXISTS idx_user_scenarios_user_id;
|
||||||
|
DROP TABLE IF EXISTS user_scenarios;
|
||||||
40
backend/migrations/004_user_scenarios.up.sql
Normal file
40
backend/migrations/004_user_scenarios.up.sql
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
-- 004_user_scenarios.up.sql
|
||||||
|
-- 用户自建情景表
|
||||||
|
|
||||||
|
CREATE TABLE user_scenarios (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
name VARCHAR(50) NOT NULL,
|
||||||
|
icon VARCHAR(20) DEFAULT '✨',
|
||||||
|
description VARCHAR(100),
|
||||||
|
prompt TEXT NOT NULL,
|
||||||
|
greeting VARCHAR(500),
|
||||||
|
language VARCHAR(10) DEFAULT 'zh-CN',
|
||||||
|
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||||
|
updated_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||||
|
|
||||||
|
CONSTRAINT unique_user_scenario UNIQUE(user_id, name),
|
||||||
|
CONSTRAINT check_name_length CHECK (char_length(name) >= 2 AND char_length(name) <= 50),
|
||||||
|
CONSTRAINT check_description_length CHECK (description IS NULL OR char_length(description) <= 100),
|
||||||
|
CONSTRAINT check_prompt_length CHECK (char_length(prompt) >= 10),
|
||||||
|
CONSTRAINT check_greeting_length CHECK (greeting IS NULL OR char_length(greeting) <= 500)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- 为用户 ID 创建索引,加速查询
|
||||||
|
CREATE INDEX idx_user_scenarios_user_id ON user_scenarios(user_id);
|
||||||
|
|
||||||
|
-- 为创建时间创建索引,用于排序
|
||||||
|
CREATE INDEX idx_user_scenarios_created_at ON user_scenarios(created_at DESC);
|
||||||
|
|
||||||
|
-- 表和列注释
|
||||||
|
COMMENT ON TABLE user_scenarios IS '用户自建情景表,存储用户创建的 AI 对话情景配置';
|
||||||
|
COMMENT ON COLUMN user_scenarios.id IS '情景唯一标识符 (UUID)';
|
||||||
|
COMMENT ON COLUMN user_scenarios.user_id IS '所属用户 ID,外键关联 users 表,用户删除时级联删除';
|
||||||
|
COMMENT ON COLUMN user_scenarios.name IS '情景名称 (2-50 字符),如"创意写作导师"';
|
||||||
|
COMMENT ON COLUMN user_scenarios.icon IS 'Emoji 图标 (最多 20 字符),支持复合 Emoji,如"🎨"';
|
||||||
|
COMMENT ON COLUMN user_scenarios.description IS '简短描述 (最多 100 字符),可选,显示在情景卡片上';
|
||||||
|
COMMENT ON COLUMN user_scenarios.prompt IS '角色 System Prompt (最少 10 字符,无上限),定义 AI 行为和对话风格';
|
||||||
|
COMMENT ON COLUMN user_scenarios.greeting IS '首句引导 (最多 500 字符),可选,AI 的开场白';
|
||||||
|
COMMENT ON COLUMN user_scenarios.language IS '默认语言代码 (如 zh-CN、en-US、ja-JP)';
|
||||||
|
COMMENT ON COLUMN user_scenarios.created_at IS '情景创建时间';
|
||||||
|
COMMENT ON COLUMN user_scenarios.updated_at IS '情景最后更新时间';
|
||||||
40
deploy.sh
40
deploy.sh
@@ -4,30 +4,52 @@ set -euo pipefail
|
|||||||
PROJECT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
PROJECT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||||
cd "$PROJECT_DIR"
|
cd "$PROJECT_DIR"
|
||||||
|
|
||||||
|
# .env 固定路径(act_runner 容器已挂载 /opt/camtalk)
|
||||||
|
ENV_FILE="/opt/camtalk/.env"
|
||||||
|
|
||||||
# 颜色输出
|
# 颜色输出
|
||||||
GREEN='\033[0;32m'
|
GREEN='\033[0;32m'
|
||||||
NC='\033[0m'
|
NC='\033[0m'
|
||||||
|
|
||||||
info() { echo -e "${GREEN}[INFO]${NC} $*"; }
|
info() { echo -e "${GREEN}[INFO]${NC} $*"; }
|
||||||
|
|
||||||
|
# .env 检查:首次部署时从 .env.example 复制模板,提示用户填写
|
||||||
|
check_env() {
|
||||||
|
if [ ! -f "$ENV_FILE" ]; then
|
||||||
|
echo "=============================================="
|
||||||
|
echo " 错误: 未找到环境变量文件"
|
||||||
|
echo " 路径: $ENV_FILE"
|
||||||
|
echo " 模板参考: backend/.env.example"
|
||||||
|
echo "=============================================="
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
check_env
|
||||||
|
|
||||||
|
# 所有 docker compose 命令统一使用 --env-file,用于解析 ${POSTGRES_USER} 等变量
|
||||||
|
DC="docker compose --env-file $ENV_FILE"
|
||||||
|
|
||||||
cmd_build() {
|
cmd_build() {
|
||||||
info "构建 Docker 镜像..."
|
info "构建 Docker 镜像..."
|
||||||
# 启用 BuildKit 加速构建
|
DOCKER_BUILDKIT=1 $DC build --parallel
|
||||||
DOCKER_BUILDKIT=1 docker compose build --parallel
|
|
||||||
info "构建完成"
|
info "构建完成"
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd_up() {
|
cmd_up() {
|
||||||
info "启动服务..."
|
info "启动服务..."
|
||||||
docker compose up -d
|
$DC up -d
|
||||||
info "服务已启动"
|
info "服务已启动"
|
||||||
info "前端: http://8.161.227.145:9000"
|
PUBLIC_IP=$(curl -s --connect-timeout 3 https://ifconfig.me 2>/dev/null || \
|
||||||
info "健康检查: http://8.161.227.145:9000/api/health"
|
curl -s --connect-timeout 3 https://api.ipify.org 2>/dev/null || \
|
||||||
|
echo "YOUR_SERVER_IP")
|
||||||
|
info "前端: http://$PUBLIC_IP:9000"
|
||||||
|
info "健康检查: http://$PUBLIC_IP:9000/api/health"
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd_down() {
|
cmd_down() {
|
||||||
info "停止服务..."
|
info "停止服务..."
|
||||||
docker compose down
|
$DC down
|
||||||
info "服务已停止"
|
info "服务已停止"
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -38,11 +60,11 @@ cmd_restart() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
cmd_logs() {
|
cmd_logs() {
|
||||||
docker compose logs -f "${@}"
|
$DC logs -f "${@}"
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd_status() {
|
cmd_status() {
|
||||||
docker compose ps
|
$DC ps
|
||||||
}
|
}
|
||||||
|
|
||||||
usage() {
|
usage() {
|
||||||
@@ -58,6 +80,8 @@ CamTalk 部署脚本
|
|||||||
restart 重启服务
|
restart 重启服务
|
||||||
logs 查看日志(可加服务名,如: $0 logs backend)
|
logs 查看日志(可加服务名,如: $0 logs backend)
|
||||||
status 查看服务状态
|
status 查看服务状态
|
||||||
|
|
||||||
|
.env 路径: $ENV_FILE
|
||||||
EOF
|
EOF
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -17,25 +17,31 @@ services:
|
|||||||
context: ./backend
|
context: ./backend
|
||||||
dockerfile: Dockerfile
|
dockerfile: Dockerfile
|
||||||
container_name: camtalk-backend
|
container_name: camtalk-backend
|
||||||
|
env_file:
|
||||||
|
- /opt/camtalk/.env
|
||||||
environment:
|
environment:
|
||||||
- APP_ENV=production
|
# 运行环境(强制生产环境)
|
||||||
- CAMTALK_STORAGE_DRIVER=postgres
|
- APP_ENV=prod
|
||||||
- CAMTALK_STORAGE_DSN=postgres://camtalk:camtalk123@postgres:5432/camtalk?sslmode=disable
|
# 三级存储配置(敏感信息通过 env_file 注入)
|
||||||
- CAMTALK_AUTH_JWT_SECRET=78uWBBAF8XEQEotKDlrnlnd4y8i4WN3E4zXmNmC8BYQ=
|
- CAMTALK_STORAGE_REDIS_ENABLED=${CAMTALK_STORAGE_REDIS_ENABLED:-true}
|
||||||
|
- CAMTALK_STORAGE_PERSISTENCE_ENABLED=${CAMTALK_STORAGE_PERSISTENCE_ENABLED:-true}
|
||||||
|
- CAMTALK_STORAGE_PERSISTENCE_DRIVER=${CAMTALK_STORAGE_PERSISTENCE_DRIVER:-postgres}
|
||||||
|
- CAMTALK_REDIS_ADDR=redis:6379
|
||||||
depends_on:
|
depends_on:
|
||||||
postgres:
|
postgres:
|
||||||
condition: service_healthy
|
condition: service_healthy
|
||||||
|
redis:
|
||||||
|
condition: service_healthy
|
||||||
networks:
|
networks:
|
||||||
- camtalk-net
|
- camtalk-net
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
|
|
||||||
postgres:
|
postgres:
|
||||||
# 轩辕镜像加速,避免 Docker Hub 拉取超时
|
|
||||||
image: docker.m.daocloud.io/library/postgres:15-alpine
|
image: docker.m.daocloud.io/library/postgres:15-alpine
|
||||||
container_name: camtalk-postgres
|
container_name: camtalk-postgres
|
||||||
|
env_file:
|
||||||
|
- /opt/camtalk/.env
|
||||||
environment:
|
environment:
|
||||||
POSTGRES_USER: camtalk
|
|
||||||
POSTGRES_PASSWORD: camtalk123
|
|
||||||
POSTGRES_DB: camtalk
|
POSTGRES_DB: camtalk
|
||||||
volumes:
|
volumes:
|
||||||
- pgdata:/var/lib/postgresql/data
|
- pgdata:/var/lib/postgresql/data
|
||||||
@@ -43,7 +49,30 @@ services:
|
|||||||
networks:
|
networks:
|
||||||
- camtalk-net
|
- camtalk-net
|
||||||
healthcheck:
|
healthcheck:
|
||||||
test: ["CMD-SHELL", "pg_isready -U camtalk -d camtalk"]
|
test: ["CMD-SHELL", "pg_isready -U $$POSTGRES_USER -d camtalk"]
|
||||||
|
interval: 5s
|
||||||
|
timeout: 3s
|
||||||
|
retries: 10
|
||||||
|
restart: unless-stopped
|
||||||
|
|
||||||
|
redis:
|
||||||
|
image: docker.m.daocloud.io/library/redis:7-alpine
|
||||||
|
container_name: camtalk-redis
|
||||||
|
env_file:
|
||||||
|
- /opt/camtalk/.env
|
||||||
|
command: >
|
||||||
|
sh -c '
|
||||||
|
if [ -n "$$CAMTALK_REDIS_PASSWORD" ]; then
|
||||||
|
exec redis-server --appendonly yes --requirepass "$$CAMTALK_REDIS_PASSWORD"
|
||||||
|
else
|
||||||
|
exec redis-server --appendonly yes
|
||||||
|
fi'
|
||||||
|
volumes:
|
||||||
|
- redisdata:/data
|
||||||
|
networks:
|
||||||
|
- camtalk-net
|
||||||
|
healthcheck:
|
||||||
|
test: ["CMD-SHELL", "if [ -n \"$$CAMTALK_REDIS_PASSWORD\" ]; then redis-cli -a \"$$CAMTALK_REDIS_PASSWORD\" ping; else redis-cli ping; fi"]
|
||||||
interval: 5s
|
interval: 5s
|
||||||
timeout: 3s
|
timeout: 3s
|
||||||
retries: 10
|
retries: 10
|
||||||
@@ -51,6 +80,7 @@ services:
|
|||||||
|
|
||||||
volumes:
|
volumes:
|
||||||
pgdata:
|
pgdata:
|
||||||
|
redisdata:
|
||||||
|
|
||||||
networks:
|
networks:
|
||||||
camtalk-net:
|
camtalk-net:
|
||||||
|
|||||||
369
docs/01-架构设计.md
Normal file
369
docs/01-架构设计.md
Normal file
@@ -0,0 +1,369 @@
|
|||||||
|
# 架构设计
|
||||||
|
|
||||||
|
## 项目概述
|
||||||
|
|
||||||
|
CamTalk 是一款**多模态实时 AI 视觉对话助手**。用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。
|
||||||
|
|
||||||
|
核心挑战在于三个维度之间的张力:
|
||||||
|
|
||||||
|
| 维度 | 关键问题 |
|
||||||
|
|------|---------|
|
||||||
|
| 视觉理解 | 如何准确理解摄像头画面中的人物、物体、场景? |
|
||||||
|
| 语音交互 | 如何让对话像真人交流一样自然、低延迟? |
|
||||||
|
| 成本控制 | 实时视频流 + LLM 推理,如何避免账单爆炸? |
|
||||||
|
|
||||||
|
## 系统架构
|
||||||
|
|
||||||
|
三层架构:**前端做轻量预处理,后端做智能编排,云端 AI 服务按需调用**。
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph TB
|
||||||
|
subgraph Browser["浏览器客户端"]
|
||||||
|
UI["UI 渲染层<br/>React 18 + TypeScript"]
|
||||||
|
Edge["边缘预处理层<br/>VAD / 关键帧检测"]
|
||||||
|
Media["媒体采集层<br/>Camera / Microphone"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Gateway["Go 网关"]
|
||||||
|
WS["WebSocket Handler<br/>连接管理 / 消息分发"]
|
||||||
|
Session["Session Manager<br/>会话状态 / 对话历史"]
|
||||||
|
Orch["AI Orchestrator<br/>Eino Graph 声明式编排"]
|
||||||
|
Auth["Auth 模块<br/>JWT / bcrypt"]
|
||||||
|
REST["REST API<br/>健康检查 / 对话管理"]
|
||||||
|
Store["Store 层<br/>Repository 接口"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph AI["云端 AI 服务"]
|
||||||
|
STT["STT<br/>Deepgram / MiMo ASR"]
|
||||||
|
LLM["LLM<br/>GPT-4o / 通义千问"]
|
||||||
|
TTS["TTS<br/>OpenAI TTS / MiMo TTS"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Storage["存储层"]
|
||||||
|
Mem["Memory<br/>进程内缓存"]
|
||||||
|
Redis["Redis<br/>会话状态"]
|
||||||
|
PG["PostgreSQL<br/>持久化存储"]
|
||||||
|
end
|
||||||
|
|
||||||
|
Media --> Edge
|
||||||
|
Edge -->|"query (image+audio)"| WS
|
||||||
|
UI <-->|"WebSocket"| WS
|
||||||
|
WS --> Session
|
||||||
|
WS --> Orch
|
||||||
|
Orch --> STT
|
||||||
|
Orch --> LLM
|
||||||
|
Orch --> TTS
|
||||||
|
Session --> Store
|
||||||
|
Store --> Mem
|
||||||
|
Store --> Redis
|
||||||
|
Store --> PG
|
||||||
|
REST --> Session
|
||||||
|
WS --> Auth
|
||||||
|
```
|
||||||
|
|
||||||
|
> 为什么单独加一层 Go 网关,而不是让前端直连 AI API?1)API Key 安全性;2)统一的速率限制和成本管控;3)多模型路由逻辑集中在一处便于维护。
|
||||||
|
|
||||||
|
## 核心交互流程
|
||||||
|
|
||||||
|
一次完整的"用户提问 → AI 回答"流程:
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
sequenceDiagram
|
||||||
|
participant B as 浏览器
|
||||||
|
participant G as Go 网关(Eino Graph)
|
||||||
|
participant S as STT
|
||||||
|
participant L as LLM(ChatModel)
|
||||||
|
participant T as TTS
|
||||||
|
|
||||||
|
B->>B: VAD 检测到语音结束
|
||||||
|
B->>G: query {image, audio}
|
||||||
|
Note over G: EinoOrchestrator 启动 Graph.Stream()
|
||||||
|
|
||||||
|
G->>S: STT Lambda:音频 → 文本
|
||||||
|
S-->>G: 识别文本
|
||||||
|
G-->>B: stt_result {text}
|
||||||
|
|
||||||
|
G->>G: History Lambda:组装提示词 + 历史 + 多模态消息
|
||||||
|
|
||||||
|
G->>L: ChatModel Node:流式推理
|
||||||
|
loop LLM 流式输出(Callback OnEndWithStreamOutput)
|
||||||
|
L-->>G: token delta
|
||||||
|
G-->>B: llm_chunk {delta}
|
||||||
|
end
|
||||||
|
|
||||||
|
G->>G: Msg2Str + Splitter Lambda:句子切分
|
||||||
|
G->>T: TTS Lambda:逐句合成
|
||||||
|
T-->>G: 音频 chunk
|
||||||
|
G-->>B: tts_audio {audio}
|
||||||
|
|
||||||
|
G->>G: Done Lambda:发送完成通知
|
||||||
|
G-->>B: llm_done {full_text, tokens}
|
||||||
|
G-->>B: tts_audio {final: true}
|
||||||
|
```
|
||||||
|
|
||||||
|
**关键优化**:Eino Graph 以 Stream 模式运行,ChatModel 的 token 流通过 Callback 的 `OnEndWithStreamOutput` 实时推送到客户端(`llm_chunk`),同时 Splitter 节点将 token 流切分为句子,TTS 节点逐句合成并推送音频。LLM 文本流和 TTS 音频流**并行推送**,用户感知延迟大幅降低。
|
||||||
|
|
||||||
|
## 技术栈
|
||||||
|
|
||||||
|
### 前端
|
||||||
|
|
||||||
|
| 技术 | 选型 | 选择理由 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 框架 | React 18 + TypeScript | 组件化开发,类型安全,生态成熟 |
|
||||||
|
| 构建 | Vite | 开发热更新快,构建产物小 |
|
||||||
|
| 实时通信 | WebSocket(原生 API) + 自封装连接管理 | 浏览器原生支持,封装心跳/重连/消息分发 |
|
||||||
|
| 语音检测 | @ricky0123/vad-web | 基于 WebRTC VAD,纯前端零延迟 |
|
||||||
|
| 媒体采集 | MediaDevices API | 浏览器原生摄像头/麦克风访问 |
|
||||||
|
|
||||||
|
### 后端
|
||||||
|
|
||||||
|
| 技术 | 选型 | 选择理由 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 语言 | Go | 高并发 goroutine 模型,适合长连接管理 |
|
||||||
|
| HTTP 框架 | Gin | 高性能 HTTP 路由,中间件生态成熟 |
|
||||||
|
| WebSocket | gorilla/websocket | Go 生态最成熟的 WebSocket 库 |
|
||||||
|
| 会话存储 | Memory / Redis / PostgreSQL 三级存储 | 进程内存零依赖,Redis 支持多实例,PG 持久化。TieredManager 自动降级 |
|
||||||
|
| AI 编排 | CloudWeGo Eino Graph | 声明式 DAG 编排,Stream 模式,Callback AOP |
|
||||||
|
| 持久化存储 | PostgreSQL | 对话历史、用户数据、会话元数据 |
|
||||||
|
| 配置管理 | Viper + godotenv | 支持 YAML + .env + 环境变量覆盖 |
|
||||||
|
| 日志 | Zap | 高性能结构化日志 |
|
||||||
|
|
||||||
|
### AI 服务
|
||||||
|
|
||||||
|
| 能力 | 默认方案 | 备选方案 |
|
||||||
|
|------|---------|---------|
|
||||||
|
| 多模态 LLM | DashScope qwen3-vl-plus | GPT-4o 等 OpenAI 兼容模型 |
|
||||||
|
| 语音识别 STT | MiMo ASR(小米) | Deepgram |
|
||||||
|
| 语音合成 TTS | MiMo TTS(小米) | OpenAI TTS |
|
||||||
|
|
||||||
|
> Go 网关的 AI 服务层统一封装不同服务商的调用接口,通过配置切换 provider。LLM 通过 Eino 框架的 `eino-ext/components/model/openai` 组件接入,支持任何 OpenAI 兼容接口。
|
||||||
|
|
||||||
|
## 后端模块
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph LR
|
||||||
|
subgraph Entry["入口层"]
|
||||||
|
Main["main.go<br/>依赖注入 / 启动"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Transport["传输层"]
|
||||||
|
WSH["WebSocket Handler<br/>连接管理 / 认证"]
|
||||||
|
APH["REST API Handlers<br/>Auth / Conversation / Health"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Business["业务层"]
|
||||||
|
SM["Session Manager<br/>会话生命周期"]
|
||||||
|
ORCH["EinoOrchestrator<br/>Eino Graph 编排"]
|
||||||
|
AS["Auth Service<br/>注册/登录/刷新/登出"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Eino_Layer["Eino 编排层"]
|
||||||
|
PG["PipelineGraph<br/>7 节点 DAG"]
|
||||||
|
CB["Callback Handler<br/>LLM token 推送"]
|
||||||
|
ST["PipelineState<br/>跨节点状态"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph AI_Layer["AI 服务层"]
|
||||||
|
STT_S["STT Service<br/>MiMo / Deepgram"]
|
||||||
|
LLM_S["ChatModel<br/>eino-ext OpenAI 兼容"]
|
||||||
|
TTS_S["TTS Service<br/>MiMo / OpenAI"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Data["数据层"]
|
||||||
|
UR["UserRepository"]
|
||||||
|
MR["MessageRepository"]
|
||||||
|
SR["SessionRepository"]
|
||||||
|
end
|
||||||
|
|
||||||
|
Main --> WSH
|
||||||
|
Main --> APH
|
||||||
|
Main --> SM
|
||||||
|
Main --> ORCH
|
||||||
|
Main --> AS
|
||||||
|
|
||||||
|
WSH --> SM
|
||||||
|
WSH --> ORCH
|
||||||
|
APH --> SM
|
||||||
|
APH --> AS
|
||||||
|
ORCH --> PG
|
||||||
|
PG --> CB
|
||||||
|
PG --> ST
|
||||||
|
PG --> STT_S
|
||||||
|
PG --> LLM_S
|
||||||
|
PG --> TTS_S
|
||||||
|
SM --> MR
|
||||||
|
SM --> SR
|
||||||
|
AS --> UR
|
||||||
|
```
|
||||||
|
|
||||||
|
| 模块 | 职责 |
|
||||||
|
|------|------|
|
||||||
|
| WebSocket Handler | 管理客户端连接生命周期,JWT 认证,conversation_id 恢复,单播消息推送 |
|
||||||
|
| Session Manager | 维护用户会话状态、对话历史,三级存储架构,30 分钟 TTL |
|
||||||
|
| Eino 编排层 | 基于 CloudWeGo Eino Graph 的声明式 AI 编排,7 节点 DAG 流水线,Stream 模式调用 |
|
||||||
|
| AI Orchestrator | EinoOrchestrator 适配器,包装 Eino Graph 实现 Orchestrator 接口 |
|
||||||
|
| AI Service Layer | AI 服务抽象层,多 provider 支持(Deepgram/MiMo/OpenAI 等) |
|
||||||
|
| Auth | 用户认证与授权,JWT 双 token 轮转,bcrypt 密码哈希 |
|
||||||
|
| Store | 持久化存储层,Repository 接口与实现(内存 + PostgreSQL) |
|
||||||
|
| REST API | 健康检查、认证、对话管理端点 |
|
||||||
|
| Logger | Zap 结构化日志 |
|
||||||
|
| Models | 数据模型定义 |
|
||||||
|
| Migrations | 数据库版本化迁移 |
|
||||||
|
| Model Router | 根据请求类型选择 AI 模型(待实现) |
|
||||||
|
| Rate Limiter | 令牌桶限流,详见 [11-令牌桶限流.md](./11-令牌桶限流.md) |
|
||||||
|
|
||||||
|
## 前端组件
|
||||||
|
|
||||||
|
| 组件 | 职责 |
|
||||||
|
|------|------|
|
||||||
|
| LandingPage | 未登录时的着陆页,内嵌 LoginModal 登录/注册弹窗 |
|
||||||
|
| CameraManager | 摄像头流采集 |
|
||||||
|
| MicManager | 麦克风音频采集 |
|
||||||
|
| EdgeProcessor | VAD + 关键帧检测 |
|
||||||
|
| WebSocketManager | WebSocket 连接生命周期管理 |
|
||||||
|
| ChatPanel | 消息展示、流式回复、文本输入、场景选择 |
|
||||||
|
| VideoPreview | 摄像头画面预览 |
|
||||||
|
| SessionSidebar | 左侧对话列表(搜索、重命名、删除、时间分组) |
|
||||||
|
| ConfigPanel | 右侧配置面板(主题、TTS 开关、detail level、语言、场景、登出) |
|
||||||
|
| Toast | 轻量通知提示 |
|
||||||
|
|
||||||
|
核心 Hook:`useVisionSession()` 封装完整的视觉对话会话(摄像头、VAD、WebSocket、消息状态、认证、场景模式)。`useSessionList()` 通过 REST API 管理对话列表 CRUD。
|
||||||
|
|
||||||
|
### 前端会话状态模型(三态)
|
||||||
|
|
||||||
|
前端 UI 存在三个会话状态,由 `isConnected` 和 `isCameraOn` 联合决定:
|
||||||
|
|
||||||
|
```
|
||||||
|
┌──────────┐ startSession() ┌──────────┐
|
||||||
|
│ initial │ ──────────────────→ │ video │
|
||||||
|
│ 初始态 │ │ 视频通话 │
|
||||||
|
└──────────┘ └──────────┘
|
||||||
|
↑ │
|
||||||
|
│ stopSession() stopVideo()
|
||||||
|
│ │
|
||||||
|
│ ▼
|
||||||
|
│ ┌──────────┐
|
||||||
|
└──────────────────────── │ textOnly │
|
||||||
|
│ 文字对话 │
|
||||||
|
└──────────┘
|
||||||
|
│
|
||||||
|
startSession()
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌──────────┐
|
||||||
|
│ video │
|
||||||
|
└──────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
| 状态 | 条件 | WebSocket | 摄像头 | 消息 | 文字输入 |
|
||||||
|
|------|------|-----------|--------|------|---------|
|
||||||
|
| `initial` | `!isConnected && messages.length === 0` | 断开 | 关闭 | 空 | 可用(自动连接) |
|
||||||
|
| `video` | `isConnected && isCameraOn` | 连接 | 开启 | 有 | 可用 |
|
||||||
|
| `textOnly` | `isConnected && !isCameraOn` | 连接 | 关闭 | 保留 | 可用 |
|
||||||
|
|
||||||
|
- **`stopVideo()`**:停止摄像头/麦克风/VAD,保持 WebSocket 连接和消息历史,用户可继续文字对话
|
||||||
|
- **`stopSession()`**:完全断开 WebSocket、清空消息、重置状态,回到初始态
|
||||||
|
|
||||||
|
## 数据库设计
|
||||||
|
|
||||||
|
### ER 关系
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
erDiagram
|
||||||
|
users ||--o{ sessions : "1:N"
|
||||||
|
users ||--o{ refresh_tokens : "1:N"
|
||||||
|
sessions ||--o{ messages : "1:N"
|
||||||
|
|
||||||
|
users {
|
||||||
|
uuid id PK
|
||||||
|
varchar username UK
|
||||||
|
varchar password_hash
|
||||||
|
timestamptz created_at
|
||||||
|
timestamptz updated_at
|
||||||
|
}
|
||||||
|
|
||||||
|
sessions {
|
||||||
|
uuid id PK
|
||||||
|
uuid user_id FK
|
||||||
|
varchar title
|
||||||
|
jsonb config
|
||||||
|
timestamptz created_at
|
||||||
|
timestamptz updated_at
|
||||||
|
}
|
||||||
|
|
||||||
|
messages {
|
||||||
|
bigserial id PK
|
||||||
|
uuid session_id FK
|
||||||
|
varchar role
|
||||||
|
text content
|
||||||
|
integer tokens_used
|
||||||
|
timestamptz created_at
|
||||||
|
}
|
||||||
|
|
||||||
|
refresh_tokens {
|
||||||
|
bigserial id PK
|
||||||
|
uuid user_id FK
|
||||||
|
varchar token_hash UK
|
||||||
|
timestamptz expires_at
|
||||||
|
timestamptz created_at
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
系统采用关系型数据库存储持久化数据,包括用户账户、对话会话、消息记录和刷新令牌。数据库表定义详见 `backend/migrations/` 目录下的 SQL 迁移文件。
|
||||||
|
|
||||||
|
### 存储策略
|
||||||
|
|
||||||
|
系统采用**三级存储架构**(TieredManager)实现会话状态管理,平衡性能与可靠性:
|
||||||
|
|
||||||
|
- **L1 Memory**:进程内缓存,提供微秒级读写性能
|
||||||
|
- **L2 Redis**:分布式缓存层,支持多实例部署,提供毫秒级访问
|
||||||
|
- **L3 PostgreSQL**:持久化存储层,确保数据可靠性
|
||||||
|
|
||||||
|
会话数据按 TTL(默认 30 分钟)在三级存储间流转,支持 Redis 故障时自动降级到 Memory + PostgreSQL 模式。配置灵活,可根据部署规模选择单级(Memory)、双级(Memory + PostgreSQL)或完整三级存储方案。
|
||||||
|
|
||||||
|
## 认证设计
|
||||||
|
|
||||||
|
系统采用 **JWT 双 token 轮转认证机制**,结合 bcrypt 密码哈希和 Refresh Token Rotation 安全策略。
|
||||||
|
|
||||||
|
核心机制包括:双 token 轮转(access_token 15 分钟 + refresh_token 7 天)、密码安全(bcrypt cost=10)、token 安全(SHA256 哈希存储、复用检测)、WebSocket 连接认证(基于 access_token 的 HTTP Upgrade 校验)等。认证流程、安全机制、配置要求等详细设计见 [10-鉴权体系.md](./10-鉴权体系.md)。
|
||||||
|
|
||||||
|
## 部署架构
|
||||||
|
|
||||||
|
系统采用分层部署架构,支持单实例和多实例水平扩展:
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph TB
|
||||||
|
User["用户浏览器"] --> Nginx
|
||||||
|
|
||||||
|
subgraph Nginx["Nginx 反向代理"]
|
||||||
|
Static["/ → 前端静态资源"]
|
||||||
|
API["/api/* → Go Gateway"]
|
||||||
|
WS_Proxy["/ws → Go Gateway"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Gateway_Pool["Go Gateway 实例"]
|
||||||
|
G1["Gateway-1"]
|
||||||
|
G2["Gateway-2"]
|
||||||
|
GN["Gateway-N"]
|
||||||
|
end
|
||||||
|
|
||||||
|
Nginx --> G1
|
||||||
|
Nginx --> G2
|
||||||
|
Nginx --> GN
|
||||||
|
|
||||||
|
G1 --> Redis
|
||||||
|
G2 --> Redis
|
||||||
|
GN --> Redis
|
||||||
|
|
||||||
|
G1 --> PG_DB["PostgreSQL"]
|
||||||
|
G2 --> PG_DB
|
||||||
|
GN --> PG_DB
|
||||||
|
|
||||||
|
G1 --> AI_Services["AI Services(外部 API)"]
|
||||||
|
G2 --> AI_Services
|
||||||
|
GN --> AI_Services
|
||||||
|
```
|
||||||
|
|
||||||
|
**跨域策略**:Nginx 将前端(`/`)、REST API(`/api/*`)、WebSocket(`/ws`)统一反代到同一域名,浏览器无跨域问题。
|
||||||
|
|
||||||
|
**开发环境**:前端 Vite :5173 通过 `server.proxy` 转发 `/ws` 和 `/api` 到后端 :8080,无需硬编码端口。
|
||||||
@@ -1,28 +0,0 @@
|
|||||||
# 项目概述
|
|
||||||
|
|
||||||
## 概述
|
|
||||||
|
|
||||||
开发一款**多模态实时对话应用**——通过摄像头与麦克风捕获用户的视觉场景与语音输入,由 AI 理解并给出自然、流畅的回应。
|
|
||||||
|
|
||||||
核心挑战在于三个维度之间的张力:
|
|
||||||
|
|
||||||
| 维度 | 关键问题 | 详见 |
|
|
||||||
|------|---------|------|
|
|
||||||
| 视觉理解 | 如何准确理解摄像头画面中的人物、物体、场景? | `07-视觉理解.md` |
|
|
||||||
| 语音交互 | 如何让对话像真人交流一样自然、低延迟? | `06-语音交互.md` |
|
|
||||||
| 成本控制 | 实时视频流 + LLM 推理,如何避免账单爆炸? | `08-成本控制.md` |
|
|
||||||
|
|
||||||
> 提升视觉精度意味着更高分辨率和更频繁的采样,但这会直接推高带宽和推理成本。架构设计需要在三者之间做好取舍。
|
|
||||||
|
|
||||||
## 项目目标
|
|
||||||
|
|
||||||
1. **用户故事规划**:明确"AI 能看、能听、能说"需要覆盖哪些场景 → `05-用户故事.md`
|
|
||||||
2. **成本控制策略**:从架构设计层面融入运营成本意识 → `08-成本控制.md`
|
|
||||||
|
|
||||||
## 交付物
|
|
||||||
|
|
||||||
- 可运行的应用程序(摄像头 + 麦克风 → AI 回应)
|
|
||||||
- 设计文档,覆盖:
|
|
||||||
- 计划实现 vs 最终实现的用户故事
|
|
||||||
- 成本控制技巧的构思 vs 实际采用的方案
|
|
||||||
- 项目架构设计与技术选型
|
|
||||||
557
docs/02-接口文档.md
Normal file
557
docs/02-接口文档.md
Normal file
@@ -0,0 +1,557 @@
|
|||||||
|
# 接口文档
|
||||||
|
|
||||||
|
## 概述
|
||||||
|
|
||||||
|
前后端通信接口契约。WebSocket 承载实时对话,REST API 支撑基础运维。
|
||||||
|
|
||||||
|
**设计原则**:
|
||||||
|
- WebSocket 为主:所有对话数据走 WebSocket
|
||||||
|
- REST 为辅:仅用于健康检查、认证、对话管理等低频操作
|
||||||
|
- 接口先行:先定义契约,再填充实现——前后端可并行开发
|
||||||
|
|
||||||
|
## 接口全景
|
||||||
|
|
||||||
|
```
|
||||||
|
浏览器 Go Gateway :8080
|
||||||
|
WebSocket Client <--> /ws?token=<jwt> (实时对话,需 JWT 认证)
|
||||||
|
HTTP Client --> GET /api/health (健康检查)
|
||||||
|
HTTP Client <--> POST /api/auth/* (注册/登录/刷新/登出)
|
||||||
|
HTTP Client <--> GET/POST/PATCH/DELETE (对话 CRUD)
|
||||||
|
/api/conversations/*
|
||||||
|
HTTP Client <--> GET /api/conversations/:id (历史消息)
|
||||||
|
/messages
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 一、WebSocket 协议
|
||||||
|
|
||||||
|
连接地址:`ws://localhost:8080/ws?token=<access_token>&conversation_id=<uuid>`
|
||||||
|
|
||||||
|
| 参数 | 必填 | 说明 |
|
||||||
|
|------|------|------|
|
||||||
|
| `token` | 是 | JWT access_token,缺失或无效时返回 401 |
|
||||||
|
| `conversation_id` | 否 | 恢复已有对话;省略则创建新对话 |
|
||||||
|
|
||||||
|
### 消息格式约定
|
||||||
|
|
||||||
|
所有 WebSocket 消息均为 JSON 文本帧,统一结构:
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface WsMessage {
|
||||||
|
type: string; // 消息类型,必填
|
||||||
|
request_id?: string; // 可选,用于请求-响应关联
|
||||||
|
timestamp?: number; // 可选,毫秒时间戳
|
||||||
|
[key: string]: any; // 类型特定字段
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 客户端 → 服务端消息
|
||||||
|
|
||||||
|
#### `query` — 发起一次视觉对话
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface QueryMessage {
|
||||||
|
type: "query";
|
||||||
|
request_id: string; // 客户端生成的 UUID
|
||||||
|
image: string; // Base64 编码的 JPEG 图像(不含 data: 前缀)
|
||||||
|
audio: string; // Base64 编码的音频片段(PCM 16kHz),文本输入时为空字符串
|
||||||
|
text?: string; // 用户手动输入的文本(有值时跳过 STT,直接使用此文本)
|
||||||
|
mime_type?: string; // 音频格式,默认 "audio/pcm"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `config` — 更新会话配置
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface ConfigMessage {
|
||||||
|
type: "config";
|
||||||
|
payload: {
|
||||||
|
tts_enabled?: boolean; // 是否开启语音合成,默认 true
|
||||||
|
detail_level?: "low" | "high"; // 图像精度,默认 "low"
|
||||||
|
language?: string; // 交互语言,默认 "zh-CN"
|
||||||
|
scenario?: string; // 场景模式:free_chat / interviewer / english_teacher / debate / interpreter
|
||||||
|
};
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `interrupt` — 打断当前回复
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface InterruptMessage {
|
||||||
|
type: "interrupt";
|
||||||
|
request_id?: string; // 可选,当前实现不使用此字段,服务端始终取消当前活跃请求
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `ping` — 心跳保活
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface PingMessage {
|
||||||
|
type: "ping";
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 服务端 → 客户端消息
|
||||||
|
|
||||||
|
#### `connected` — 连接建立确认
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface ConnectedMessage {
|
||||||
|
type: "connected";
|
||||||
|
session_id: string;
|
||||||
|
conversation_id: string;
|
||||||
|
config: {
|
||||||
|
tts_enabled: boolean;
|
||||||
|
detail_level: "low" | "high";
|
||||||
|
language: string;
|
||||||
|
scenario: string;
|
||||||
|
};
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `stt_result` — 语音识别结果
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface SttResultMessage {
|
||||||
|
type: "stt_result";
|
||||||
|
request_id: string;
|
||||||
|
text: string; // 识别出的文本
|
||||||
|
is_final: boolean; // 当前实现始终为 true
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `llm_chunk` — LLM 流式响应片段
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface LlmChunkMessage {
|
||||||
|
type: "llm_chunk";
|
||||||
|
request_id: string;
|
||||||
|
content: string; // 当前 token 片段
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `llm_done` — LLM 响应完成
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface LlmDoneMessage {
|
||||||
|
type: "llm_done";
|
||||||
|
request_id: string;
|
||||||
|
full_text: string; // 完整响应文本
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `tts_audio` — TTS 音频片段
|
||||||
|
|
||||||
|
音频流式推送,每个消息携带一个句子的音频数据。
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface TtsAudioMessage {
|
||||||
|
type: "tts_audio";
|
||||||
|
request_id: string;
|
||||||
|
audio: string; // Base64 编码的音频数据
|
||||||
|
format: string; // 音频格式
|
||||||
|
sample_rate: number; // 采样率(Hz)
|
||||||
|
sequence: number; // 句子序号,从 0 开始递增
|
||||||
|
is_final: boolean; // 是否为最后一个句子
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**音频格式约束**:
|
||||||
|
|
||||||
|
| 字段 | 值 | 说明 |
|
||||||
|
|------|-----|------|
|
||||||
|
| `format` | `"pcm"` | 线性 PCM,小端序 |
|
||||||
|
| `sample_rate` | `24000` | 24kHz 采样率 |
|
||||||
|
| 位深度 | 16-bit | 单声道 |
|
||||||
|
| 句子划分 | 按标点符号(。!?;:)分割 | 服务端按句分割 LLM 响应,并行合成 |
|
||||||
|
|
||||||
|
#### `error` — 错误通知
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface ErrorMessage {
|
||||||
|
type: "error";
|
||||||
|
request_id?: string; // 关联的请求 ID,全局错误时为空
|
||||||
|
code: string; // 错误码,见下文错误码表
|
||||||
|
message: string; // 人类可读的错误描述
|
||||||
|
details?: any; // 可选的详细错误信息
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `pong` — 心跳响应
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface PongMessage {
|
||||||
|
type: "pong";
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 连接管理
|
||||||
|
|
||||||
|
- **心跳机制**:客户端每 30 秒发送 `ping`,服务端回复 `pong`;60 秒无活动则服务端断开连接
|
||||||
|
- **重连策略**:客户端断线后指数退避重连(1s → 2s → 4s → ... → 最大 30s)
|
||||||
|
- **并发控制**:同一连接同时只能有一个活跃的 `query` 请求;新请求到来时自动取消旧请求
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 二、REST API
|
||||||
|
|
||||||
|
所有 REST 端点均使用 JSON 格式。
|
||||||
|
|
||||||
|
### 2.1 健康检查
|
||||||
|
|
||||||
|
#### `GET /api/health`
|
||||||
|
|
||||||
|
检查服务健康状态。
|
||||||
|
|
||||||
|
**响应**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"status": "healthy",
|
||||||
|
"timestamp": "2024-01-15T10:30:00Z",
|
||||||
|
"dependencies": {
|
||||||
|
"database": "healthy",
|
||||||
|
"redis": "healthy"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.2 认证 API
|
||||||
|
|
||||||
|
#### `POST /api/auth/register` — 用户注册
|
||||||
|
|
||||||
|
**请求**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"username": "alice",
|
||||||
|
"email": "alice@example.com",
|
||||||
|
"password": "SecurePass123!"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"user": {
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"username": "alice",
|
||||||
|
"email": "alice@example.com",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z"
|
||||||
|
},
|
||||||
|
"access_token": "eyJhbGc...",
|
||||||
|
"refresh_token": "eyJhbGc...",
|
||||||
|
"expires_in": 7200
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**错误**:
|
||||||
|
- `400 INVALID_INPUT`: 参数验证失败
|
||||||
|
- `409 USER_EXISTS`: 用户名或邮箱已存在
|
||||||
|
|
||||||
|
#### `POST /api/auth/login` — 用户登录
|
||||||
|
|
||||||
|
**请求**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"username": "alice",
|
||||||
|
"password": "SecurePass123!"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"user": {
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"username": "alice",
|
||||||
|
"email": "alice@example.com"
|
||||||
|
},
|
||||||
|
"access_token": "eyJhbGc...",
|
||||||
|
"refresh_token": "eyJhbGc...",
|
||||||
|
"expires_in": 7200
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**错误**:
|
||||||
|
- `400 INVALID_INPUT`: 参数缺失
|
||||||
|
- `401 INVALID_CREDENTIALS`: 用户名或密码错误
|
||||||
|
|
||||||
|
#### `POST /api/auth/refresh` — 刷新 Access Token
|
||||||
|
|
||||||
|
**请求头**:
|
||||||
|
```
|
||||||
|
Authorization: Bearer <refresh_token>
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"access_token": "eyJhbGc...",
|
||||||
|
"refresh_token": "eyJhbGc...",
|
||||||
|
"expires_in": 7200
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**错误**:
|
||||||
|
- `401 INVALID_TOKEN`: Refresh Token 无效或过期
|
||||||
|
|
||||||
|
#### `POST /api/auth/logout` — 用户登出
|
||||||
|
|
||||||
|
**请求头**:
|
||||||
|
```
|
||||||
|
Authorization: Bearer <access_token>
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"message": "Logged out successfully"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.3 对话管理 API
|
||||||
|
|
||||||
|
所有端点均需 JWT 认证(`Authorization: Bearer <access_token>`)。
|
||||||
|
|
||||||
|
#### `GET /api/conversations` — 获取对话列表
|
||||||
|
|
||||||
|
**查询参数**:
|
||||||
|
- `page`: 页码,从 1 开始,默认 1
|
||||||
|
- `page_size`: 每页条数,默认 20,最大 100
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"conversations": [
|
||||||
|
{
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"title": "关于植物的对话",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z",
|
||||||
|
"updated_at": "2024-01-15T11:45:00Z",
|
||||||
|
"message_count": 12
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"total": 42,
|
||||||
|
"page": 1,
|
||||||
|
"page_size": 20
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `POST /api/conversations` — 创建新对话
|
||||||
|
|
||||||
|
**请求**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"title": "新的对话"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(201 Created):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"title": "新的对话",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z",
|
||||||
|
"updated_at": "2024-01-15T10:30:00Z",
|
||||||
|
"message_count": 0
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `GET /api/conversations/:id` — 获取对话详情
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"title": "关于植物的对话",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z",
|
||||||
|
"updated_at": "2024-01-15T11:45:00Z",
|
||||||
|
"message_count": 12
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**错误**:
|
||||||
|
- `404 NOT_FOUND`: 对话不存在或无权访问
|
||||||
|
|
||||||
|
#### `PATCH /api/conversations/:id` — 更新对话
|
||||||
|
|
||||||
|
**请求**:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"title": "修改后的标题"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"title": "修改后的标题",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z",
|
||||||
|
"updated_at": "2024-01-15T12:00:00Z",
|
||||||
|
"message_count": 12
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `DELETE /api/conversations/:id` — 删除对话
|
||||||
|
|
||||||
|
**响应**(204 No Content):无响应体
|
||||||
|
|
||||||
|
**错误**:
|
||||||
|
- `404 NOT_FOUND`: 对话不存在或无权访问
|
||||||
|
|
||||||
|
#### `GET /api/conversations/:id/messages` — 获取对话消息
|
||||||
|
|
||||||
|
**查询参数**:
|
||||||
|
- `page`: 页码,从 1 开始,默认 1
|
||||||
|
- `page_size`: 每页条数,默认 50,最大 100
|
||||||
|
|
||||||
|
**响应**(200 OK):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"messages": [
|
||||||
|
{
|
||||||
|
"id": "660e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"conversation_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"role": "user",
|
||||||
|
"content": "这是什么植物?",
|
||||||
|
"image_url": "/api/images/abc123.jpg",
|
||||||
|
"created_at": "2024-01-15T10:30:00Z"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "770e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"conversation_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"role": "assistant",
|
||||||
|
"content": "这是一株向日葵...",
|
||||||
|
"created_at": "2024-01-15T10:30:15Z"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"total": 12,
|
||||||
|
"page": 1,
|
||||||
|
"page_size": 50
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**消息字段说明**:
|
||||||
|
- `role`: `"user"` 或 `"assistant"`
|
||||||
|
- `image_url`: 仅 `user` 消息可能包含,指向存储的图像
|
||||||
|
- `content`: 消息文本内容
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 三、错误码表
|
||||||
|
|
||||||
|
所有错误均使用以下格式:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"code": "ERROR_CODE",
|
||||||
|
"message": "Human-readable error description",
|
||||||
|
"details": {}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### WebSocket 错误码
|
||||||
|
|
||||||
|
| 错误码 | 说明 | HTTP 状态码(若适用)|
|
||||||
|
|--------|------|---------------------|
|
||||||
|
| `INVALID_MESSAGE` | 消息格式错误或缺少必填字段 | - |
|
||||||
|
| `SESSION_NOT_FOUND` | 会话不存在 | - |
|
||||||
|
| `RATE_LIMITED` | 请求频率过高 | 429 |
|
||||||
|
| `IMAGE_TOO_LARGE` | 图像超过大小限制(5MB)| - |
|
||||||
|
| `AUDIO_TOO_LARGE` | 音频超过大小限制(10MB)| - |
|
||||||
|
| `STT_ERROR` | 语音识别服务错误 | - |
|
||||||
|
| `LLM_ERROR` | LLM 服务错误 | - |
|
||||||
|
| `LLM_TIMEOUT` | LLM 响应超时(60 秒)| - |
|
||||||
|
| `TTS_ERROR` | 语音合成服务错误 | - |
|
||||||
|
| `CONCURRENT_REQUEST` | 同一连接已有进行中的请求 | - |
|
||||||
|
| `INTERNAL_ERROR` | 服务器内部错误 | 500 |
|
||||||
|
|
||||||
|
### REST API 错误码
|
||||||
|
|
||||||
|
| 错误码 | 说明 | HTTP 状态码 |
|
||||||
|
|--------|------|-------------|
|
||||||
|
| `INVALID_INPUT` | 请求参数验证失败 | 400 |
|
||||||
|
| `INVALID_TOKEN` | JWT Token 无效或过期 | 401 |
|
||||||
|
| `INVALID_CREDENTIALS` | 用户名或密码错误 | 401 |
|
||||||
|
| `UNAUTHORIZED` | 未认证或认证失败 | 401 |
|
||||||
|
| `FORBIDDEN` | 无权访问资源 | 403 |
|
||||||
|
| `NOT_FOUND` | 资源不存在 | 404 |
|
||||||
|
| `USER_EXISTS` | 用户名或邮箱已存在 | 409 |
|
||||||
|
| `RATE_LIMITED` | 请求频率过高 | 429 |
|
||||||
|
| `INTERNAL_ERROR` | 服务器内部错误 | 500 |
|
||||||
|
| `SERVICE_UNAVAILABLE` | 依赖服务不可用 | 503 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 四、数据模型
|
||||||
|
|
||||||
|
### 用户(User)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface User {
|
||||||
|
id: string; // UUID
|
||||||
|
username: string; // 用户名,唯一
|
||||||
|
email: string; // 邮箱,唯一
|
||||||
|
created_at: string; // ISO 8601 时间戳
|
||||||
|
updated_at: string; // ISO 8601 时间戳
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 对话(Conversation)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface Conversation {
|
||||||
|
id: string; // UUID
|
||||||
|
user_id: string; // 所属用户 ID
|
||||||
|
title: string; // 对话标题
|
||||||
|
created_at: string; // ISO 8601 时间戳
|
||||||
|
updated_at: string; // ISO 8601 时间戳
|
||||||
|
message_count: number; // 消息数量
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 消息(Message)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
interface Message {
|
||||||
|
id: string; // UUID
|
||||||
|
conversation_id: string; // 所属对话 ID
|
||||||
|
role: "user" | "assistant";
|
||||||
|
content: string; // 消息文本内容
|
||||||
|
image_url?: string; // 可选,用户消息的关联图像 URL
|
||||||
|
created_at: string; // ISO 8601 时间戳
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### JWT Token 载荷
|
||||||
|
|
||||||
|
**Access Token**(有效期 120 分钟):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"user_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"username": "alice",
|
||||||
|
"type": "access",
|
||||||
|
"exp": 1705318200,
|
||||||
|
"iat": 1705311000
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**Refresh Token**(有效期 7 天):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"user_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||||
|
"type": "refresh",
|
||||||
|
"exp": 1705915800,
|
||||||
|
"iat": 1705311000
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 附录:版本历史
|
||||||
|
|
||||||
|
- **v1.0**(2024-01-15):初始版本,定义 WebSocket 协议和 REST API
|
||||||
|
- **v1.1**(2024-01-20):新增文本输入模式(`query.text` 字段)
|
||||||
|
- **v1.2**(2024-01-25):新增场景模式配置(`config.scenario` 字段)
|
||||||
|
- **v2.0**(2026-06-21):重构为纯接口契约规范,移除实现细节
|
||||||
254
docs/02-系统架构.md
254
docs/02-系统架构.md
@@ -1,254 +0,0 @@
|
|||||||
# 系统架构
|
|
||||||
|
|
||||||
## 概述
|
|
||||||
|
|
||||||
三层架构:**前端做轻量预处理,后端做智能编排,云端 AI 服务按需调用**。在保证交互体验的同时控制成本。
|
|
||||||
|
|
||||||
## 三层架构
|
|
||||||
|
|
||||||
| 层级 | 职责 | 关键约束 |
|
|
||||||
|------|------|---------|
|
|
||||||
| **客户端(浏览器)** | 媒体采集、边缘预处理、UI 渲染 | 浏览器资源有限,模型需轻量 |
|
|
||||||
| **Go 网关** | 会话管理、AI 服务编排、流式管道 | 高并发、低延迟、状态管理 |
|
|
||||||
| **AI 服务** | LLM 推理、语音识别、语音合成 | 按量计费,需控制调用频率 |
|
|
||||||
|
|
||||||
> 为什么要单独加一层 Go 网关,而不是让前端直连 AI API?1)API Key 安全性;2)统一的速率限制和成本管控;3)多模型路由逻辑集中在一处便于维护。
|
|
||||||
|
|
||||||
## 技术栈
|
|
||||||
|
|
||||||
### 前端
|
|
||||||
|
|
||||||
| 技术 | 选型 | 选择理由 |
|
|
||||||
|------|------|---------|
|
|
||||||
| 框架 | React 18 + TypeScript | 组件化开发,类型安全,生态成熟 |
|
|
||||||
| 构建 | Vite | 开发热更新快,构建产物小 |
|
|
||||||
| 实时通信 | WebSocket(原生 API) + 自封装连接管理 | 浏览器原生支持,封装心跳/重连/消息分发 |
|
|
||||||
| 边缘推理 | ONNX Runtime Web | 浏览器端跑轻量模型(VAD、关键帧检测) |
|
|
||||||
| 语音检测 | @ricky0123/vad-web | 基于 WebRTC VAD,纯前端零延迟 |
|
|
||||||
| 媒体采集 | MediaDevices API | 浏览器原生摄像头/麦克风访问 |
|
|
||||||
|
|
||||||
### 后端
|
|
||||||
|
|
||||||
| 技术 | 选型 | 选择理由 |
|
|
||||||
|------|------|---------|
|
|
||||||
| 语言 | Go | 高并发 goroutine 模型,适合长连接管理 |
|
|
||||||
| HTTP 框架 | Gin | 高性能 HTTP 路由,中间件生态成熟 |
|
|
||||||
| WebSocket | gorilla/websocket | Go 生态最成熟的 WebSocket 库 |
|
|
||||||
| 会话存储 | Redis(规划中) / Memory(MVP 默认) | 高速 KV 存储,MVP 阶段使用进程内存,可通过配置切换到 Redis |
|
|
||||||
| 持久化存储 | PostgreSQL(规划中) | 对话历史、用量统计、用户偏好(MVP 阶段未实现) |
|
|
||||||
| 配置管理 | Viper + godotenv | 支持 YAML + .env + 环境变量覆盖,详见 `03-接口文档.md` 第六章 |
|
|
||||||
| 日志 | Zap | 高性能结构化日志 |
|
|
||||||
|
|
||||||
### AI 服务
|
|
||||||
|
|
||||||
| 能力 | 主选方案 | 备选方案 | 选型考量 |
|
|
||||||
|------|---------|---------|---------|
|
|
||||||
| 多模态 LLM | GPT-4o(默认) | 通义千问等 OpenAI 兼容模型 | 通过 OpenAI 兼容接口,可灵活切换 |
|
|
||||||
| 语音识别 STT | Deepgram(默认) | MiMo ASR(小米) | 支持多 provider 切换 |
|
|
||||||
| 语音合成 TTS | OpenAI TTS(默认) | MiMo TTS(小米) | 支持多 provider 切换 |
|
|
||||||
|
|
||||||
> 不必绑定单一厂商。Go 网关的 AI 服务层统一封装不同服务商的调用接口,通过配置切换 provider。
|
|
||||||
|
|
||||||
## 核心交互流程
|
|
||||||
|
|
||||||
一次完整的"用户提问 → AI 回答"流程:
|
|
||||||
|
|
||||||
```
|
|
||||||
Browser Go Gateway STT LLM TTS
|
|
||||||
| | | | |
|
|
||||||
|-- VAD 检测到语音结束 --->| | | |
|
|
||||||
| | | | |
|
|
||||||
|-- [音频+图像] -------->| | | |
|
|
||||||
| |--- 音频流 ------->| | |
|
|
||||||
| |<-- 流式文本 ------| | |
|
|
||||||
| | | | |
|
|
||||||
| |--- [图像+文本+上下文] -------->| |
|
|
||||||
| |<-- 流式回答文本 --------------| |
|
|
||||||
|<-- 推送回答文本 --------| | | |
|
|
||||||
| |--- 回答文本 ---------------------------->|
|
|
||||||
| |<-- 流式音频 --------------------------------|
|
|
||||||
|<-- 推送音频流 ----------| | | |
|
|
||||||
| | | | |
|
|
||||||
|-> 播放音频 + 渲染文字 | | | |
|
|
||||||
```
|
|
||||||
|
|
||||||
**关键优化**:LLM 文本流和 TTS 音频流是**并行推送**的——客户端先展示文字,同时开始播放语音,用户感知延迟大幅降低。
|
|
||||||
|
|
||||||
## 后端模块
|
|
||||||
|
|
||||||
| 模块 | 职责 | 关键实现 |
|
|
||||||
|------|------|---------|
|
|
||||||
| WebSocket Handler | 管理客户端连接生命周期,单播消息推送 | goroutine per connection |
|
|
||||||
| Session Manager | 维护用户会话状态、对话历史 | Memory(MVP 默认)/ Redis(可切换),30 分钟 TTL(详见 `03-接口文档.md` 第五章) |
|
|
||||||
| AI Orchestrator | 编排 STT→LLM→TTS 流式并行管道 | context 取消 + 超时控制 + 句子切分 |
|
|
||||||
| AI Service Layer | AI 服务抽象层(STT/LLM/TTS) | 多 provider 支持(Deepgram/MiMo/OpenAI 等) |
|
|
||||||
| REST API | 健康检查、会话管理端点 | Gin 路由 |
|
|
||||||
| Error Handler | 统一错误码定义与发送 | 错误码枚举 |
|
|
||||||
| Logger | 日志初始化封装 | Zap 结构化日志 |
|
|
||||||
| Models | 数据模型定义 | WebSocket 消息、会话、配置等 |
|
|
||||||
| Model Router | 根据请求类型选择 AI 模型(规划中) | 规则引擎 + 成本阈值 |
|
|
||||||
| Rate Limiter | 防止单用户过度消耗 API 额度(规划中) | 令牌桶算法 |
|
|
||||||
|
|
||||||
AI Orchestrator 核心接口(`internal/orchestrator/orchestrator.go`):
|
|
||||||
|
|
||||||
```go
|
|
||||||
// Orchestrator AI 编排器接口。
|
|
||||||
type Orchestrator interface {
|
|
||||||
ProcessQuery(ctx context.Context, sessionID string, req models.WsQuery,
|
|
||||||
history []models.Message, sender Sender) error
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
Pipeline 实现(`internal/orchestrator/pipeline.go`)流程:
|
|
||||||
1. Base64 解码音频/图片
|
|
||||||
2. 调用 `stt.Recognize()` → 发送 `stt_result`
|
|
||||||
3. 调用 `llm.ChatStream()` 获取流式输出,goroutine 消费 token → 发送 `llm_chunk` + 句子切分
|
|
||||||
4. 另一 goroutine 从句子 channel 读取 → 调用 `tts.SynthesizeStream()` → 发送 `tts_audio`
|
|
||||||
5. 流结束 → 发送 `llm_done`
|
|
||||||
6. TTS 失败静默跳过,STT/LLM 失败发送对应 error 消息
|
|
||||||
|
|
||||||
> **关键优化**:LLM 文本流和 TTS 音频流**并行推送**——客户端先逐 token 展示文字,同时 TTS 逐句子合成并推送音频,用户感知延迟大幅降低。详细的 AI 服务层接口和编排策略见 `03-接口文档.md` 第三、四章。
|
|
||||||
|
|
||||||
## 前端组件
|
|
||||||
|
|
||||||
| 组件 | 职责 |
|
|
||||||
|------|------|
|
|
||||||
| CameraManager | 摄像头流采集 |
|
|
||||||
| MicManager | 麦克风音频采集 |
|
|
||||||
| EdgeProcessor | VAD + 关键帧检测(Canvas 像素比较) |
|
|
||||||
| WebSocketManager | WS 连接生命周期管理 |
|
|
||||||
| ChatPanel | 消息展示 |
|
|
||||||
| VideoPreview | 摄像头画面预览 |
|
|
||||||
| ConfigPanel | 右侧抽屉式配置面板(主题、TTS 开关、detail level、语言) |
|
|
||||||
| Toast | 轻量通知提示(3 秒自动消失) |
|
|
||||||
|
|
||||||
核心 Hook:`useVisionSession()` 封装一次完整的视觉对话会话(摄像头、VAD、WebSocket、消息状态)。
|
|
||||||
|
|
||||||
```typescript
|
|
||||||
function useVisionSession() {
|
|
||||||
const [messages, setMessages] = useState<Message[]>([]);
|
|
||||||
const wsRef = useWebSocket(`${window.location.protocol === "https:" ? "wss:" : "ws:"}//${window.location.host}/ws`);
|
|
||||||
const videoRef = useRef<HTMLVideoElement>(null);
|
|
||||||
const { captureFrame } = useCamera(videoRef);
|
|
||||||
|
|
||||||
const { isSpeaking } = useVAD({
|
|
||||||
onSpeechEnd: async (audio) => {
|
|
||||||
const frame = captureFrame();
|
|
||||||
wsRef.current?.send(JSON.stringify({
|
|
||||||
type: "query",
|
|
||||||
image: frame.toDataURL("image/jpeg", 0.7),
|
|
||||||
audio: encodeAudio(audio)
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
useEffect(() => {
|
|
||||||
wsRef.current?.on("message", (data) => {
|
|
||||||
const { text, audio } = JSON.parse(data);
|
|
||||||
setMessages(prev => [...prev, { role: "assistant", text }]);
|
|
||||||
if (audio) playAudio(audio);
|
|
||||||
});
|
|
||||||
}, []);
|
|
||||||
|
|
||||||
return { messages, videoRef, isSpeaking };
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
## 存储策略(分阶段)
|
|
||||||
|
|
||||||
| 阶段 | 存储方案 | 持久化内容 | 理由 |
|
|
||||||
|------|---------|-----------|------|
|
|
||||||
| MVP | Memory(进程内) | 无 | 快速验证核心功能,重启丢数据可接受。Redis 实现已就绪,可通过 `storage.driver` 配置切换 |
|
|
||||||
| 上线 | Redis + PostgreSQL | 对话历史、用户偏好、用量统计 | 用户需要查看历史,运营需要成本数据 |
|
|
||||||
| 规模化 | Redis + PG + 对象存储 | 图像帧、音频片段归档 | 大文件不适合存关系库 |
|
|
||||||
|
|
||||||
冷热分离:Redis 存"热数据"(当前对话上下文,微秒级读写),PostgreSQL 存"冷数据"(历史记录)。
|
|
||||||
|
|
||||||
### PostgreSQL 表设计
|
|
||||||
|
|
||||||
```sql
|
|
||||||
CREATE TABLE sessions (
|
|
||||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
|
||||||
user_id UUID NOT NULL,
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now(),
|
|
||||||
updated_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE TABLE messages (
|
|
||||||
id BIGSERIAL PRIMARY KEY,
|
|
||||||
session_id UUID REFERENCES sessions(id),
|
|
||||||
role VARCHAR(16) NOT NULL, -- "user" | "assistant"
|
|
||||||
content TEXT NOT NULL,
|
|
||||||
image_url TEXT,
|
|
||||||
tokens_used INTEGER DEFAULT 0,
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE TABLE usage_daily (
|
|
||||||
user_id UUID NOT NULL,
|
|
||||||
date DATE NOT NULL,
|
|
||||||
llm_tokens BIGINT DEFAULT 0,
|
|
||||||
stt_seconds REAL DEFAULT 0,
|
|
||||||
tts_chars INTEGER DEFAULT 0,
|
|
||||||
estimated_cost NUMERIC(10,4) DEFAULT 0,
|
|
||||||
PRIMARY KEY (user_id, date)
|
|
||||||
);
|
|
||||||
```
|
|
||||||
|
|
||||||
## 部署架构
|
|
||||||
|
|
||||||
```
|
|
||||||
用户浏览器
|
|
||||||
↓
|
|
||||||
Nginx(同源反代 + 负载均衡)
|
|
||||||
├── / → 前端静态资源(CDN 或本地 dist)
|
|
||||||
├── /api/* → Go Gateway(REST API)
|
|
||||||
└── /ws → Go Gateway(WebSocket)
|
|
||||||
├── Gateway-1 ──→ Redis
|
|
||||||
├── Gateway-2 ──→ Redis
|
|
||||||
└── Gateway-N ──→ AI Services(外部 API)
|
|
||||||
```
|
|
||||||
|
|
||||||
**跨域策略**:Nginx 将前端和后端统一到同一域名下,浏览器无跨域问题。
|
|
||||||
|
|
||||||
### Nginx 配置
|
|
||||||
|
|
||||||
```nginx
|
|
||||||
server {
|
|
||||||
listen 80;
|
|
||||||
server_name camtalk.example.com;
|
|
||||||
|
|
||||||
# 前端静态资源
|
|
||||||
location / {
|
|
||||||
root /var/www/camtalk/dist;
|
|
||||||
try_files $uri $uri/ /index.html;
|
|
||||||
}
|
|
||||||
|
|
||||||
# REST API 反代
|
|
||||||
location /api/ {
|
|
||||||
proxy_pass http://127.0.0.1:8080;
|
|
||||||
proxy_set_header Host $host;
|
|
||||||
proxy_set_header X-Real-IP $remote_addr;
|
|
||||||
}
|
|
||||||
|
|
||||||
# WebSocket 反代
|
|
||||||
location /ws {
|
|
||||||
proxy_pass http://127.0.0.1:8080;
|
|
||||||
proxy_http_version 1.1;
|
|
||||||
proxy_set_header Upgrade $http_upgrade;
|
|
||||||
proxy_set_header Connection "upgrade";
|
|
||||||
proxy_set_header Host $host;
|
|
||||||
proxy_set_header X-Real-IP $remote_addr;
|
|
||||||
proxy_read_timeout 86400s; # 长连接超时 24h
|
|
||||||
proxy_send_timeout 86400s;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
> WebSocket 是长连接,Nginx 必须配置 `Upgrade` 和 `Connection` 头。`proxy_read_timeout` 需要覆盖心跳间隔(客户端 30s ping),否则 Nginx 会主动断开空闲连接。
|
|
||||||
|
|
||||||
### 开发环境
|
|
||||||
|
|
||||||
开发时前端(Vite :5173)和后端(Gin :8080)不同端口。前端 WebSocket 地址基于 `window.location.host` 动态构建,通过 Vite `server.proxy` 转发到后端,无需硬编码端口。
|
|
||||||
|
|
||||||
`vite.config.ts` 中配置了 `/ws`(WebSocket)和 `/api`(REST)的代理,目标为 `http://localhost:8080`。
|
|
||||||
@@ -4,15 +4,29 @@
|
|||||||
|
|
||||||
本文档记录项目中各项技术的**选型过程、替代方案对比和决策理由**。技术选型没有"绝对正确",只有"更适合"。
|
本文档记录项目中各项技术的**选型过程、替代方案对比和决策理由**。技术选型没有"绝对正确",只有"更适合"。
|
||||||
|
|
||||||
**定位**:持久化部分是拓展选型,不阻塞 MVP(MVP 用内存存储即可)。前端边缘处理部分是 MVP 阶段就需要确定的技术栈。AI 服务栈(STT/LLM/TTS)已确定默认选型,可通过配置灵活切换。
|
各技术选型章节包含关键术语解释,帮助快速理解技术概念。
|
||||||
|
|
||||||
|
### 后端核心技术栈
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **Go (Golang)** | 高并发后端语言,Google 开发,杀手锏是 goroutine——极轻量协程,一个程序可轻松开几万个,每个只占几 KB 内存,适合管理大量 WebSocket 长连接 |
|
||||||
|
| **gorilla/websocket** | Go WebSocket 库,Go 标准库无内置 WebSocket 支持,此库是社区最成熟的选择,处理了协议握手、帧解析等底层细节 |
|
||||||
|
| **Viper** | Go 配置管理库,读取 JSON/YAML/TOML 配置,支持环境变量覆盖,方便开发/测试/生产环境用不同配置 |
|
||||||
|
| **Zap** | Go 结构化日志库,Uber 开源,输出 JSON 格式日志,方便工具搜索分析,性能远超标准库 log |
|
||||||
|
|
||||||
```
|
```
|
||||||
技术选型
|
技术选型
|
||||||
|
├── AI 编排框架
|
||||||
|
│ └── CloudWeGo Eino Graph(声明式 DAG 编排,替代手写 goroutine 管道)
|
||||||
├── AI 服务栈
|
├── AI 服务栈
|
||||||
│ ├── STT: Deepgram(默认) / MiMo ASR
|
│ ├── STT: MiMo ASR(默认) / Deepgram
|
||||||
│ ├── LLM: GPT-4o(默认) / 通义千问等 OpenAI 兼容模型
|
│ ├── LLM: DashScope qwen3-vl-plus(默认) / GPT-4o 等 OpenAI 兼容模型
|
||||||
│ └── TTS: OpenAI TTS(默认) / MiMo TTS
|
│ └── TTS: MiMo TTS(默认) / OpenAI TTS
|
||||||
├── 持久化层 → 数据库选型: PostgreSQL(规划中,MVP 阶段使用内存存储)
|
├── 持久化层
|
||||||
|
│ ├── 数据库: PostgreSQL(pgx/v5,手写 SQL)
|
||||||
|
│ ├── 迁移: 嵌入式 SQL 文件,自动执行
|
||||||
|
│ └── 存储模式: 三级存储 TieredManager(L1 Memory → L2 Redis → L3 PostgreSQL)
|
||||||
├── 认证与用户系统
|
├── 认证与用户系统
|
||||||
│ ├── 认证方案: JWT (HS256), access 15min + refresh 7day
|
│ ├── 认证方案: JWT (HS256), access 15min + refresh 7day
|
||||||
│ ├── JWT 库: golang-jwt/jwt/v5
|
│ ├── JWT 库: golang-jwt/jwt/v5
|
||||||
@@ -20,48 +34,116 @@
|
|||||||
│ ├── 数据库驱动: pgx/v5(手写 SQL,不用 ORM)
|
│ ├── 数据库驱动: pgx/v5(手写 SQL,不用 ORM)
|
||||||
│ └── 前端 Token 存储: localStorage
|
│ └── 前端 Token 存储: localStorage
|
||||||
└── 前端边缘处理层
|
└── 前端边缘处理层
|
||||||
├── 边缘推理: ONNX Runtime Web(规划中,MVP 使用 Canvas 像素比较)
|
├── 关键帧检测: Canvas 像素比较(160x120 降采样)
|
||||||
├── 语音检测: @ricky0123/vad-web
|
├── 语音检测: @ricky0123/vad-web
|
||||||
└── 媒体采集: MediaDevices API
|
└── 媒体采集: MediaDevices API
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 一、AI 服务栈选型
|
## 一、AI 编排框架选型
|
||||||
|
|
||||||
|
### 关键术语
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **Eino** | 字节跳动开源的 Go AI 应用开发框架(CloudWeGo Eino),提供 Graph DAG 编排、组件抽象(ChatModel/Tool 等)、流式处理和 Callback AOP 机制 |
|
||||||
|
| **compose.Graph** | Eino 的 DAG 编排器,声明式有向无环图,节点可以是 Lambda、ChatModel、ToolsNode 等,边定义数据流向 |
|
||||||
|
| **Lambda** | Graph 中的可组合函数单元,四种模式:InvokableLambda(同步)、StreamableLambda(流式输出)、CollectableLambda(流式输入)、TransformableLambda(双向流式) |
|
||||||
|
| **StreamReader** | Eino 的流式数据抽象 `schema.StreamReader[T]`,类似 io.Reader 的语义,`Recv()` 读取一帧,`io.EOF` 表示流结束 |
|
||||||
|
| **Callback** | Eino 的 AOP 机制,类似中间件的钩子,支持节点生命周期回调(OnStart/OnEnd/OnError/OnEndWithStreamOutput) |
|
||||||
|
|
||||||
|
> 更多 Eino 相关概念详见 [10-Eino框架与编排设计.md](10-Eino框架与编排设计.md)
|
||||||
|
|
||||||
|
### 候选方案对比
|
||||||
|
|
||||||
|
| 框架 | 语言 | 特点 | CamTalk 适用性 |
|
||||||
|
|------|------|------|---------------|
|
||||||
|
| **CloudWeGo Eino** | Go | Go 原生、类型安全、流式原生、Graph DAG 编排 | ✅ 完美匹配 |
|
||||||
|
| LangChain Go | Go | 生态丰富但较重,抽象层多 | ❌ 过度抽象 |
|
||||||
|
| 自研编排 | Go | 完全可控 | ❌ 维护成本高 |
|
||||||
|
|
||||||
|
### 选择 Eino 的理由
|
||||||
|
|
||||||
|
| 维度 | 手写 goroutine(旧方案) | Eino Graph(新方案) |
|
||||||
|
|------|------------------------|---------------------|
|
||||||
|
| 编排方式 | 手动 `go func()` + `sync.WaitGroup` | 声明式 DAG,类型安全 |
|
||||||
|
| 流式处理 | 自定义 `chan` 传递 | `StreamReader` + `Pipe`,自动转换 |
|
||||||
|
| 错误处理 | 各节点独立处理,不一致 | Graph 级别统一错误传播 |
|
||||||
|
| 回调/AOP | 日志散落各处 | `callbacks.Handler` 统一注入 |
|
||||||
|
| 配置灵活性 | Pipeline 创建时固定 | 每请求 `Option` 动态注入 |
|
||||||
|
| 可测试性 | 需启动 goroutine | `Graph.Invoke()` 直接测试 |
|
||||||
|
| 扩展性 | 修改 Pipeline 代码 | 添加节点 + 边,无侵入 |
|
||||||
|
|
||||||
|
### 核心依赖
|
||||||
|
|
||||||
|
```
|
||||||
|
github.com/cloudwego/eino v0.9.9 # 核心框架
|
||||||
|
github.com/cloudwego/eino-ext/components/model/openai v0.1.13 # OpenAI 兼容 ChatModel
|
||||||
|
```
|
||||||
|
|
||||||
|
**核心理由**:
|
||||||
|
1. Go 原生,泛型支持,编译时类型检查
|
||||||
|
2. 原生流式处理(`StreamReader`),适合 LLM token 级推送
|
||||||
|
3. Graph 支持分支、并行、循环,满足当前和未来需求
|
||||||
|
4. Callback 机制实现 AOP(日志、指标、消息推送)
|
||||||
|
5. eino-ext 提供 OpenAI ChatModel 实现,直接对接 DashScope
|
||||||
|
|
||||||
|
> 详细的 Eino 框架使用文档见 [11-Eino框架技术文档](11-Eino框架技术文档.md),重构方案见 [10-Eino重构方案](10-Eino重构方案.md),实施记录见 [12-Eino重构实施记录](12-Eino重构实施记录.md)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 二、AI 服务栈选型
|
||||||
|
|
||||||
|
### 关键术语
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **多模态 LLM** | 能读文字又能看图片的大语言模型,如 GPT-4o(OpenAI)、Claude Sonnet(Anthropic),给照片+问题能"看懂"照片再回答 |
|
||||||
|
| **STT** | Speech-to-Text,语音转文字。流式识别延迟可低于 500ms |
|
||||||
|
| **TTS** | Text-to-Speech,文字转语音。支持流式——边生成边读,不必等全部生成完 |
|
||||||
|
|
||||||
### STT(语音识别)
|
### STT(语音识别)
|
||||||
|
|
||||||
| 方案 | 延迟 | 成本 | 特点 |
|
| 方案 | 延迟 | 成本 | 特点 |
|
||||||
|------|------|------|------|
|
|------|------|------|------|
|
||||||
| **Deepgram**(默认) | <500ms | 按分钟计费 | 流式识别,延迟极低,WebSocket 接口 |
|
| **MiMo ASR**(默认) | ~1s | 按量计费 | 国产替代,兼容 OpenAI chat/completions 格式,HTTP 非流式 |
|
||||||
| **MiMo ASR**(小米) | ~1s | 按量计费 | 国产替代,兼容 OpenAI chat/completions 格式,HTTP 非流式 |
|
| **Deepgram** | <500ms | 按分钟计费 | 流式识别,延迟极低,WebSocket 接口 |
|
||||||
| Whisper API | 1-3s | 按分钟计费 | 准确率高,支持多语言 |
|
| Whisper API | 1-3s | 按分钟计费 | 准确率高,支持多语言 |
|
||||||
| FunASR | <500ms | 自部署免费 | 阿里开源,中文优化 |
|
| FunASR | <500ms | 自部署免费 | 阿里开源,中文优化 |
|
||||||
|
|
||||||
当前默认使用 Deepgram nova-2,可通过 `ai.stt.provider` 配置切换到 MiMo ASR。
|
当前默认使用 MiMo ASR(mimo-v2.5-asr),可通过 `ai.stt.provider` 配置切换到 Deepgram。
|
||||||
|
|
||||||
### LLM(多模态大模型)
|
### LLM(多模态大模型)
|
||||||
|
|
||||||
| 方案 | 成本 | 特点 |
|
| 方案 | 成本 | 特点 |
|
||||||
|------|------|------|
|
|------|------|------|
|
||||||
| **GPT-4o**(默认) | $2.5/1M tokens | 视觉理解能力强,API 成熟,流式推理 |
|
| **DashScope qwen3-vl-plus**(默认) | 按量计费 | 阿里云,通过 OpenAI 兼容接口调用,视觉理解能力强 |
|
||||||
| 通义千问 qwen3-vl-plus | 按量计费 | 阿里云,通过 OpenAI 兼容接口调用 |
|
| GPT-4o | $2.5/1M tokens | OpenAI,API 成熟,流式推理 |
|
||||||
| Claude Sonnet | $3/1M tokens | Anthropic,长上下文能力强 |
|
| Claude Sonnet | $3/1M tokens | Anthropic,长上下文能力强 |
|
||||||
|
|
||||||
代码通过 OpenAI 兼容接口调用,可灵活切换到任何兼容服务商。配置 `ai.llm.provider`、`ai.llm.model`、`ai.llm.endpoint` 即可。
|
LLM 通过 Eino 框架的 `eino-ext/components/model/openai` ChatModel 组件接入,支持任何 OpenAI 兼容接口。配置 `ai.llm.provider`、`ai.llm.model`、`ai.llm.endpoint` 即可切换。
|
||||||
|
|
||||||
### TTS(语音合成)
|
### TTS(语音合成)
|
||||||
|
|
||||||
| 方案 | 成本 | 特点 |
|
| 方案 | 成本 | 特点 |
|
||||||
|------|------|------|
|
|------|------|------|
|
||||||
| **OpenAI TTS**(默认) | $15/1M 字符 | 音质自然,支持流式,默认模型 tts-1,语音 alloy |
|
| **MiMo TTS**(默认) | 按量计费 | 国产替代,通过配置切换,模型 mimo-v2.5-tts |
|
||||||
| MiMo TTS(小米) | 按量计费 | 国产替代,通过配置切换 |
|
| OpenAI TTS | $15/1M 字符 | 音质自然,支持流式,默认模型 tts-1,语音 alloy |
|
||||||
|
|
||||||
当前默认使用 OpenAI TTS(tts-1, alloy),可通过 `ai.tts.provider` 配置切换。
|
当前默认使用 MiMo TTS(mimo-v2.5-tts),可通过 `ai.tts.provider` 配置切换到 OpenAI TTS。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 二、持久化层选型(规划中,MVP 阶段使用内存存储)
|
## 三、持久化层选型
|
||||||
|
|
||||||
|
### 关键术语
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **PostgreSQL** | 关系型数据库,支持 JSONB(JSON 二进制格式,可建索引)、窗口函数、CTE 等高级特性 |
|
||||||
|
| **Redis** | 内存 KV 数据库,数据放在内存里,读写微秒级。支持 TTL 过期自动清理 |
|
||||||
|
| **MVCC** | Multi-Version Concurrency Control,多版本并发控制,PostgreSQL 用此实现高并发读写而不阻塞 |
|
||||||
|
|
||||||
### 数据特征分析
|
### 数据特征分析
|
||||||
|
|
||||||
@@ -146,17 +228,19 @@ ORDER BY created_at DESC
|
|||||||
LIMIT 20;
|
LIMIT 20;
|
||||||
```
|
```
|
||||||
|
|
||||||
### 冷热分离架构
|
### 冷热分离架构(三级存储)
|
||||||
|
|
||||||
```
|
```
|
||||||
Go Gateway
|
Go Gateway (TieredManager)
|
||||||
├── 写入路径 → Redis(实时会话状态)
|
├── L1: Memory(进程内缓存,微秒级)
|
||||||
│ → PostgreSQL(对话历史 + 用量)
|
├── L2: Redis(分布式缓存,毫秒级)
|
||||||
└── 读取路径 → Redis(当前上下文,快)
|
└── L3: PostgreSQL(持久化存储,冷数据)
|
||||||
→ PostgreSQL(历史记录,慢)
|
|
||||||
|
读取路径:L1 → L2 → L3,逐级回源,命中后向上回填
|
||||||
|
写入路径:L1 → L2(同步) → L3(异步)
|
||||||
```
|
```
|
||||||
|
|
||||||
建议异步写入——实时消息先写 Redis(快),异步批量刷入 PostgreSQL(慢),不影响对话体验。
|
`TieredManager` 自动管理三级存储,后台 goroutine 每 30 秒 ping Redis 健康状态,Redis 故障时自动降级为 L1+L3 模式。
|
||||||
|
|
||||||
### 决策流程
|
### 决策流程
|
||||||
|
|
||||||
@@ -173,7 +257,19 @@ Go Gateway
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 二、前端边缘处理层选型
|
## 四、前端边缘处理层选型
|
||||||
|
|
||||||
|
### 关键术语
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **React 18** | 组件化 UI 框架,Facebook 开源,把页面拆成组件搭积木拼装。18 版本支持并发渲染 |
|
||||||
|
| **TypeScript** | 带类型的 JavaScript,在 JS 基础上增加类型声明,编译阶段就能发现类型错误 |
|
||||||
|
| **Vite** | 前端构建工具,利用浏览器原生 ES Module,开发时毫秒级热更新(HMR),构建产物小 |
|
||||||
|
| **WebSocket** | 浏览器与服务器的双向通道。HTTP 是"一问一答",WebSocket 像打电话——接通后双方随时互发消息,适合实时对话场景 |
|
||||||
|
| **ONNX Runtime Web** | 浏览器端 AI 推理引擎,微软定义的通用模型格式 ONNX 的运行引擎,可在浏览器中用 WASM 加速跑轻量模型(如 VAD、关键帧检测),零延迟、不耗服务器资源 |
|
||||||
|
| **VAD** | Voice Activity Detection,语音活动检测,检测"人有没有在说话"。WebRTC 内置了高效的 VAD 算法 |
|
||||||
|
| **MediaDevices API** | 浏览器摄像头/麦克风接口,`navigator.mediaDevices.getUserMedia()` 是浏览器音视频采集的唯一标准入口,无需插件 |
|
||||||
|
|
||||||
### 总览
|
### 总览
|
||||||
|
|
||||||
@@ -223,7 +319,16 @@ vad-web 是"够用且最轻"的平衡点——直接包装浏览器原生 WebRTC
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 四、认证与用户系统选型
|
## 五、认证与用户系统选型
|
||||||
|
|
||||||
|
### 关键术语
|
||||||
|
|
||||||
|
| 名词 | 解释 |
|
||||||
|
|------|------|
|
||||||
|
| **JWT** | JSON Web Token,无状态 token,服务端不存 session,分布式友好 |
|
||||||
|
| **HS256** | HMAC-SHA256,JWT 对称签名算法,用同一密钥签名和验证 |
|
||||||
|
| **bcrypt** | 密码哈希算法,自适应 cost factor,抗暴力破解 |
|
||||||
|
| **pgx** | Go 生态性能最优的 PostgreSQL 驱动,原生协议实现,内置连接池 pgxpool |
|
||||||
|
|
||||||
### 总览
|
### 总览
|
||||||
|
|
||||||
1622
docs/03-接口文档.md
1622
docs/03-接口文档.md
File diff suppressed because it is too large
Load Diff
@@ -14,6 +14,8 @@
|
|||||||
麦克风 → VAD → STT → LLM → TTS → 扬声器
|
麦克风 → VAD → STT → LLM → TTS → 扬声器
|
||||||
```
|
```
|
||||||
|
|
||||||
|
> 后端 AI 编排基于 Eino Graph 声明式 DAG 实现:`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`。详见 [11-Eino框架技术文档](11-Eino框架技术文档.md)。
|
||||||
|
|
||||||
## 环节一:VAD(语音活动检测)
|
## 环节一:VAD(语音活动检测)
|
||||||
|
|
||||||
从持续音频流中检测"人什么时候在说话",避免将环境噪音当作有效输入。**浏览器端完成**,节省 ~70% 带宽。
|
从持续音频流中检测"人什么时候在说话",避免将环境噪音当作有效输入。**浏览器端完成**,节省 ~70% 带宽。
|
||||||
@@ -41,12 +43,12 @@ vad.start();
|
|||||||
|
|
||||||
| 方案 | 延迟 | 成本 | 特点 |
|
| 方案 | 延迟 | 成本 | 特点 |
|
||||||
|------|------|------|------|
|
|------|------|------|------|
|
||||||
|
| **MiMo ASR**(默认) | ~1s | 按量计费 | 国产替代,兼容 OpenAI 格式,HTTP 非流式 |
|
||||||
|
| **Deepgram** | <500ms | 按分钟计费 | 流式识别,延迟极低 |
|
||||||
| Whisper API | 1-3s | 按分钟计费 | 准确率高,支持多语言 |
|
| Whisper API | 1-3s | 按分钟计费 | 准确率高,支持多语言 |
|
||||||
| **Deepgram**(默认) | <500ms | 按分钟计费 | 流式识别,延迟极低 |
|
|
||||||
| **MiMo ASR**(小米) | ~1s | 按量计费 | 国产替代,兼容 OpenAI 格式,HTTP 非流式 |
|
|
||||||
| 浏览器原生 | ~1s | 免费 | 中文效果一般 |
|
| 浏览器原生 | ~1s | 免费 | 中文效果一般 |
|
||||||
|
|
||||||
当前实现为**一次性语音识别**(非流式):前端 VAD 检测到用户说完后,将完整音频片段发送到后端,后端调用 `stt.Recognize()` 一次性返回识别结果。流式 STT 为未来优化方向。
|
当前实现为**一次性语音识别**(非流式):前端 VAD 检测到用户说完后,将完整音频片段发送到后端,后端通过 Eino Graph 的 STT Lambda 节点调用 `stt.Recognize()` 一次性返回识别结果。流式 STT 为未来优化方向。
|
||||||
|
|
||||||
音频编码格式:前端 `audio.ts` 将 Float32Array 转为 Int16 PCM(16kHz, pcm_s16le)再编码为 Base64。
|
音频编码格式:前端 `audio.ts` 将 Float32Array 转为 Int16 PCM(16kHz, pcm_s16le)再编码为 Base64。
|
||||||
|
|
||||||
@@ -56,12 +58,12 @@ vad.start();
|
|||||||
|
|
||||||
句子切分规则:按中文标点(`。!?`)、英文标点(`. ! ?`)和换行符切分。
|
句子切分规则:按中文标点(`。!?`)、英文标点(`. ! ?`)和换行符切分。
|
||||||
|
|
||||||
当前实现参数:Voice `"alloy"`、Speed `1.0`、OutputFmt `"mp3"`、SampleRate `24000`。
|
当前实现参数:Voice `"mimo_default"`(可通过配置切换)、Speed `1.0`、OutputFmt `"mp3"`、SampleRate `24000`。
|
||||||
|
|
||||||
方案选择:
|
方案选择:
|
||||||
- **OpenAI TTS**(默认):音质好,延迟中等,按字符计费,模型 tts-1
|
- **MiMo TTS**(默认):国产替代,模型 mimo-v2.5-tts,通过配置切换
|
||||||
- **MiMo TTS**(小米):国产替代,通过配置切换
|
- **OpenAI TTS**:音质好,延迟中等,按字符计费,模型 tts-1
|
||||||
- **Edge TTS**(规划中):微软免费方案,音质不错,延迟略高
|
- **Edge TTS**(待实现):微软免费方案,音质不错,延迟略高
|
||||||
|
|
||||||
## 延迟优化要点
|
## 延迟优化要点
|
||||||
|
|
||||||
@@ -26,39 +26,32 @@
|
|||||||
| 用户触发 | 高 | 低 | 只在用户提问时拍照 |
|
| 用户触发 | 高 | 低 | 只在用户提问时拍照 |
|
||||||
| 本地预筛选 | 中 | 高 | 用轻量模型判断"是否值得问 LLM" |
|
| 本地预筛选 | 中 | 高 | 用轻量模型判断"是否值得问 LLM" |
|
||||||
|
|
||||||
```typescript
|
**实现细节**:参见 `frontend/src/lib/sampling.ts` 中的 SamplingController,根据 VAD 状态在空闲模式(5s/帧)和活跃模式(1s/帧)之间切换。
|
||||||
// 混合策略:定时低频 + 事件高频(sampling.ts)
|
|
||||||
const IDLE_INTERVAL = 5000; // 空闲 5 秒一帧
|
|
||||||
const ACTIVE_INTERVAL = 1000; // 用户说话时 1 秒一帧
|
|
||||||
|
|
||||||
// SamplingController 根据 VAD 状态切换采样间隔
|
|
||||||
// detail_level 通过 session config 静态配置,不随说话状态动态变化
|
|
||||||
```
|
|
||||||
|
|
||||||
## 策略二:端云协同——把计算推到边缘
|
## 策略二:端云协同——把计算推到边缘
|
||||||
|
|
||||||
不是所有计算都需要上云。可前置到客户端的计算:
|
不是所有计算都需要上云。可前置到客户端的计算:
|
||||||
|
|
||||||
- **VAD 语音检测**:浏览器端完成,减少无效音频上传(节省 ~70% 带宽)
|
- **VAD 语音检测**:浏览器端完成,减少无效音频上传(节省 ~70% 带宽)
|
||||||
- **人脸/物体检测**(规划中):用 ONNX Runtime 跑轻量模型(如 YOLOv8-nano ~6MB,推理 ~30ms),只在检测到新物体时触发 LLM。当前 MVP 使用 Canvas 像素比较做关键帧检测
|
- **人脸/物体检测**(待实现):用 ONNX Runtime 跑轻量模型(如 YOLOv8-nano ~6MB,推理 ~30ms),只在检测到新物体时触发 LLM。当前 MVP 使用 Canvas 像素比较做关键帧检测
|
||||||
- **重复画面过滤**:计算帧间相似度,对话模式 similarity > 0.9 跳过,观察模式 similarity < 0.85 触发
|
- **重复画面过滤**:计算帧间相似度,对话模式 similarity > 0.9 跳过,观察模式 similarity < 0.85 触发
|
||||||
- **敏感内容过滤**(规划中):NSFW 检测前置,避免无效 API 调用
|
- **敏感内容过滤**(待实现):NSFW 检测前置,避免无效 API 调用
|
||||||
|
|
||||||
## 策略三:模型分级——用对模型做对事(规划中)
|
## 策略三:模型分级——用对模型做对事(待实现)
|
||||||
|
|
||||||
不是每个问题都需要最贵的模型:
|
不是每个问题都需要最贵的模型:
|
||||||
|
|
||||||
```
|
```
|
||||||
用户提问 → 问题复杂度判断
|
用户提问 → 问题复杂度判断
|
||||||
├── 简单识别 → GPT-4o-mini ($0.15/1M tokens)
|
├── 简单识别 → 轻量模型(如 qwen-turbo)
|
||||||
├── 深度分析 → GPT-4o ($2.5/1M tokens)
|
├── 深度分析 → qwen3-vl-plus(默认,按量计费)
|
||||||
└── 代码/推理 → o1 ($15/1M tokens)
|
└── 代码/推理 → 更强模型(如 o1)
|
||||||
```
|
```
|
||||||
|
|
||||||
> 当前 MVP 阶段使用单一模型(默认 GPT-4o),模型分级路由为未来优化方向。通过配置 `ai.llm.model` 可手动切换模型。
|
> 当前 MVP 阶段使用单一模型(默认 DashScope qwen3-vl-plus),模型分级路由为未来优化方向。LLM 通过 Eino ChatModel 接入,支持任何 OpenAI 兼容接口。
|
||||||
|
|
||||||
## 策略四:缓存与复用(规划中)
|
## 策略四:缓存与复用(待实现)
|
||||||
|
|
||||||
- **语义缓存**(规划中):相似问题直接返回缓存结果(如反复问"这是什么")
|
- **语义缓存**(待实现):相似问题直接返回缓存结果(如反复问"这是什么")
|
||||||
- **上下文复用**:连续对话中,未变化的图像不必重复发送(已通过重复画面过滤实现)
|
- **上下文复用**:连续对话中,未变化的图像不必重复发送(已通过重复画面过滤实现)
|
||||||
- **对话历史裁剪**:前端按 `MAX_HISTORY_ROUNDS = 10` 裁剪,后端按 `defaultHistorySize = 20` 裁剪,限制每轮的固定 token 开销
|
- **对话历史裁剪**:前端按 `MAX_HISTORY_ROUNDS = 10` 裁剪,后端按 `defaultHistorySize = 20` 裁剪,限制每轮的固定 token 开销
|
||||||
523
docs/08-Eino框架与编排设计.md
Normal file
523
docs/08-Eino框架与编排设计.md
Normal file
@@ -0,0 +1,523 @@
|
|||||||
|
# CamTalk Eino 框架与编排设计
|
||||||
|
|
||||||
|
## 1. 概述
|
||||||
|
|
||||||
|
### 1.1 为什么选择 Eino
|
||||||
|
|
||||||
|
[CloudWeGo Eino](https://github.com/cloudwego/eino) 是字节跳动 CloudWeGo 团队开源的 AI 应用开发框架,提供基于图(Graph)的编排能力、组件抽象和流式处理支持。
|
||||||
|
|
||||||
|
CamTalk 使用 Eino 替代原有的手写 goroutine 管道,实现 STT → LLM → TTS 的声明式编排。
|
||||||
|
|
||||||
|
**技术选型对比:**
|
||||||
|
|
||||||
|
| 维度 | 手写 goroutine(旧方案) | Eino Graph(新方案) |
|
||||||
|
|------|------------------------|---------------------|
|
||||||
|
| 编排方式 | 手动 `go func()` + `sync.WaitGroup` | 声明式 DAG,类型安全 |
|
||||||
|
| 流式处理 | 自定义 `chan` 传递 | `StreamReader` + `Pipe`,自动转换 |
|
||||||
|
| 错误处理 | 各节点独立处理,不一致 | Graph 级别统一错误传播 |
|
||||||
|
| 回调/AOP | 日志散落各处 | `callbacks.Handler` 统一注入 |
|
||||||
|
| 配置灵活性 | Pipeline 创建时固定 | 每请求 `Option` 动态注入 |
|
||||||
|
| 可测试性 | 需启动 goroutine | `Graph.Invoke()` 直接测试 |
|
||||||
|
| 扩展性 | 修改 Pipeline 代码 | 添加节点 + 边,无侵入 |
|
||||||
|
| 并发安全 | 手动 `sync` | State 自动加锁 |
|
||||||
|
|
||||||
|
**选择 Eino 的核心理由:**
|
||||||
|
1. Go 原生,泛型支持,编译时类型检查
|
||||||
|
2. 原生流式处理(`StreamReader`),适合 LLM token 级推送
|
||||||
|
3. Graph 支持分支、并行、循环,满足当前和未来需求
|
||||||
|
4. Callback 机制实现 AOP(日志、指标、消息推送)
|
||||||
|
5. eino-ext 提供 OpenAI ChatModel 实现,直接对接 DashScope
|
||||||
|
|
||||||
|
### 1.2 旧方案的问题
|
||||||
|
|
||||||
|
当前后端 AI 编排层(`internal/orchestrator/pipeline.go`)为手写 goroutine 管道存在以下问题:
|
||||||
|
|
||||||
|
1. **编排逻辑硬编码**:STT→LLM→TTS 流程写死,扩展困难
|
||||||
|
2. **并发控制粗糙**:手动 goroutine 调度,缺乏结构化流式传递
|
||||||
|
3. **无回调/AOP 机制**:日志、指标、追踪散落各处
|
||||||
|
4. **配置耦合**:模型名、TTS 参数等硬编码在结构体
|
||||||
|
5. **错误处理不一致**:TTS 错误静默吞掉,STT/LLM 错误通过 Sender 发送
|
||||||
|
|
||||||
|
### 1.3 核心依赖版本
|
||||||
|
|
||||||
|
```go
|
||||||
|
github.com/cloudwego/eino v0.9.9
|
||||||
|
github.com/cloudwego/eino-ext/components/model/openai v0.1.13
|
||||||
|
```
|
||||||
|
|
||||||
|
## 2. Eino 核心概念
|
||||||
|
|
||||||
|
### 2.1 Lambda
|
||||||
|
|
||||||
|
Lambda 是 Graph 中的可组合函数单元,支持四种模式:
|
||||||
|
|
||||||
|
| 模式 | 函数签名 | 构造方法 | 说明 |
|
||||||
|
|------|---------|---------|------|
|
||||||
|
| Invoke | `I → O` | `compose.InvokableLambda()` | 同步调用 |
|
||||||
|
| Stream | `I → StreamReader[O]` | `compose.StreamableLambda()` | 流式输出 |
|
||||||
|
| Collect | `StreamReader[I] → O` | `compose.CollectableLambda()` | 流式输入 |
|
||||||
|
| Transform | `StreamReader[I] → StreamReader[O]` | `compose.TransformableLambda()` | 双向流式 |
|
||||||
|
|
||||||
|
**返回类型**:所有 Lambda 构造函数返回 `*compose.Lambda`。
|
||||||
|
|
||||||
|
### 2.2 Graph
|
||||||
|
|
||||||
|
Graph 是有向无环图(DAG)编排器,支持:
|
||||||
|
- **节点**:Lambda、ChatModel、ToolsNode 等
|
||||||
|
- **边**:`g.AddEdge(from, to)` 定义数据流向
|
||||||
|
- **分支**:`g.AddBranch()` 条件路由
|
||||||
|
- **State**:`compose.WithGenLocalState()` 跨节点共享状态
|
||||||
|
|
||||||
|
```go
|
||||||
|
g := compose.NewGraph[PipelineInput, PipelineOutput]()
|
||||||
|
g.AddLambdaNode("stt", sttLambda)
|
||||||
|
g.AddChatModelNode("llm", chatModel)
|
||||||
|
g.AddEdge(compose.START, "stt")
|
||||||
|
g.AddEdge("stt", "llm")
|
||||||
|
g.AddEdge("llm", compose.END)
|
||||||
|
|
||||||
|
runnable, err := g.Compile(ctx)
|
||||||
|
output, err := runnable.Invoke(ctx, input) // 同步调用
|
||||||
|
stream, err := runnable.Stream(ctx, input) // 流式调用
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.3 ChatModel
|
||||||
|
|
||||||
|
ChatModel 是 LLM 组件抽象,接口定义:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type BaseChatModel interface {
|
||||||
|
Generate(ctx, []*schema.Message, ...Option) (*schema.Message, error)
|
||||||
|
Stream(ctx, []*schema.Message, ...Option) (*schema.StreamReader[*schema.Message], error)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
CamTalk 使用 `eino-ext/components/model/openai` 实现,通过 `BaseURL` 对接 DashScope:
|
||||||
|
|
||||||
|
```go
|
||||||
|
chatModel, _ := openai.NewChatModel(ctx, &openai.ChatModelConfig{
|
||||||
|
APIKey: cfg.AI.LLM.APIKey,
|
||||||
|
Model: cfg.AI.LLM.Model,
|
||||||
|
BaseURL: cfg.AI.LLM.Endpoint, // "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.4 StreamReader
|
||||||
|
|
||||||
|
`schema.StreamReader[T]` 是 Eino 的流式数据抽象:
|
||||||
|
- `sr.Recv()` 读取一帧,`io.EOF` 表示流结束
|
||||||
|
- `schema.Pipe[T](bufSize)` 创建 `StreamReader` + `StreamWriter` 对
|
||||||
|
- 框架自动处理 `T ↔ StreamReader[T]` 的转换(装箱/concat)
|
||||||
|
|
||||||
|
### 2.5 Callback
|
||||||
|
|
||||||
|
Callback 是 Eino 的 AOP 机制,支持节点生命周期钩子:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type Handler interface {
|
||||||
|
OnStart(ctx, *RunInfo, CallbackInput) context.Context
|
||||||
|
OnEnd(ctx, *RunInfo, CallbackOutput) context.Context
|
||||||
|
OnError(ctx, *RunInfo, error) context.Context
|
||||||
|
OnStartWithStreamInput(ctx, *RunInfo, *StreamReader[CallbackInput]) context.Context
|
||||||
|
OnEndWithStreamOutput(ctx, *RunInfo, *StreamReader[CallbackOutput]) context.Context
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
CamTalk 使用 `utils/callbacks.NewHandlerHelper()` 构建 typed handler:
|
||||||
|
- `ModelCallbackHandler.OnEndWithStreamOutput`:逐 token 推送 `llm_chunk`
|
||||||
|
|
||||||
|
### 2.6 State
|
||||||
|
|
||||||
|
Graph 全局状态,通过 `WithGenLocalState` 注册:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type PipelineState struct {
|
||||||
|
FullResponse strings.Builder
|
||||||
|
TranscribedText string
|
||||||
|
TokenUsage *TokenUsage
|
||||||
|
}
|
||||||
|
|
||||||
|
g := compose.NewGraph[I, O](compose.WithGenLocalState(func(ctx context.Context) *PipelineState {
|
||||||
|
return &PipelineState{}
|
||||||
|
}))
|
||||||
|
```
|
||||||
|
|
||||||
|
节点通过 `compose.ProcessState` 读写 State。
|
||||||
|
|
||||||
|
## 3. CamTalk Graph 设计
|
||||||
|
|
||||||
|
### 3.1 拓扑结构
|
||||||
|
|
||||||
|
```
|
||||||
|
START → STT → History → ChatModel → Splitter → TTS → Done → END
|
||||||
|
```
|
||||||
|
|
||||||
|
| 节点 | 类型 | 输入 → 输出 | 职责 |
|
||||||
|
|------|------|------------|------|
|
||||||
|
| STT | InvokableLambda | `PipelineInput → STTOutput` | 语音识别,写入 State |
|
||||||
|
| History | InvokableLambda | `STTOutput → []*schema.Message` | 组装提示词和历史 |
|
||||||
|
| ChatModel | ChatModel(原生) | `[]*schema.Message → StreamReader[*Message]` | LLM 流式推理 |
|
||||||
|
| Splitter | TransformableLambda | `StreamReader[string] → StreamReader[[]string]` | 句子切分 |
|
||||||
|
| TTS | InvokableLambda | `[]string → struct{}` | 语音合成,推送音频 |
|
||||||
|
| Done | InvokableLambda | `struct{} → PipelineOutput` | 发送 llm_done |
|
||||||
|
|
||||||
|
### 3.2 数据类型定义
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Graph 统一输入
|
||||||
|
type PipelineInput struct {
|
||||||
|
AudioData []byte // base64 解码后的音频(可选)
|
||||||
|
ImageData []byte // base64 解码后的图像(可选)
|
||||||
|
Text string // 直接文本输入(可选,跳过 STT)
|
||||||
|
SessionID string
|
||||||
|
RequestID string
|
||||||
|
Language string // zh / en
|
||||||
|
Scenario string // free_chat, interviewer, etc.
|
||||||
|
}
|
||||||
|
|
||||||
|
// Graph 统一输出
|
||||||
|
type PipelineOutput struct {
|
||||||
|
TranscribedText string // STT 结果
|
||||||
|
FullResponse string // LLM 完整回复
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pipeline State(跨节点共享)
|
||||||
|
type PipelineState struct {
|
||||||
|
FullResponse strings.Builder
|
||||||
|
TranscribedText string
|
||||||
|
TokenUsage *TokenUsage
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.3 流式模式
|
||||||
|
|
||||||
|
Graph 使用 **Stream 模式**调用:
|
||||||
|
- 内部所有节点以 Transform 模式运行
|
||||||
|
- ChatModel 的 `Stream()` 方法实现真正的 token 级流式
|
||||||
|
- 适配器消费 `StreamReader[PipelineOutput]` 触发整条链路
|
||||||
|
|
||||||
|
### 3.4 消息推送机制
|
||||||
|
|
||||||
|
| 消息 | 推送方式 | 时机 |
|
||||||
|
|------|---------|------|
|
||||||
|
| `stt_result` | Lambda 内部直接调用 Sender | STT 完成后 |
|
||||||
|
| `llm_chunk` | Callback `OnEndWithStreamOutput` | ChatModel 逐 token |
|
||||||
|
| `tts_audio` | Lambda 内部直接调用 Sender | TTS 逐句合成 |
|
||||||
|
| `llm_done` | Lambda 内部直接调用 Sender | Done 节点执行时 |
|
||||||
|
|
||||||
|
**Context 注入**:Sender、RequestID、SessionID、PipelineState 通过 `context.WithValue` 传递。
|
||||||
|
|
||||||
|
### 3.5 多模态支持
|
||||||
|
|
||||||
|
History 节点将图片构建为 `schema.Message.UserInputMultiContent`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
systemMsg.UserInputMultiContent = []schema.MessageInputPart{
|
||||||
|
{
|
||||||
|
Type: schema.ChatMessagePartTypeImageURL,
|
||||||
|
Image: &schema.MessageInputImage{
|
||||||
|
MessagePartCommon: schema.MessagePartCommon{
|
||||||
|
Base64Data: &base64Str,
|
||||||
|
MIMEType: "image/jpeg",
|
||||||
|
},
|
||||||
|
Detail: schema.ImageURLDetailAuto,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 4. 实现要点
|
||||||
|
|
||||||
|
### 4.1 目录结构
|
||||||
|
|
||||||
|
```
|
||||||
|
backend/internal/eino/
|
||||||
|
├── types.go # PipelineInput/Output、STTOutput、TokenUsage
|
||||||
|
├── state.go # PipelineState(跨节点状态)
|
||||||
|
├── callback.go # Callback handler(LLM token 推送)
|
||||||
|
├── graph.go # Graph 构建与编译
|
||||||
|
├── adapter.go # EinoOrchestrator(Orchestrator 接口适配器)
|
||||||
|
├── nodes_stt.go # STT Lambda
|
||||||
|
├── nodes_history.go # 历史组装 Lambda
|
||||||
|
├── nodes_splitter.go # 句子分割 Transform Lambda
|
||||||
|
├── nodes_tts.go # TTS Lambda
|
||||||
|
├── nodes_done.go # Done Lambda
|
||||||
|
└── graph_test.go # 单元测试
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.2 关键节点实现
|
||||||
|
|
||||||
|
#### STT Lambda(可选跳过)
|
||||||
|
|
||||||
|
```go
|
||||||
|
func sttLambda(sttSvc stt.Service) func(ctx context.Context, input PipelineInput) (STTOutput, error) {
|
||||||
|
return func(ctx context.Context, input PipelineInput) (STTOutput, error) {
|
||||||
|
// 文本模式:跳过 STT
|
||||||
|
if input.Text != "" {
|
||||||
|
return STTOutput{Text: input.Text, Language: input.Language}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 调用 STT 服务
|
||||||
|
result, err := sttSvc.Recognize(ctx, input.AudioData, stt.Options{
|
||||||
|
Language: input.Language,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return STTOutput{}, fmt.Errorf("STT error: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return STTOutput{Text: result.Text, Language: result.Language}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Splitter Transform Lambda(句子切分)
|
||||||
|
|
||||||
|
```go
|
||||||
|
func splitterLambda() func(ctx, *schema.StreamReader[*schema.Message]) (*schema.StreamReader[[]string], error) {
|
||||||
|
return func(ctx context.Context, stream *schema.StreamReader[*schema.Message]) (*schema.StreamReader[[]string], error) {
|
||||||
|
sr, sw := schema.Pipe[[]string](8)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer sw.Close()
|
||||||
|
var buffer []rune
|
||||||
|
|
||||||
|
for {
|
||||||
|
chunk, err := stream.Recv()
|
||||||
|
if err != nil {
|
||||||
|
if err == io.EOF {
|
||||||
|
if len(buffer) > 0 {
|
||||||
|
sw.Send([]string{string(buffer)}, nil)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sw.Send(nil, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range chunk.Content {
|
||||||
|
buffer = append(buffer, r)
|
||||||
|
if isSentenceDelimiter(r) {
|
||||||
|
sw.Send([]string{string(buffer)}, nil)
|
||||||
|
buffer = buffer[:0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return sr, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### TTS Lambda(并行合成)
|
||||||
|
|
||||||
|
```go
|
||||||
|
func ttsLambda(ttsSvc tts.Service, sender orchestrator.Sender) func(ctx, []string) (struct{}, error) {
|
||||||
|
return func(ctx context.Context, sentences []string) (struct{}, error) {
|
||||||
|
for _, sentence := range sentences {
|
||||||
|
if sentence == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// 调用 TTS 服务
|
||||||
|
audioData, err := ttsSvc.Synthesize(ctx, sentence, tts.Options{})
|
||||||
|
if err != nil {
|
||||||
|
// TTS 失败不中断流程,仅记录日志
|
||||||
|
log.Warn("TTS synthesis failed", zap.Error(err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// 推送音频到客户端
|
||||||
|
sender.SendTTSAudio(orchestrator.TTSAudioPayload{
|
||||||
|
Audio: audioData,
|
||||||
|
Format: "mp3",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return struct{}{}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.3 Callback 集成
|
||||||
|
|
||||||
|
```go
|
||||||
|
// ModelCallbackHandler 用于 LLM token 推送
|
||||||
|
type ModelCallbackHandler struct {
|
||||||
|
sender orchestrator.Sender
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *ModelCallbackHandler) OnEndWithStreamOutput(
|
||||||
|
ctx context.Context,
|
||||||
|
info *callbacks.RunInfo,
|
||||||
|
output *schema.StreamReader[*schema.Message],
|
||||||
|
) context.Context {
|
||||||
|
// 逐 token 推送到客户端
|
||||||
|
for {
|
||||||
|
msg, err := output.Recv()
|
||||||
|
if err == io.EOF {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
h.sender.SendLLMChunk(orchestrator.LLMChunkPayload{
|
||||||
|
Content: msg.Content,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.4 按请求动态配置
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 运行时 Option:每请求可变
|
||||||
|
func WithModelName(name string) compose.Option {
|
||||||
|
return compose.WithChatModelOption(model.WithModel(name))
|
||||||
|
}
|
||||||
|
|
||||||
|
func WithTemperature(temp float32) compose.Option {
|
||||||
|
return compose.WithChatModelOption(model.WithTemperature(temp))
|
||||||
|
}
|
||||||
|
|
||||||
|
// WebSocket Handler 中的调用
|
||||||
|
func (c *Client) handleQuery(req QueryRequest) {
|
||||||
|
opts := []compose.Option{}
|
||||||
|
|
||||||
|
if req.Model != "" {
|
||||||
|
opts = append(opts, WithModelName(req.Model))
|
||||||
|
}
|
||||||
|
|
||||||
|
output, err := c.pipeline.Stream(ctx, PipelineInput{...}, opts...)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.5 注意事项
|
||||||
|
|
||||||
|
#### 值类型 vs 指针类型
|
||||||
|
Graph 泛型参数必须使用值类型(`PipelineInput`/`PipelineOutput`),所有 Lambda 的输入输出也使用值类型。框架在 Transform 模式下会自动处理 `T` 和 `StreamReader[T]` 的转换。
|
||||||
|
|
||||||
|
#### Callback 运行时传入
|
||||||
|
Callback 通过 `Stream()` 的 option 传入,不在 `Compile()` 时注册:
|
||||||
|
|
||||||
|
```go
|
||||||
|
streamReader, err := runnable.Stream(ctx, input, compose.WithCallbacks(handler))
|
||||||
|
```
|
||||||
|
|
||||||
|
#### eino-ext 与 DashScope 兼容性
|
||||||
|
eino-ext OpenAI ChatModel 通过 `BaseURL` 对接 DashScope 兼容接口。需注意:
|
||||||
|
- 多模态图片使用 `Base64Data` + `MIMEType` 格式
|
||||||
|
- `Timeout` 控制单次请求超时
|
||||||
|
- 流式输出通过 `Stream()` 方法获取 `StreamReader[*schema.Message]`
|
||||||
|
|
||||||
|
#### 框架自动类型转换
|
||||||
|
Eino 框架在编排场景中自动处理以下转换:
|
||||||
|
- **T → StreamReader[T]**:将完整值装箱为单帧流(非流式 → 假流式)
|
||||||
|
- **StreamReader[T] → T**:将流 concat 为完整值(流式 → 非流式)
|
||||||
|
|
||||||
|
这使得不同流式模式的节点可以无缝连接。
|
||||||
|
|
||||||
|
## 5. 测试策略
|
||||||
|
|
||||||
|
### 5.1 单元测试
|
||||||
|
|
||||||
|
```go
|
||||||
|
func TestPipelineGraph_WithTextInput(t *testing.T) {
|
||||||
|
mockLLM := &mockChatModel{responses: []string{"你好!"}}
|
||||||
|
mockSender := &mockSender{}
|
||||||
|
|
||||||
|
graph, err := NewPipelineGraph(ctx, &GraphOption{
|
||||||
|
ChatModel: mockLLM,
|
||||||
|
Sender: mockSender,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
output, err := graph.Invoke(ctx, PipelineInput{
|
||||||
|
Text: "你好",
|
||||||
|
SessionID: "test-session",
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, "你好!", output.FullResponse)
|
||||||
|
assert.True(t, mockSender.LLMDoneSent)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPipelineGraph_WithAudioInput(t *testing.T) {
|
||||||
|
mockSTT := &mockSTT{text: "你好"}
|
||||||
|
mockLLM := &mockChatModel{responses: []string{"你好!"}}
|
||||||
|
mockTTS := &mockTTS{audio: []byte("fake-audio")}
|
||||||
|
mockSender := &mockSender{}
|
||||||
|
|
||||||
|
graph, _ := NewPipelineGraph(ctx, &GraphOption{
|
||||||
|
ChatModel: mockLLM,
|
||||||
|
STTService: mockSTT,
|
||||||
|
TTSService: mockTTS,
|
||||||
|
Sender: mockSender,
|
||||||
|
})
|
||||||
|
|
||||||
|
output, err := graph.Invoke(ctx, PipelineInput{
|
||||||
|
AudioData: []byte("fake-audio-data"),
|
||||||
|
SessionID: "test-session",
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, mockSender.TTSAudioSent)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5.2 集成测试
|
||||||
|
|
||||||
|
- 启动真实 OpenAI API 调用(使用测试 key)
|
||||||
|
- 验证 WebSocket 消息序列:`stt_result` → `llm_chunk` × N → `llm_done` → `tts_audio` × N
|
||||||
|
- 验证 interrupt 取消功能
|
||||||
|
- 验证多并发请求隔离
|
||||||
|
|
||||||
|
## 6. 未来扩展路径
|
||||||
|
|
||||||
|
基于 Eino Graph 的重构完成后,可无缝扩展:
|
||||||
|
|
||||||
|
1. **ReAct Agent**:Graph 添加 Branch 节点,实现 LLM → Tool → LLM 循环
|
||||||
|
2. **多模态理解**:添加视觉分析 Lambda 节点(图像描述 → 上下文注入)
|
||||||
|
3. **Model Router**:Graph 前置分支节点,按场景/成本路由不同 LLM
|
||||||
|
4. **Rate Limiter**:通过 Callback 的 OnStart 实现令牌桶
|
||||||
|
5. **Checkpoint/Resume**:利用 Eino 的 CheckpointStore 实现断点续传
|
||||||
|
6. **Multi-Agent**:利用 ADK 的 Supervisor/SequentialAgent 编排复杂对话流程
|
||||||
|
|
||||||
|
## 附录:关键 Eino API 参考
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 构建 Graph
|
||||||
|
g := compose.NewGraph[I, O](opts...)
|
||||||
|
g.AddChatModelNode(key, chatModel)
|
||||||
|
g.AddLambdaNode(key, lambda, opts...)
|
||||||
|
g.AddEdge(from, to)
|
||||||
|
g.AddBranch(from, branchFunc, mapping)
|
||||||
|
|
||||||
|
// 编译
|
||||||
|
runnable, err := g.Compile(ctx, opts...)
|
||||||
|
|
||||||
|
// 执行四种模式
|
||||||
|
output, err := runnable.Invoke(ctx, input, opts...)
|
||||||
|
stream, err := runnable.Stream(ctx, input, opts...)
|
||||||
|
output, err := runnable.Collect(ctx, inputStream, opts...)
|
||||||
|
stream, err := runnable.Transform(ctx, inputStream, opts...)
|
||||||
|
|
||||||
|
// Lambda 四种构造器
|
||||||
|
lambda := compose.InvokableLambda(fn) // I → O
|
||||||
|
lambda := compose.StreamableLambda(fn) // I → StreamReader[O]
|
||||||
|
lambda := compose.CollectableLambda(fn) // StreamReader[I] → O
|
||||||
|
lambda := compose.TransformableLambda(fn) // StreamReader[I] → StreamReader[O]
|
||||||
|
|
||||||
|
// Stream 操作
|
||||||
|
sr, sw := schema.Pipe[T](bufSize)
|
||||||
|
sw.Send(chunk, err)
|
||||||
|
chunk, err := sr.Recv()
|
||||||
|
sw.Close()
|
||||||
|
|
||||||
|
// Option
|
||||||
|
compose.WithCallbacks(handler)
|
||||||
|
compose.WithCallbacks(handler).DesignateNode("node_key")
|
||||||
|
compose.WithChatModelOption(model.WithTemperature(0.7))
|
||||||
|
compose.WithGenLocalState(genFunc)
|
||||||
|
```
|
||||||
348
docs/09-情景切换.md
Normal file
348
docs/09-情景切换.md
Normal file
@@ -0,0 +1,348 @@
|
|||||||
|
# 情景切换功能
|
||||||
|
|
||||||
|
## 功能概述
|
||||||
|
|
||||||
|
情景切换功能允许用户选择不同的对话场景,AI 会根据选择的情景扮演不同的角色:
|
||||||
|
|
||||||
|
| 情景 | AI 角色 | 主要功能 |
|
||||||
|
|------|---------|---------|
|
||||||
|
| 🎯 模拟面试官 | 资深面试官 | 提出面试问题,评估候选人能力,给出反馈 |
|
||||||
|
| 📚 英语老师 | 英语外教 | 全英文对话,纠正语法错误,引导深入交流 |
|
||||||
|
| ⚔️ 辩论对手 | 辩论选手 | 站在反方立场,用逻辑和证据反驳观点 |
|
||||||
|
| 🌐 同声翻译 | 翻译员 | 实时中英互译,口语化翻译,无额外解释 |
|
||||||
|
| 💬 自由对话 | 视觉助手 | 通用视觉对话助手(默认) |
|
||||||
|
|
||||||
|
### 核心特性
|
||||||
|
|
||||||
|
1. **情景首句引导**:切换情景后,AI 自动发送第一句话引导用户进入角色
|
||||||
|
2. **情景提示卡片**:对话顶部显示当前情景模式的蓝色提示卡片
|
||||||
|
3. **增强 System Prompt**:每个情景有详细的角色定位、交互规则和约束
|
||||||
|
4. **多语言支持**:完整支持中文、英文、日文界面
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 技术实现
|
||||||
|
|
||||||
|
### 后端实现
|
||||||
|
|
||||||
|
#### 1. 情景 Prompt 定义
|
||||||
|
|
||||||
|
**文件**: `backend/internal/ai/llm/scenarios.go`
|
||||||
|
|
||||||
|
- 扩展 `scenarioPrompt` 结构体,新增首句引导字段(GreetingZH/EN/JA)
|
||||||
|
- 增强所有情景的 System Prompt(添加角色定位、交互规则、约束)
|
||||||
|
- 新增函数 `GetScenarioGreeting(scenarioID, language string) string`
|
||||||
|
|
||||||
|
**示例 Prompt**(模拟面试官):
|
||||||
|
|
||||||
|
```go
|
||||||
|
"interviewer": {
|
||||||
|
ZH: `你是一位资深面试官。你通过摄像头观察面试者...
|
||||||
|
|
||||||
|
【角色定位】
|
||||||
|
- 你是面试官,不是助手或顾问
|
||||||
|
- 你的目标是评估候选人的能力
|
||||||
|
- 保持专业、客观、礼貌
|
||||||
|
|
||||||
|
【交互规则】
|
||||||
|
1. 每次只问一个问题,等用户回答后再追问
|
||||||
|
2. 问题要有层次:自我介绍 → 专业问题 → 情景题
|
||||||
|
3. 对用户的回答给出简短点评,然后追问
|
||||||
|
...`,
|
||||||
|
GreetingZH: "你好!我是今天的面试官。让我们先从自我介绍开始...",
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 2. 首句引导推送
|
||||||
|
|
||||||
|
**文件**: `backend/internal/ws/handler.go`
|
||||||
|
|
||||||
|
在处理 `config` 消息时,如果切换到非自由对话情景,自动返回首句引导:
|
||||||
|
|
||||||
|
```go
|
||||||
|
case "config":
|
||||||
|
// ... 更新配置 ...
|
||||||
|
|
||||||
|
// 如果切换了情景(非自由对话),返回首句引导
|
||||||
|
if scenarioID != "" && scenarioID != "free_chat" {
|
||||||
|
greeting := llm.GetScenarioGreeting(scenarioID, sess.Config.Language)
|
||||||
|
if greeting != "" {
|
||||||
|
// 发送 llm_chunk 和 llm_done 消息
|
||||||
|
// 追加到历史记录
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 3. State 初始化
|
||||||
|
|
||||||
|
**文件**: `backend/internal/eino/adapter.go`
|
||||||
|
|
||||||
|
从 `PipelineInput` 复制元数据到 `PipelineState`,确保情景配置正确传递到所有节点:
|
||||||
|
|
||||||
|
```go
|
||||||
|
state := genLocalState(ctx)
|
||||||
|
state.SessionID = input.SessionID
|
||||||
|
state.RequestID = input.RequestID
|
||||||
|
state.ImageData = input.ImageData
|
||||||
|
state.Scenario = input.Scenario // 关键:复制情景配置
|
||||||
|
state.Language = input.Language
|
||||||
|
state.DetailLevel = sess.Config.DetailLevel
|
||||||
|
state.TTSEnabled = input.TTSEnabled
|
||||||
|
ctx = WithPipelineState(ctx, state)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 前端实现
|
||||||
|
|
||||||
|
#### 1. 情景提示卡片
|
||||||
|
|
||||||
|
**文件**: `frontend/src/components/ChatPanel/index.tsx`
|
||||||
|
|
||||||
|
在对话列表顶部(非空状态 + 非自由对话模式)添加情景提示卡片:
|
||||||
|
|
||||||
|
```tsx
|
||||||
|
{messages.length > 0 && !isFreeChat && (
|
||||||
|
<div className="chat-panel__scenario-hint">
|
||||||
|
<div className="scenario-hint-card">
|
||||||
|
<span className="scenario-hint-card__icon">
|
||||||
|
{scenarios.find(s => s.id === activeScenario)?.icon}
|
||||||
|
</span>
|
||||||
|
<div className="scenario-hint-card__text">
|
||||||
|
<strong>{t(scenarios.find(s => s.id === activeScenario)?.nameKey || "")}</strong>
|
||||||
|
<p>{t(`scenario.${activeScenario}.hint`)}</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
```
|
||||||
|
|
||||||
|
**显示效果**:
|
||||||
|
- 蓝色渐变背景(135deg 从蓝到紫)
|
||||||
|
- 左侧大图标 + 右侧标题和说明
|
||||||
|
- 最大宽度 520px,响应式布局
|
||||||
|
- 柔和阴影和半透明边框
|
||||||
|
|
||||||
|
#### 2. WebSocket 消息发送
|
||||||
|
|
||||||
|
**文件**: `frontend/src/hooks/useVisionSession.ts`
|
||||||
|
|
||||||
|
发送 config 消息时包含 `scenario` 字段:
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
send({
|
||||||
|
type: "config",
|
||||||
|
payload: {
|
||||||
|
tts_enabled: config.ttsEnabled,
|
||||||
|
detail_level: config.detailLevel,
|
||||||
|
language: config.language,
|
||||||
|
scenario: config.scenario, // 情景配置
|
||||||
|
},
|
||||||
|
});
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 3. 样式实现
|
||||||
|
|
||||||
|
**文件**: `frontend/src/App.css`
|
||||||
|
|
||||||
|
情景提示卡片样式:
|
||||||
|
|
||||||
|
```css
|
||||||
|
.scenario-hint-card {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 12px;
|
||||||
|
padding: 12px 16px;
|
||||||
|
border-radius: var(--radius-sm);
|
||||||
|
background: linear-gradient(135deg, rgba(59, 130, 246, 0.08) 0%, rgba(99, 102, 241, 0.08) 100%);
|
||||||
|
border: 1px solid rgba(59, 130, 246, 0.2);
|
||||||
|
box-shadow: 0 2px 8px rgba(59, 130, 246, 0.06);
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 4. 多语言翻译
|
||||||
|
|
||||||
|
**文件**: `frontend/src/lib/i18n/{zh-CN,en-US,ja-JP}.ts`
|
||||||
|
|
||||||
|
新增翻译 key:
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
"scenario.interviewer.hint": "AI 会扮演面试官,逐步提出专业问题并点评你的回答",
|
||||||
|
"scenario.englishTeacher.hint": "AI 会用英语对话,纠正语法错误并引导深入交流",
|
||||||
|
"scenario.debate.hint": "AI 会站在反方立场,用逻辑和证据反驳你的观点",
|
||||||
|
"scenario.interpreter.hint": "AI 会实时翻译你的话(中英互译),无解释评论",
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 数据流
|
||||||
|
|
||||||
|
### WebSocket 协议
|
||||||
|
|
||||||
|
**客户端 → 服务端**(config 消息):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"type": "config",
|
||||||
|
"payload": {
|
||||||
|
"tts_enabled": true,
|
||||||
|
"detail_level": "low",
|
||||||
|
"language": "zh-CN",
|
||||||
|
"scenario": "interviewer"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**服务端 → 客户端**(首句引导):
|
||||||
|
|
||||||
|
```json
|
||||||
|
// llm_chunk
|
||||||
|
{
|
||||||
|
"type": "llm_chunk",
|
||||||
|
"request_id": "scenario_greeting",
|
||||||
|
"delta": "你好!我是今天的面试官...",
|
||||||
|
"role": "assistant"
|
||||||
|
}
|
||||||
|
|
||||||
|
// llm_done
|
||||||
|
{
|
||||||
|
"type": "llm_done",
|
||||||
|
"request_id": "scenario_greeting",
|
||||||
|
"full_text": "你好!我是今天的面试官...",
|
||||||
|
"tokens_used": {"prompt": 0, "completion": 0, "total": 0}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### System Prompt 构建流程
|
||||||
|
|
||||||
|
```
|
||||||
|
sess.Config.Scenario = "interviewer"
|
||||||
|
↓
|
||||||
|
PipelineInput.Scenario = "interviewer"
|
||||||
|
↓
|
||||||
|
PipelineState.Scenario = "interviewer" (adapter.go 复制)
|
||||||
|
↓
|
||||||
|
nodes_history.go 读取 state.Scenario
|
||||||
|
↓
|
||||||
|
scenarioPrompt := llm.GetScenarioPrompt("interviewer", "zh-CN")
|
||||||
|
↓
|
||||||
|
systemPrompt := llm.BuildSystemPrompt(language, detailLevel, scenarioPrompt)
|
||||||
|
↓
|
||||||
|
messages[0] = {Role: "system", Content: systemPrompt}
|
||||||
|
↓
|
||||||
|
ChatModel 接收到情景 Prompt
|
||||||
|
↓
|
||||||
|
LLM 按情景角色生成回复
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 使用指南
|
||||||
|
|
||||||
|
### 快速验证
|
||||||
|
|
||||||
|
1. **打开浏览器** → http://localhost:5173
|
||||||
|
2. **登录系统**
|
||||||
|
3. **切换情景** → 右侧配置面板 → 对话情景 → 模拟面试官
|
||||||
|
4. **观察现象**:
|
||||||
|
- ✨ AI 立即说:"你好!我是今天的面试官。让我们先从自我介绍开始..."
|
||||||
|
- ✨ 对话框顶部显示蓝色提示卡片
|
||||||
|
5. **验证效果** → 发送:"你是谁?"
|
||||||
|
- ✅ **正确回复**:"我是今天的面试官..."
|
||||||
|
- ❌ **错误回复**:"我是通义千问..."
|
||||||
|
|
||||||
|
### 功能测试清单
|
||||||
|
|
||||||
|
| 测试项 | 操作步骤 | 预期结果 |
|
||||||
|
|--------|---------|---------|
|
||||||
|
| **首句引导** | 切换到"模拟面试官" | AI 自动说:"你好!我是今天的面试官..." |
|
||||||
|
| **情景生效** | 问 "你是谁?" | AI 回答:"我是今天的面试官..." |
|
||||||
|
| **提示卡片** | 发送一条消息后查看顶部 | 显示蓝色卡片:"🎯 模拟面试官 \| AI 会扮演面试官..." |
|
||||||
|
| **语言联动** | 切换到"英语老师" | 语言自动切换到 en-US,AI 用英语回复 |
|
||||||
|
| **持久化** | 切换情景后刷新页面 | 情景配置保持,首句仍在历史中 |
|
||||||
|
| **多情景** | 依次测试所有情景 | 每个情景 AI 回复风格明显不同 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 故障排查
|
||||||
|
|
||||||
|
### 如果情景不生效
|
||||||
|
|
||||||
|
1. **检查后端日志**:
|
||||||
|
```bash
|
||||||
|
grep "config updated" /tmp/camtalk_server.log | tail -5
|
||||||
|
grep "历史组装完成" /tmp/camtalk_server.log | tail -5
|
||||||
|
```
|
||||||
|
|
||||||
|
- 如果 `scenario=` 是空的,说明前端未发送或后端未接收
|
||||||
|
- 如果 `scenario=interviewer` 正确,但 AI 回复仍是通用的,可能是 LLM 模型问题
|
||||||
|
|
||||||
|
2. **检查前端 WebSocket 消息**(浏览器 DevTools → Network → WS):
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"type": "config",
|
||||||
|
"payload": {
|
||||||
|
"scenario": "interviewer" // 确认存在
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
3. **检查会话配置是否保存**:
|
||||||
|
- 切换情景后,LocalStorage 中应该有 `camtalk_config`
|
||||||
|
- 内容应包含 `"scenario": "interviewer"`
|
||||||
|
|
||||||
|
4. **清除缓存重试**:
|
||||||
|
```bash
|
||||||
|
# 浏览器:清除 LocalStorage
|
||||||
|
# 后端:重启服务
|
||||||
|
# 前端:刷新页面
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 后续优化建议
|
||||||
|
|
||||||
|
### P2(强烈推荐)
|
||||||
|
|
||||||
|
1. **情景切换时创建新会话**
|
||||||
|
- 避免历史对话干扰新情景
|
||||||
|
- 弹窗确认:"切换情景会创建新会话,当前对话将保存。是否继续?"
|
||||||
|
- 实现难度:⭐⭐
|
||||||
|
- 用户价值:⭐⭐⭐⭐
|
||||||
|
|
||||||
|
2. **进一步增强 System Prompt**
|
||||||
|
- 增加示例对话(Few-shot Prompting)
|
||||||
|
- 增加"禁止事项"列表
|
||||||
|
- 实现难度:⭐
|
||||||
|
- 效果提升:⭐⭐⭐
|
||||||
|
|
||||||
|
### P3(可选)
|
||||||
|
|
||||||
|
1. **情景专属 UI 主题色**
|
||||||
|
- 面试官 → 深蓝色
|
||||||
|
- 英语老师 → 绿色
|
||||||
|
- 辩论 → 红色
|
||||||
|
- 翻译 → 紫色
|
||||||
|
|
||||||
|
2. **切换动画与音效**
|
||||||
|
- 切换时播放短音效
|
||||||
|
- 聊天面板淡出淡入动画
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 修改文件清单
|
||||||
|
|
||||||
|
### 后端(3 个文件)
|
||||||
|
|
||||||
|
- `backend/internal/eino/adapter.go` — 修复 State 初始化
|
||||||
|
- `backend/internal/ws/handler.go` — 添加首句引导
|
||||||
|
- `backend/internal/ai/llm/scenarios.go` — 增强 Prompt + 首句
|
||||||
|
|
||||||
|
### 前端(5 个文件)
|
||||||
|
|
||||||
|
- `frontend/src/hooks/useVisionSession.ts` — 修复 scenario 发送
|
||||||
|
- `frontend/src/components/ChatPanel/index.tsx` — 添加提示卡片
|
||||||
|
- `frontend/src/App.css` — 卡片样式
|
||||||
|
- `frontend/src/lib/i18n/zh-CN.ts` — 中文翻译
|
||||||
|
- `frontend/src/lib/i18n/en-US.ts` — 英文翻译
|
||||||
|
- `frontend/src/lib/i18n/ja-JP.ts` — 日文翻译
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
# 技术名词解释
|
|
||||||
|
|
||||||
对架构文档中技术选型表里出现的所有关键名词的简明解释。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 前端相关
|
|
||||||
|
|
||||||
| 名词 | 一句话 | 展开 |
|
|
||||||
|------|--------|------|
|
|
||||||
| **React 18** | 组件化 UI 框架 | Facebook 开源,把页面拆成组件搭积木拼装。18 版本支持并发渲染。 |
|
|
||||||
| **TypeScript** | 带类型的 JavaScript | 在 JS 基础上增加类型声明,编译阶段就能发现类型错误。 |
|
|
||||||
| **Vite** | 前端构建工具 | 利用浏览器原生 ES Module,开发时毫秒级热更新(HMR),构建产物小。 |
|
|
||||||
| **WebSocket** | 浏览器与服务器的双向通道 | HTTP 是"一问一答",WebSocket 像打电话——接通后双方随时互发消息,适合实时对话场景。 |
|
|
||||||
| **ONNX Runtime Web** | 浏览器端 AI 推理引擎 | 微软定义的通用模型格式 ONNX 的运行引擎,可在浏览器中用 WASM 加速跑轻量模型(如 VAD、关键帧检测),零延迟、不耗服务器资源。 |
|
|
||||||
| **VAD** | 语音活动检测 | Voice Activity Detection,检测"人有没有在说话"。WebRTC 内置了高效的 VAD 算法,本项目用 @ricky0123/vad-web 包装。 |
|
|
||||||
| **MediaDevices API** | 浏览器摄像头/麦克风接口 | `navigator.mediaDevices.getUserMedia()` 是浏览器音视频采集的唯一标准入口,无需插件。 |
|
|
||||||
|
|
||||||
## 后端相关
|
|
||||||
|
|
||||||
| 名词 | 一句话 | 展开 |
|
|
||||||
|------|--------|------|
|
|
||||||
| **Go (Golang)** | 高并发后端语言 | Google 开发,杀手锏是 goroutine——极轻量协程,一个程序可轻松开几万个,每个只占几 KB 内存,适合管理大量 WebSocket 长连接。 |
|
|
||||||
| **gorilla/websocket** | Go WebSocket 库 | Go 标准库无内置 WebSocket 支持,此库是社区最成熟的选择,处理了协议握手、帧解析等底层细节。 |
|
|
||||||
| **Redis** | 内存 KV 数据库 | 数据放在内存里,读写微秒级。本项目用于会话状态和对话上下文缓存,支持 TTL 过期自动清理。多 Gateway 实例通过 Redis 共享状态。 |
|
|
||||||
| **Viper** | Go 配置管理 | 读取 JSON/YAML/TOML 配置,支持环境变量覆盖,方便开发/测试/生产环境用不同配置。 |
|
|
||||||
| **Zap** | Go 结构化日志 | Uber 开源,输出 JSON 格式日志,方便工具搜索分析,性能远超标准库 log。 |
|
|
||||||
|
|
||||||
## AI 服务相关
|
|
||||||
|
|
||||||
| 名词 | 一句话 | 展开 |
|
|
||||||
|------|--------|------|
|
|
||||||
| **多模态 LLM** | 能读文字又能看图片的大语言模型 | GPT-4o(OpenAI)/ Claude Sonnet(Anthropic),给照片+问题能"看懂"照片再回答。 |
|
|
||||||
| **STT** | 语音转文字 | Speech-to-Text。Deepgram 流式识别延迟 <500ms。备选 FunASR(阿里开源,可自部署)。 |
|
|
||||||
| **TTS** | 文字转语音 | Text-to-Speech。OpenAI TTS 音质接近真人。Edge TTS 免费。支持流式——边生成边读,不必等全部生成完。 |
|
|
||||||
| **GPT-4o-mini** | 轻量分类模型 | 又快又便宜的小模型,用于模型路由——先用小模型判断问题复杂度,简单问题走小模型省 API 费用。 |
|
|
||||||
@@ -1,5 +0,0 @@
|
|||||||
1.视频录制
|
|
||||||
2.对话翻译
|
|
||||||
3.对话总结
|
|
||||||
4.手动对话功能
|
|
||||||
5.视频框大小可调整,可最小化然后拖动
|
|
||||||
1111
docs/10-鉴权体系.md
Normal file
1111
docs/10-鉴权体系.md
Normal file
File diff suppressed because it is too large
Load Diff
814
docs/11-令牌桶限流.md
Normal file
814
docs/11-令牌桶限流.md
Normal file
@@ -0,0 +1,814 @@
|
|||||||
|
# 令牌桶限流设计
|
||||||
|
|
||||||
|
## 概述
|
||||||
|
|
||||||
|
CamTalk 采用令牌桶(Token Bucket)算法实现按用户维度的速率限制,核心目标是**控制 AI 调用成本**,同时为 REST API 提供防暴力破解保护。
|
||||||
|
|
||||||
|
**设计原则**:
|
||||||
|
- **成本优先**:主要限流对象是 WebSocket `query` 消息(每次触发 STT + LLM + TTS 完整调用链)
|
||||||
|
- **用户隔离**:Per-user 维度限流,单用户超限不影响其他用户
|
||||||
|
- **弹性突发**:令牌桶允许合理的突发请求,优于固定窗口的滑动限流
|
||||||
|
- **存储适配**:内存 + Redis 双实现,单实例零依赖,多实例分布式一致
|
||||||
|
|
||||||
|
## 整体架构
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph TB
|
||||||
|
subgraph Entry["入口层"]
|
||||||
|
WS["WebSocket Handler<br/>query 消息"]
|
||||||
|
REST["REST API<br/>login / register"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph LimiterModule["Rate Limiter 模块"]
|
||||||
|
Interface["Limiter 接口<br/>Allow(userID) → (bool, retryAfter)"]
|
||||||
|
MemBucket["TokenBucket<br/>内存令牌桶"]
|
||||||
|
RedisBucket["RedisTokenBucket<br/>Redis 令牌桶(Lua 脚本)"]
|
||||||
|
Middleware["RateLimitMiddleware<br/>Gin 中间件"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph Storage["存储层"]
|
||||||
|
MemSync["sync.RWMutex<br/>进程内 map"]
|
||||||
|
Redis["Redis<br/>分布式计数"]
|
||||||
|
end
|
||||||
|
|
||||||
|
WS -->|"限流检查"| Interface
|
||||||
|
REST -->|"中间件"| Middleware
|
||||||
|
Middleware --> Interface
|
||||||
|
Interface --> MemBucket
|
||||||
|
Interface --> RedisBucket
|
||||||
|
MemBucket --> MemSync
|
||||||
|
RedisBucket --> Redis
|
||||||
|
```
|
||||||
|
|
||||||
|
## 令牌桶算法
|
||||||
|
|
||||||
|
### 原理
|
||||||
|
|
||||||
|
令牌桶以固定速率向桶中添加令牌,桶有最大容量上限。每次请求消耗一个令牌,桶空时拒绝请求。
|
||||||
|
|
||||||
|
```
|
||||||
|
桶容量(capacity) = 允许的突发请求数上限
|
||||||
|
填充速率(rate) = 每秒补充的令牌数
|
||||||
|
|
||||||
|
时间线示例(capacity=5, rate=0.2):
|
||||||
|
t=0s 桶满 5 令牌 → 用户连续发 5 个 query 全部通过
|
||||||
|
t=0s 桶空 → 第 6 个 query 被拒绝,retryAfter=5s
|
||||||
|
t=5s 桶补充 1 令牌 → 可再发 1 个 query
|
||||||
|
t=10s 桶补充 1 令牌 → 可再发 1 个 query
|
||||||
|
```
|
||||||
|
|
||||||
|
### 算法公式
|
||||||
|
|
||||||
|
```
|
||||||
|
elapsed = now - lastRefill
|
||||||
|
newTokens = elapsed * rate
|
||||||
|
currentTokens = min(capacity, lastTokens + newTokens)
|
||||||
|
|
||||||
|
if currentTokens >= 1:
|
||||||
|
currentTokens -= 1
|
||||||
|
allowed = true
|
||||||
|
else:
|
||||||
|
allowed = false
|
||||||
|
retryAfter = (1 - currentTokens) / rate
|
||||||
|
```
|
||||||
|
|
||||||
|
## 核心组件
|
||||||
|
|
||||||
|
### 1. Limiter 接口
|
||||||
|
|
||||||
|
**文件位置**:`backend/internal/ratelimit/limiter.go`
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Limiter 速率限制器接口。
|
||||||
|
type Limiter interface {
|
||||||
|
// Allow 判断 key 是否允许执行一次操作。
|
||||||
|
// key 通常为 "userID:action" 格式。
|
||||||
|
// 返回 (allowed, retryAfter)。retryAfter 表示需要等待的时间。
|
||||||
|
Allow(ctx context.Context, key string) (bool, time.Duration)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- key 为字符串,不限定格式,由调用方决定维度(用户 ID、IP 地址等)
|
||||||
|
- 返回 `retryAfter` 供客户端/服务端设置 `Retry-After` header
|
||||||
|
- 接受 `context.Context` 支持超时和取消(Redis 实现需要)
|
||||||
|
|
||||||
|
### 2. 内存令牌桶(TokenBucket)
|
||||||
|
|
||||||
|
**文件位置**:`backend/internal/ratelimit/bucket.go`
|
||||||
|
|
||||||
|
```go
|
||||||
|
// TokenBucket 内存令牌桶,适用于单实例部署。
|
||||||
|
type TokenBucket struct {
|
||||||
|
capacity int // 桶容量
|
||||||
|
rate float64 // 每秒填充令牌数
|
||||||
|
tokens float64 // 当前令牌数
|
||||||
|
lastRefill time.Time // 上次填充时间
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// Limiter 管理多个用户的令牌桶。
|
||||||
|
type Limiter struct {
|
||||||
|
buckets map[string]*TokenBucket
|
||||||
|
config Config
|
||||||
|
mu sync.RWMutex
|
||||||
|
stopOnce sync.Once
|
||||||
|
done chan struct{}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**并发安全**:
|
||||||
|
- 每个桶内部用 `sync.Mutex` 保护读写
|
||||||
|
- 桶 map 用 `sync.RWMutex` 保护(读多写少场景)
|
||||||
|
- 用户首次请求时惰性创建桶
|
||||||
|
|
||||||
|
**内存回收**:
|
||||||
|
- 后台 goroutine 定期扫描,清理超过 10 分钟无活动的桶
|
||||||
|
- 避免长期运行后内存泄漏
|
||||||
|
|
||||||
|
### 3. Redis 令牌桶(RedisTokenBucket)
|
||||||
|
|
||||||
|
**文件位置**:`backend/internal/ratelimit/redis_bucket.go`
|
||||||
|
|
||||||
|
使用 Redis Lua 脚本保证原子性,避免竞态条件:
|
||||||
|
|
||||||
|
```lua
|
||||||
|
-- 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
|
||||||
|
retry_after = (1 - tokens) / rate
|
||||||
|
end
|
||||||
|
|
||||||
|
-- 回写状态
|
||||||
|
redis.call('HMSET', key, 'tokens', tokens, 'last_refill', now)
|
||||||
|
redis.call('EXPIRE', key, ttl)
|
||||||
|
|
||||||
|
return {allowed, tostring(retry_after)}
|
||||||
|
```
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- 每个用户的限流状态存储为一个 Redis Hash(`tokens` + `last_refill`)
|
||||||
|
- TTL 自动过期,无需手动清理
|
||||||
|
- Lua 脚本保证"读取-计算-回写"原子执行
|
||||||
|
|
||||||
|
### 4. Gin 中间件
|
||||||
|
|
||||||
|
**文件位置**:`backend/internal/ratelimit/middleware.go`
|
||||||
|
|
||||||
|
```go
|
||||||
|
// RateLimitMiddleware 返回 Gin 中间件,按 key 维度限流。
|
||||||
|
// keyFunc 从请求中提取限流 key(如 IP、用户 ID)。
|
||||||
|
func RateLimitMiddleware(limiter Limiter, keyFunc func(*gin.Context) string) gin.HandlerFunc
|
||||||
|
```
|
||||||
|
|
||||||
|
**使用方式**:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 按 IP 限流(登录/注册,未登录用户无 userID)
|
||||||
|
loginGroup.POST("/login",
|
||||||
|
ratelimit.Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return c.ClientIP() + ":login"
|
||||||
|
}),
|
||||||
|
authHandler.Login,
|
||||||
|
)
|
||||||
|
|
||||||
|
// 按用户 ID 限流(已认证的 API)
|
||||||
|
authorized.POST("/conversations",
|
||||||
|
ratelimit.Middleware(limiter, func(c *gin.Context) string {
|
||||||
|
return c.GetString("user_id") + ":conversation"
|
||||||
|
}),
|
||||||
|
convHandler.Create,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
**错误响应**:
|
||||||
|
|
||||||
|
REST API 返回 HTTP 429:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"code": "RATE_LIMITED",
|
||||||
|
"message": "too many requests, retry after 5s"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
同时设置 `Retry-After` header:
|
||||||
|
|
||||||
|
```
|
||||||
|
HTTP/1.1 429 Too Many Requests
|
||||||
|
Retry-After: 5
|
||||||
|
```
|
||||||
|
|
||||||
|
## 限流接入点
|
||||||
|
|
||||||
|
### WebSocket query 消息(核心)
|
||||||
|
|
||||||
|
在 `ws/handler.go` 的 `case "query"` 分支中,orchestrator 调用前检查:
|
||||||
|
|
||||||
|
```go
|
||||||
|
case "query":
|
||||||
|
// ... 解析消息 ...
|
||||||
|
|
||||||
|
// 限流检查
|
||||||
|
if limiter != nil {
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, userID+":query")
|
||||||
|
if !allowed {
|
||||||
|
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
|
||||||
|
fmt.Errorf("rate limited, retry after %s", retryAfter))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ... 继续处理 query ...
|
||||||
|
```
|
||||||
|
|
||||||
|
### REST API 登录/注册
|
||||||
|
|
||||||
|
在 `api/auth.go` 的路由注册中添加中间件:
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup, limiter ratelimit.Limiter) {
|
||||||
|
auth := rg.Group("/auth")
|
||||||
|
if limiter != nil {
|
||||||
|
auth.POST("/register",
|
||||||
|
ratelimit.Middleware(limiter, ipKeyFunc("register")),
|
||||||
|
h.Register,
|
||||||
|
)
|
||||||
|
auth.POST("/login",
|
||||||
|
ratelimit.Middleware(limiter, ipKeyFunc("login")),
|
||||||
|
h.Login,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
auth.POST("/register", h.Register)
|
||||||
|
auth.POST("/login", h.Login)
|
||||||
|
}
|
||||||
|
auth.POST("/refresh", h.Refresh)
|
||||||
|
auth.POST("/logout", h.Logout)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 不限流的端点
|
||||||
|
|
||||||
|
| 端点 | 原因 |
|
||||||
|
|------|------|
|
||||||
|
| `ping` / `pong` | 心跳保活,无 AI 调用成本 |
|
||||||
|
| `config` | 配置更新,无 AI 调用成本 |
|
||||||
|
| `interrupt` | 中断请求,取消操作不应被限流 |
|
||||||
|
| `GET /api/health` | 健康检查,运维必需 |
|
||||||
|
| `POST /api/auth/refresh` | Token 刷新,限流会导致用户被迫重新登录 |
|
||||||
|
| `POST /api/auth/logout` | 登出,限流会导致用户无法正常退出 |
|
||||||
|
| `GET /api/conversations` | 查询列表,无 AI 调用成本 |
|
||||||
|
|
||||||
|
## 配置设计
|
||||||
|
|
||||||
|
### 配置文件
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# backend/config/config.yaml 新增
|
||||||
|
ratelimit:
|
||||||
|
enabled: true
|
||||||
|
# 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 次
|
||||||
|
```
|
||||||
|
|
||||||
|
### 配置结构体
|
||||||
|
|
||||||
|
```go
|
||||||
|
// config/config.go 新增
|
||||||
|
|
||||||
|
type RateLimitConfig struct {
|
||||||
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
Query BucketConfig `mapstructure:"query"`
|
||||||
|
Login BucketConfig `mapstructure:"login"`
|
||||||
|
Register BucketConfig `mapstructure:"register"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type BucketConfig struct {
|
||||||
|
Capacity int `mapstructure:"capacity"` // 桶容量(突发上限)
|
||||||
|
Rate float64 `mapstructure:"rate"` // 每秒填充令牌数
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 默认值
|
||||||
|
|
||||||
|
```go
|
||||||
|
// setDefaults 新增
|
||||||
|
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)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 参数选择建议
|
||||||
|
|
||||||
|
| 场景 | capacity | rate | 含义 |
|
||||||
|
|------|----------|------|------|
|
||||||
|
| WebSocket query | 10 | 0.2 | 突发 10 个,之后每 5 秒 1 个 |
|
||||||
|
| 登录 | 5 | 0.1 | 突发 5 次,之后每 10 秒 1 次 |
|
||||||
|
| 注册 | 3 | 0.05 | 突发 3 次,之后每 20 秒 1 次 |
|
||||||
|
|
||||||
|
> **调参原则**:capacity 决定"能忍多少次突发",rate 决定"稳态下多久能再请求一次"。query 的 rate 建议根据 AI 调用成本和目标月预算反推。
|
||||||
|
|
||||||
|
## 依赖注入
|
||||||
|
|
||||||
|
### main.go 初始化
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 初始化限流器
|
||||||
|
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()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 注入到各模块
|
||||||
|
|
||||||
|
```go
|
||||||
|
// WebSocket Handler —— 新增 limiter 参数
|
||||||
|
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter))
|
||||||
|
|
||||||
|
// Auth REST —— 新增 limiter 参数
|
||||||
|
authHandler := api.NewAuthHandler(authService, tokenMgr)
|
||||||
|
authHandler.RegisterRoutes(apiGroup, limiter)
|
||||||
|
```
|
||||||
|
|
||||||
|
## 错误码
|
||||||
|
|
||||||
|
复用已有错误码 `RATE_LIMITED`(`backend/internal/errors/codes.go`):
|
||||||
|
|
||||||
|
| 传输层 | HTTP 状态码 | 错误格式 |
|
||||||
|
|--------|-----------|---------|
|
||||||
|
| REST API | 429 Too Many Requests | `{code: "RATE_LIMITED", message: "too many requests, retry after Xs"}` |
|
||||||
|
| WebSocket | — | `{type: "error", code: "RATE_LIMITED", request_id: "...", message: "..."}` |
|
||||||
|
|
||||||
|
## 文件结构
|
||||||
|
|
||||||
|
```
|
||||||
|
backend/internal/ratelimit/
|
||||||
|
├── limiter.go # Limiter 接口 + Config 类型定义
|
||||||
|
├── bucket.go # 内存令牌桶实现
|
||||||
|
├── bucket_test.go # 内存令牌桶单元测试
|
||||||
|
├── redis_bucket.go # Redis 令牌桶实现(Lua 脚本)
|
||||||
|
├── redis_bucket_test.go# Redis 令牌桶单元测试
|
||||||
|
└── middleware.go # Gin 中间件
|
||||||
|
```
|
||||||
|
|
||||||
|
## 测试用例
|
||||||
|
|
||||||
|
### 单元测试
|
||||||
|
|
||||||
|
**内存令牌桶**(`bucket_test.go`):
|
||||||
|
- 首次请求通过
|
||||||
|
- 连续消耗至桶空
|
||||||
|
- 桶空后拒绝,返回正确 retryAfter
|
||||||
|
- 等待后令牌补充,请求通过
|
||||||
|
- 并发安全性(多个 goroutine 同时 Allow)
|
||||||
|
- 桶容量边界(capacity=0, capacity=1)
|
||||||
|
- 填充速率边界(rate=0, rate 极大值)
|
||||||
|
- 不活跃桶的内存回收
|
||||||
|
|
||||||
|
**Redis 令牌桶**(`redis_bucket_test.go`):
|
||||||
|
- 与内存实现行为一致性
|
||||||
|
- Lua 脚本原子性
|
||||||
|
- key TTL 自动过期
|
||||||
|
- 并发安全性(多个客户端同时请求)
|
||||||
|
|
||||||
|
### 集成测试
|
||||||
|
|
||||||
|
- 限流关闭时不拦截请求
|
||||||
|
- 限流开启后,REST API 登录超限返回 429
|
||||||
|
- 限流开启后,WebSocket query 超限返回 `RATE_LIMITED` 错误
|
||||||
|
- 单实例内存限流 vs 多实例 Redis 限流行为一致
|
||||||
|
- 重启后内存限流重置,Redis 限流保持
|
||||||
|
|
||||||
|
## 扩展点
|
||||||
|
|
||||||
|
### 1. 多级限流
|
||||||
|
|
||||||
|
可扩展为多级限流策略:
|
||||||
|
|
||||||
|
```
|
||||||
|
全局限流(全用户共享) → 用户级限流(当前实现) → 端点级限流(不同 API 不同限制)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. 动态调参
|
||||||
|
|
||||||
|
通过配置热更新或管理 API 动态调整限流参数,无需重启:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 预留接口
|
||||||
|
type DynamicLimiter interface {
|
||||||
|
Limiter
|
||||||
|
UpdateConfig(action string, cfg BucketConfig) error
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. 按用户等级差异化
|
||||||
|
|
||||||
|
不同用户等级使用不同的限流参数:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
ratelimit:
|
||||||
|
query:
|
||||||
|
capacity: 10 # 免费用户
|
||||||
|
rate: 0.2
|
||||||
|
query_premium:
|
||||||
|
capacity: 30 # 付费用户
|
||||||
|
rate: 1.0
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4. 滑动窗口限流
|
||||||
|
|
||||||
|
令牌桶适合允许突发的场景。如果需要更平滑的限流,可增加滑动窗口实现:
|
||||||
|
|
||||||
|
```go
|
||||||
|
type SlidingWindowLimiter struct {
|
||||||
|
windowSize time.Duration
|
||||||
|
maxRequests int
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5. 分布式全局限流
|
||||||
|
|
||||||
|
当前 Redis 实现是 Per-Instance 独立计数。如需全局精确限流,可改为 Redis 全局计数器(所有实例共享同一个 key)。
|
||||||
|
|
||||||
|
## 监控指标
|
||||||
|
|
||||||
|
### 关键指标
|
||||||
|
|
||||||
|
- **限流触发率**:被拒绝请求数 / 总请求数
|
||||||
|
- **各端点限流分布**:query / login / register 各自的触发率
|
||||||
|
- **等待时长分布**:retryAfter 的 P50/P99
|
||||||
|
- **桶状态**:各用户桶的平均令牌数(反映使用模式)
|
||||||
|
|
||||||
|
### 告警规则
|
||||||
|
|
||||||
|
- **限流触发率突增**:可能表示异常流量或攻击
|
||||||
|
- **单用户持续被限流**:可能表示客户端 bug(死循环请求)
|
||||||
|
|
||||||
|
## 实际实现要点
|
||||||
|
|
||||||
|
### 文件结构
|
||||||
|
|
||||||
|
```
|
||||||
|
backend/internal/ratelimit/
|
||||||
|
├── limiter.go # Limiter 接口定义
|
||||||
|
├── bucket.go # 内存令牌桶实现 (MemoryLimiter + TokenBucket)
|
||||||
|
├── bucket_test.go # 内存令牌桶单元测试(11 个测试用例)
|
||||||
|
├── redis_bucket.go # Redis 令牌桶实现(Lua 脚本)
|
||||||
|
├── redis_bucket_test.go # Redis 令牌桶单元测试
|
||||||
|
└── middleware.go # Gin 中间件实现
|
||||||
|
```
|
||||||
|
|
||||||
|
### TokenBucket 实现细节
|
||||||
|
|
||||||
|
**核心数据结构**(`bucket.go:12-18`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
type TokenBucket struct {
|
||||||
|
capacity int // 桶容量
|
||||||
|
rate float64 // 每秒填充令牌数
|
||||||
|
tokens float64 // 当前令牌数(浮点数支持小数令牌)
|
||||||
|
lastRefill time.Time // 上次填充时间
|
||||||
|
mu sync.Mutex // 保护并发访问
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**并发安全**(`bucket.go:31-56`):
|
||||||
|
- 每个桶内部使用 `sync.Mutex` 保护 `tokens` 和 `lastRefill` 字段
|
||||||
|
- `allow()` 方法的"读取-计算-回写"操作原子执行
|
||||||
|
- 桶 map 使用 `sync.RWMutex` 保护,读多写少优化(`bucket.go:62`)
|
||||||
|
- 双重检查锁(`bucket.go:106-113`):先尝试读锁获取桶,不存在时升级写锁创建
|
||||||
|
|
||||||
|
**内存回收机制**(`bucket.go:132-161`):
|
||||||
|
- 后台 goroutine 每 10 分钟扫描一次(`cleanup()` 方法)
|
||||||
|
- 删除超过 10 分钟无活动的桶(`lastRefill` 超时判断)
|
||||||
|
- 通过 `done` channel 和 `sync.Once` 保证优雅停止
|
||||||
|
|
||||||
|
**惰性创建**(`bucket.go:95-121`):
|
||||||
|
- 用户首次请求时才创建桶,避免预分配内存
|
||||||
|
- `getOrCreateBucket()` 使用读写锁分离,优化热路径性能
|
||||||
|
|
||||||
|
### Gin 中间件实现
|
||||||
|
|
||||||
|
**实际代码**(`middleware.go:12-42`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
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()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- `nil` limiter 自动跳过限流(支持配置关闭)
|
||||||
|
- 空 key 跳过限流(支持匿名端点)
|
||||||
|
- `retryAfter` 向上取整到秒(符合 HTTP 标准)
|
||||||
|
- `c.Abort()` 阻止后续 handler 执行
|
||||||
|
|
||||||
|
### WebSocket 限流接入
|
||||||
|
|
||||||
|
**实际接入点**(`internal/ws/handler.go:230-240`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
case "query":
|
||||||
|
// ... 解析消息 ...
|
||||||
|
|
||||||
|
// 限流检查
|
||||||
|
if limiter != nil {
|
||||||
|
key := fmt.Sprintf("%s:query", userID)
|
||||||
|
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||||
|
if !allowed {
|
||||||
|
// 限流触发时自动记录 Warn 日志(在 limiter 内部使用 trace.FromContext)
|
||||||
|
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
|
||||||
|
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ... 继续处理 query ...
|
||||||
|
```
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- key 格式:`userID:query`(用户级限流)
|
||||||
|
- 拒绝时发送 `RATE_LIMITED` 错误到客户端
|
||||||
|
- 限流触发时 `RedisLimiter.Allow` 内部自动记录 Warn 日志(带 trace_id,详见 `docs/13-日志追踪.md`)
|
||||||
|
- 不阻塞其他消息类型(`ping`/`config`/`interrupt` 不限流)
|
||||||
|
|
||||||
|
### 配置加载与依赖注入
|
||||||
|
|
||||||
|
**配置文件路径**:
|
||||||
|
- 基础配置:`backend/config/config.yaml`
|
||||||
|
- 开发环境:`backend/config/config.dev.yaml`
|
||||||
|
- 生产环境:`backend/config/config.prod.yaml`
|
||||||
|
|
||||||
|
**实际配置示例**(`config.yaml:63-76`):
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
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 次
|
||||||
|
```
|
||||||
|
|
||||||
|
**依赖注入实现**(`cmd/server/main.go:200-214`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 初始化限流器
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**自动选择策略**:
|
||||||
|
1. 配置关闭(`enabled: false`)→ `limiter = nil`(完全跳过限流)
|
||||||
|
2. Redis 可用 → `NewRedisLimiter`(分布式一致)
|
||||||
|
3. Redis 不可用 → `NewMemoryLimiter`(单实例零依赖)
|
||||||
|
|
||||||
|
**注入到模块**:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// WebSocket Handler
|
||||||
|
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter))
|
||||||
|
|
||||||
|
// REST API
|
||||||
|
authHandler.RegisterRoutes(apiGroup, limiter)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 实际测试用例
|
||||||
|
|
||||||
|
**内存令牌桶测试**(`bucket_test.go`,11 个用例):
|
||||||
|
|
||||||
|
| 测试用例 | 验证内容 |
|
||||||
|
|---------|---------|
|
||||||
|
| `TestTokenBucket_Allow_FirstRequest` | 首次请求通过 |
|
||||||
|
| `TestTokenBucket_Allow_ConsumeUntilEmpty` | 连续消耗至桶空 |
|
||||||
|
| `TestTokenBucket_Allow_RetryAfterCorrect` | `retryAfter` 计算准确性 |
|
||||||
|
| `TestTokenBucket_Allow_RefillAfterWait` | 等待后令牌补充 |
|
||||||
|
| `TestTokenBucket_Allow_CapacityLimit` | 桶容量上限限制 |
|
||||||
|
| `TestTokenBucket_Allow_ConcurrentSafe` | 100 并发请求正确性 |
|
||||||
|
| `TestTokenBucket_Allow_ZeroCapacity` | 边界:`capacity=0` |
|
||||||
|
| `TestTokenBucket_Allow_ZeroRate` | 边界:`rate=0` |
|
||||||
|
| `TestMemoryLimiter_Allow_DifferentKeys` | 不同用户隔离 |
|
||||||
|
| `TestMemoryLimiter_Cleanup` | 不活跃桶自动清理 |
|
||||||
|
| `TestMemoryLimiter_Stop` | 多次 `Stop()` 不 panic |
|
||||||
|
|
||||||
|
**中间件测试**(`middleware_test.go`,7 个用例):
|
||||||
|
|
||||||
|
| 测试用例 | 验证内容 |
|
||||||
|
|---------|---------|
|
||||||
|
| `TestMiddleware_Allow` | 允许时正常响应 |
|
||||||
|
| `TestMiddleware_Deny` | 拒绝时返回 429 + `Retry-After` header |
|
||||||
|
| `TestMiddleware_NilLimiter` | `nil` limiter 放行 |
|
||||||
|
| `TestMiddleware_EmptyKey` | 空 key 放行 |
|
||||||
|
| `TestMiddleware_KeyFunc` | `keyFunc` 正确提取 key |
|
||||||
|
| `TestMiddleware_RetryAfterRounding` | `retryAfter` 向上取整 |
|
||||||
|
|
||||||
|
**并发安全性验证**(`bucket_test.go:83-107`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Redis Lua 脚本实现
|
||||||
|
|
||||||
|
**实际脚本**(`redis_bucket.go:15-53`):
|
||||||
|
|
||||||
|
```lua
|
||||||
|
-- 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)}
|
||||||
|
```
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- 使用 Hash 存储两个字段:`tokens`(当前令牌数)+ `last_refill`(上次填充时间)
|
||||||
|
- 原子性:整个脚本在 Redis 单线程中执行,无竞态条件
|
||||||
|
- 自动过期:每次操作设置 TTL(默认 10 分钟),无需手动清理
|
||||||
|
- 与内存实现算法一致(便于单元测试验证行为等价性)
|
||||||
|
|
||||||
|
### 编译期接口检查
|
||||||
|
|
||||||
|
**接口契约**(`bucket.go:172`,`middleware_test.go:31`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
// 确保 MemoryLimiter 实现了 Limiter 接口
|
||||||
|
var _ Limiter = (*MemoryLimiter)(nil)
|
||||||
|
|
||||||
|
// 确保 mockLimiter 实现了 Limiter 接口
|
||||||
|
var _ Limiter = (*mockLimiter)(nil)
|
||||||
|
```
|
||||||
|
|
||||||
|
编译器会在类型不匹配时报错,避免运行时接口错误。
|
||||||
|
|
||||||
|
### 环境变量覆盖
|
||||||
|
|
||||||
|
配置文件中的 `ratelimit` 配置可通过环境变量覆盖:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CAMTALK_RATELIMIT_ENABLED=true
|
||||||
|
export CAMTALK_RATELIMIT_QUERY_CAPACITY=20
|
||||||
|
export CAMTALK_RATELIMIT_QUERY_RATE=0.5
|
||||||
|
```
|
||||||
|
|
||||||
|
环境变量优先级高于配置文件(Viper 配置绑定)。
|
||||||
|
|
||||||
|
## 参考资料
|
||||||
|
|
||||||
|
- [Token Bucket 算法](https://en.wikipedia.org/wiki/Token_bucket)
|
||||||
|
- [Redis Rate Limiting](https://redis.io/glossaries/rate-limiting/)
|
||||||
|
- [Cloudflare - How we built rate limiting capable of scaling to millions of domains](https://blog.cloudflare.com/counting-things-a-lot-of-different-things/)
|
||||||
@@ -1,750 +0,0 @@
|
|||||||
# 持久化与用户系统设计
|
|
||||||
|
|
||||||
## 概述
|
|
||||||
|
|
||||||
本文档定义用户注册/登录、JWT 认证、对话历史持久化的完整设计方案。核心目标:**用户登录后可在对话列表中选择历史对话继续交谈**。
|
|
||||||
|
|
||||||
### 设计决策
|
|
||||||
|
|
||||||
| 决策项 | 选择 | 理由 |
|
|
||||||
|--------|------|------|
|
|
||||||
| 认证方式 | JWT(access + refresh 双 token) | 无状态,适合分布式部署 |
|
|
||||||
| 注册方式 | 用户名 + 密码 | MVP 最简方案 |
|
|
||||||
| 密码存储 | bcrypt hash | 行业标准,抗彩虹表 |
|
|
||||||
| 对话恢复 | 对话列表选择 | 用户可见所有历史对话,自主选择继续或新建 |
|
|
||||||
| 对话标题 | 自动取首条用户消息前 20 字符 | 零成本,自然可读 |
|
|
||||||
| 图像持久化 | 不存储 | 节省空间,文字历史已足够 |
|
|
||||||
| 登录后行为 | 先选对话,再进聊天 | 明确的入口,避免困惑 |
|
|
||||||
| WS 认证 | URL query 参数 `?token=xxx` | HTTP Upgrade 无法带 Authorization header |
|
|
||||||
| Token 策略 | access 15min + refresh 7day | 安全性与体验平衡 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 一、数据库设计
|
|
||||||
|
|
||||||
### 1.1 ER 关系
|
|
||||||
|
|
||||||
```
|
|
||||||
users 1──N sessions 1──N messages
|
|
||||||
│
|
|
||||||
└── refresh_tokens (1──N, token 轮转管理)
|
|
||||||
```
|
|
||||||
|
|
||||||
### 1.2 表结构
|
|
||||||
|
|
||||||
```sql
|
|
||||||
-- 用户表
|
|
||||||
CREATE TABLE users (
|
|
||||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
|
||||||
username VARCHAR(64) NOT NULL UNIQUE,
|
|
||||||
password_hash VARCHAR(256) NOT NULL, -- bcrypt hash
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now(),
|
|
||||||
updated_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE INDEX idx_users_username ON users(username);
|
|
||||||
|
|
||||||
-- 会话(对话)表
|
|
||||||
CREATE TABLE sessions (
|
|
||||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
|
||||||
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
||||||
title VARCHAR(128) DEFAULT '新对话',
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now(),
|
|
||||||
updated_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE INDEX idx_sessions_user_id ON sessions(user_id, updated_at DESC);
|
|
||||||
|
|
||||||
-- 消息表
|
|
||||||
CREATE TABLE messages (
|
|
||||||
id BIGSERIAL PRIMARY KEY,
|
|
||||||
session_id UUID NOT NULL REFERENCES sessions(id) ON DELETE CASCADE,
|
|
||||||
role VARCHAR(16) NOT NULL, -- "user" | "assistant"
|
|
||||||
content TEXT NOT NULL,
|
|
||||||
tokens_used INTEGER DEFAULT 0,
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE INDEX idx_messages_session_id ON messages(session_id, id);
|
|
||||||
|
|
||||||
-- 刷新令牌表
|
|
||||||
CREATE TABLE refresh_tokens (
|
|
||||||
id BIGSERIAL PRIMARY KEY,
|
|
||||||
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
||||||
token_hash VARCHAR(256) NOT NULL UNIQUE, -- SHA256(refresh_token)
|
|
||||||
expires_at TIMESTAMPTZ NOT NULL,
|
|
||||||
created_at TIMESTAMPTZ DEFAULT now()
|
|
||||||
);
|
|
||||||
|
|
||||||
CREATE INDEX idx_refresh_tokens_user ON refresh_tokens(user_id);
|
|
||||||
CREATE INDEX idx_refresh_tokens_hash ON refresh_tokens(token_hash);
|
|
||||||
```
|
|
||||||
|
|
||||||
### 1.3 与现有设计的差异
|
|
||||||
|
|
||||||
| 变更 | 原设计(`02-系统架构.md`) | 新设计 | 理由 |
|
|
||||||
|------|--------------------------|--------|------|
|
|
||||||
| `sessions.user_id` | `NOT NULL` 无外键 | `REFERENCES users(id) ON DELETE CASCADE` | 关联用户,级联删除 |
|
|
||||||
| `sessions.title` | 无 | `VARCHAR(128) DEFAULT '新对话'` | 对话列表展示 |
|
|
||||||
| `messages.image_url` | 有 | 移除 | 不存储图像 |
|
|
||||||
| `usage_daily` | 有 | MVP 暂不实现 | 按需后加 |
|
|
||||||
| 新增 `users` | 无 | 新增 | 用户系统核心 |
|
|
||||||
| 新增 `refresh_tokens` | 无 | 新增 | JWT refresh 机制 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 二、JWT 认证设计
|
|
||||||
|
|
||||||
### 2.1 Token 结构
|
|
||||||
|
|
||||||
**access_token**:
|
|
||||||
- payload: `{user_id, username, exp (15min), iat, iss: "camtalk"}`
|
|
||||||
- 签名算法: HS256(对称密钥,从配置读取)
|
|
||||||
- 存储位置: 前端 localStorage
|
|
||||||
|
|
||||||
**refresh_token**:
|
|
||||||
- payload: `{user_id, token_id (UUID), exp (7day), iat, iss: "camtalk"}`
|
|
||||||
- 存储位置: 前端 localStorage + 数据库 `refresh_tokens` 表(存 SHA256 hash)
|
|
||||||
|
|
||||||
### 2.2 认证流程
|
|
||||||
|
|
||||||
#### 注册
|
|
||||||
|
|
||||||
```
|
|
||||||
用户 ──POST /api/auth/register──> 检查 username 唯一性
|
|
||||||
bcrypt hash 密码
|
|
||||||
INSERT users
|
|
||||||
↓
|
|
||||||
生成 access_token + refresh_token
|
|
||||||
存 SHA256(refresh_token) 到 DB
|
|
||||||
↓
|
|
||||||
返回 {user, access_token, refresh_token}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 登录
|
|
||||||
|
|
||||||
```
|
|
||||||
用户 ──POST /api/auth/login──> 查 users 表 by username
|
|
||||||
bcrypt.CompareHashAndPassword
|
|
||||||
↓
|
|
||||||
生成 access_token + refresh_token
|
|
||||||
存 SHA256(refresh_token) 到 DB
|
|
||||||
↓
|
|
||||||
返回 {user, access_token, refresh_token}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 刷新
|
|
||||||
|
|
||||||
```
|
|
||||||
用户 ──POST /api/auth/refresh──> 校验 refresh_token 签名和过期
|
|
||||||
查 DB 验证 hash 存在
|
|
||||||
↓
|
|
||||||
撤销旧 refresh_token(DELETE)
|
|
||||||
生成新的 access + refresh
|
|
||||||
存新 refresh_token hash
|
|
||||||
↓
|
|
||||||
返回 {access_token, refresh_token}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 登出
|
|
||||||
|
|
||||||
```
|
|
||||||
用户 ──POST /api/auth/logout──> 撤销 refresh_token (DELETE from DB)
|
|
||||||
前端清除 localStorage
|
|
||||||
```
|
|
||||||
|
|
||||||
### 2.3 Go 实现接口
|
|
||||||
|
|
||||||
```go
|
|
||||||
// internal/auth/jwt.go
|
|
||||||
|
|
||||||
type Claims struct {
|
|
||||||
UserID string `json:"user_id"`
|
|
||||||
Username string `json:"username"`
|
|
||||||
jwt.RegisteredClaims
|
|
||||||
}
|
|
||||||
|
|
||||||
type TokenManager struct {
|
|
||||||
secret []byte
|
|
||||||
accessTTL time.Duration // 15min
|
|
||||||
refreshTTL time.Duration // 7day
|
|
||||||
}
|
|
||||||
|
|
||||||
// GeneratePair 生成 access + refresh token 对。
|
|
||||||
func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh string, err error)
|
|
||||||
|
|
||||||
// ValidateAccess 校验 access_token,返回 Claims。
|
|
||||||
func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error)
|
|
||||||
|
|
||||||
// ValidateRefresh 校验 refresh_token 签名和过期(不查 DB,DB 校验由 service 层负责)。
|
|
||||||
func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error)
|
|
||||||
|
|
||||||
// HashToken 计算 token 的 SHA256 hash(用于 DB 存储)。
|
|
||||||
func HashToken(token string) string
|
|
||||||
```
|
|
||||||
|
|
||||||
```go
|
|
||||||
// internal/auth/middleware.go
|
|
||||||
|
|
||||||
// AuthMiddleware Gin 中间件:从 Authorization: Bearer <token> 提取并校验。
|
|
||||||
// 校验通过后将 Claims 写入 gin.Context。
|
|
||||||
func AuthMiddleware(tm *TokenManager) gin.HandlerFunc {
|
|
||||||
return func(c *gin.Context) {
|
|
||||||
auth := c.GetHeader("Authorization")
|
|
||||||
if !strings.HasPrefix(auth, "Bearer ") {
|
|
||||||
c.AbortWithStatusJSON(401, gin.H{"error": "missing token"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
claims, err := tm.ValidateAccess(strings.TrimPrefix(auth, "Bearer "))
|
|
||||||
if err != nil {
|
|
||||||
c.AbortWithStatusJSON(401, gin.H{"error": "invalid token"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
c.Set("claims", claims)
|
|
||||||
c.Set("user_id", claims.UserID)
|
|
||||||
c.Next()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 三、REST API 设计
|
|
||||||
|
|
||||||
### 3.1 认证 API(新增)
|
|
||||||
|
|
||||||
#### 注册
|
|
||||||
|
|
||||||
```
|
|
||||||
POST /api/auth/register
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{"username": "alice", "password": "s3cret123"}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 201 Created
|
|
||||||
{
|
|
||||||
"user": {"id": "uuid", "username": "alice", "created_at": "2026-06-14T10:00:00Z"},
|
|
||||||
"access_token": "eyJ...",
|
|
||||||
"refresh_token": "eyJ..."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
错误码:`USERNAME_TAKEN`(409)、`INVALID_INPUT`(400,用户名/密码格式不合规)
|
|
||||||
|
|
||||||
#### 登录
|
|
||||||
|
|
||||||
```
|
|
||||||
POST /api/auth/login
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{"username": "alice", "password": "s3cret123"}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 200 OK
|
|
||||||
{
|
|
||||||
"user": {"id": "uuid", "username": "alice"},
|
|
||||||
"access_token": "eyJ...",
|
|
||||||
"refresh_token": "eyJ..."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
错误码:`INVALID_CREDENTIALS`(401)
|
|
||||||
|
|
||||||
#### 刷新 Token
|
|
||||||
|
|
||||||
```
|
|
||||||
POST /api/auth/refresh
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{"refresh_token": "eyJ..."}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 200 OK
|
|
||||||
{
|
|
||||||
"access_token": "eyJ...",
|
|
||||||
"refresh_token": "eyJ..."
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
错误码:`INVALID_TOKEN`(401)
|
|
||||||
|
|
||||||
#### 登出
|
|
||||||
|
|
||||||
```
|
|
||||||
POST /api/auth/logout
|
|
||||||
Authorization: Bearer <access_token>
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{"refresh_token": "eyJ..."}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:`204 No Content`
|
|
||||||
|
|
||||||
### 3.2 对话管理 API(新增)
|
|
||||||
|
|
||||||
所有端点需要 `Authorization: Bearer <access_token>` header。
|
|
||||||
|
|
||||||
#### 获取对话列表
|
|
||||||
|
|
||||||
```
|
|
||||||
GET /api/conversations?page=1&size=20
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 200 OK
|
|
||||||
{
|
|
||||||
"conversations": [
|
|
||||||
{
|
|
||||||
"id": "uuid",
|
|
||||||
"title": "这是一朵红色的玫瑰花",
|
|
||||||
"last_message": "它看起来很美丽。",
|
|
||||||
"message_count": 6,
|
|
||||||
"updated_at": "2026-06-14T10:30:00Z"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"total": 42,
|
|
||||||
"page": 1,
|
|
||||||
"size": 20
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 创建新对话
|
|
||||||
|
|
||||||
```
|
|
||||||
POST /api/conversations
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 201 Created
|
|
||||||
{
|
|
||||||
"id": "uuid",
|
|
||||||
"title": "新对话",
|
|
||||||
"created_at": "2026-06-14T10:00:00Z"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 获取对话详情
|
|
||||||
|
|
||||||
```
|
|
||||||
GET /api/conversations/:id
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 200 OK
|
|
||||||
{
|
|
||||||
"id": "uuid",
|
|
||||||
"title": "这是一朵红色的玫瑰花",
|
|
||||||
"created_at": "2026-06-14T10:00:00Z",
|
|
||||||
"updated_at": "2026-06-14T10:30:00Z",
|
|
||||||
"config": {"tts_enabled": true, "detail_level": "low", "language": "zh-CN"}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 更新对话标题
|
|
||||||
|
|
||||||
```
|
|
||||||
PATCH /api/conversations/:id
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{"title": "新的标题"}
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:`200 OK` + 更新后的对话详情
|
|
||||||
|
|
||||||
#### 删除对话
|
|
||||||
|
|
||||||
```
|
|
||||||
DELETE /api/conversations/:id
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:`204 No Content`(级联删除 messages)
|
|
||||||
|
|
||||||
#### 获取对话历史消息
|
|
||||||
|
|
||||||
```
|
|
||||||
GET /api/conversations/:id/messages?limit=50&before=<message_id>
|
|
||||||
```
|
|
||||||
|
|
||||||
响应:
|
|
||||||
|
|
||||||
```json
|
|
||||||
// 200 OK
|
|
||||||
{
|
|
||||||
"messages": [
|
|
||||||
{"id": 1, "role": "user", "content": "这是什么花?", "created_at": "..."},
|
|
||||||
{"id": 2, "role": "assistant", "content": "这是一朵红色的玫瑰。", "tokens_used": 42, "created_at": "..."}
|
|
||||||
],
|
|
||||||
"has_more": false
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 3.3 现有 API 变更
|
|
||||||
|
|
||||||
| 端点 | 变更 |
|
|
||||||
|------|------|
|
|
||||||
| `GET /api/health` | 不变 |
|
|
||||||
| `POST /api/sessions` | **废弃**,使用 `POST /api/conversations` 替代 |
|
|
||||||
| `DELETE /api/sessions/{id}` | **废弃**,使用 `DELETE /api/conversations/:id` 替代 |
|
|
||||||
|
|
||||||
### 3.4 新增错误码
|
|
||||||
|
|
||||||
| 错误码 | HTTP 状态 | 含义 |
|
|
||||||
|--------|-----------|------|
|
|
||||||
| `USERNAME_TAKEN` | 409 | 用户名已被注册 |
|
|
||||||
| `INVALID_CREDENTIALS` | 401 | 用户名或密码错误 |
|
|
||||||
| `INVALID_TOKEN` | 401 | JWT 无效或已过期 |
|
|
||||||
| `INVALID_INPUT` | 400 | 请求参数不合规(用户名/密码长度等) |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 四、Session Manager 改造
|
|
||||||
|
|
||||||
### 4.1 接口扩展
|
|
||||||
|
|
||||||
```go
|
|
||||||
// internal/session/manager.go
|
|
||||||
|
|
||||||
type Manager interface {
|
|
||||||
// ===== 原有方法(签名变更) =====
|
|
||||||
|
|
||||||
// Create 创建新会话,关联 user_id。
|
|
||||||
Create(ctx context.Context, userID string, config models.SessionConfig) (string, error)
|
|
||||||
|
|
||||||
Get(ctx context.Context, sessionID string) (*models.Session, error)
|
|
||||||
UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error
|
|
||||||
GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error)
|
|
||||||
AppendMessage(ctx context.Context, sessionID string, msg models.Message) error
|
|
||||||
SetActiveRequest(ctx context.Context, sessionID string, requestID string) error
|
|
||||||
GetActiveRequestID(ctx context.Context, sessionID string) (string, error)
|
|
||||||
ClearActiveRequest(ctx context.Context, sessionID string) error
|
|
||||||
Touch(ctx context.Context, sessionID string) error
|
|
||||||
Destroy(ctx context.Context, sessionID string) error
|
|
||||||
ActiveCount() int
|
|
||||||
|
|
||||||
// ===== 新增方法 =====
|
|
||||||
|
|
||||||
// ListByUser 获取用户的对话列表(分页)。
|
|
||||||
ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error)
|
|
||||||
|
|
||||||
// UpdateTitle 更新对话标题。
|
|
||||||
UpdateTitle(ctx context.Context, sessionID string, title string) error
|
|
||||||
|
|
||||||
// LoadFromDB 从 PostgreSQL 加载历史消息到热存储(Redis/内存)。
|
|
||||||
// 用户选择历史对话继续交谈时调用。
|
|
||||||
LoadFromDB(ctx context.Context, sessionID string) error
|
|
||||||
}
|
|
||||||
|
|
||||||
// ConversationSummary 对话列表项。
|
|
||||||
type ConversationSummary struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Title string `json:"title"`
|
|
||||||
LastMessage string `json:"last_message"`
|
|
||||||
MessageCount int `json:"message_count"`
|
|
||||||
UpdatedAt time.Time `json:"updated_at"`
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 4.2 Model 变更
|
|
||||||
|
|
||||||
```go
|
|
||||||
// internal/models/models.go
|
|
||||||
|
|
||||||
type Session struct {
|
|
||||||
ID string `json:"session_id"`
|
|
||||||
UserID string `json:"user_id"` // 新增
|
|
||||||
Title string `json:"title"` // 新增
|
|
||||||
CreatedAt time.Time `json:"created_at"`
|
|
||||||
UpdatedAt time.Time `json:"updated_at"` // 新增
|
|
||||||
Config SessionConfig `json:"config"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type User struct {
|
|
||||||
ID string `json:"id"`
|
|
||||||
Username string `json:"username"`
|
|
||||||
PasswordHash string `json:"-"` // 不序列化到 JSON
|
|
||||||
CreatedAt time.Time `json:"created_at"`
|
|
||||||
UpdatedAt time.Time `json:"updated_at"`
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 4.3 冷热数据策略
|
|
||||||
|
|
||||||
```
|
|
||||||
当前活跃会话: Redis/内存(热) ←→ PostgreSQL(冷,write-through)
|
|
||||||
历史会话加载: PostgreSQL → Redis/内存(按需恢复)
|
|
||||||
```
|
|
||||||
|
|
||||||
**Write-through 保证持久化**:每次 `AppendMessage` 同时写入 PostgreSQL,确保服务重启不丢数据。
|
|
||||||
|
|
||||||
**历史对话恢复流程**:
|
|
||||||
1. 用户从对话列表选择一个历史对话
|
|
||||||
2. 前端带 `conversation_id` 建立 WebSocket 连接
|
|
||||||
3. 后端调用 `sessionManager.LoadFromDB(conversationID)` 将历史消息从 PostgreSQL 加载到 Redis/内存
|
|
||||||
4. 后续对话正常走热存储路径
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 五、WebSocket 认证集成
|
|
||||||
|
|
||||||
### 5.1 连接流程
|
|
||||||
|
|
||||||
```
|
|
||||||
前端 后端
|
|
||||||
| |
|
|
||||||
|-- WS /ws?token=<access> ---->|
|
|
||||||
| &conversation_id=<uuid> |
|
|
||||||
| |-- 校验 access_token
|
|
||||||
| |-- 校验 conversation_id 归属
|
|
||||||
| |-- LoadFromDB(如果是历史对话)
|
|
||||||
| |-- 创建新 session(如果 conversation_id 为空)
|
|
||||||
|<-- connected {session_id} ---|
|
|
||||||
| |
|
|
||||||
|-- query {image, audio} ----->| (正常对话流程)
|
|
||||||
```
|
|
||||||
|
|
||||||
### 5.2 Go 实现
|
|
||||||
|
|
||||||
```go
|
|
||||||
// internal/ws/handler.go
|
|
||||||
|
|
||||||
func (h *Handler) HandleWS(c *gin.Context) {
|
|
||||||
// 1. 提取并校验 access_token
|
|
||||||
tokenStr := c.Query("token")
|
|
||||||
if tokenStr == "" {
|
|
||||||
c.JSON(401, gin.H{"error": "missing token"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
claims, err := h.tokenManager.ValidateAccess(tokenStr)
|
|
||||||
if err != nil {
|
|
||||||
c.JSON(401, gin.H{"error": "invalid token"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2. 提取 conversation_id(可选)
|
|
||||||
conversationID := c.Query("conversation_id")
|
|
||||||
|
|
||||||
// 3. 升级 WebSocket
|
|
||||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// 4. 获取或创建 session
|
|
||||||
var sessionID string
|
|
||||||
if conversationID != "" {
|
|
||||||
// 验证该对话属于当前用户
|
|
||||||
sess, err := h.sessionMgr.Get(c, conversationID)
|
|
||||||
if err != nil || sess.UserID != claims.UserID {
|
|
||||||
conn.WriteJSON(models.WsError{Type: "error", Code: "SESSION_NOT_FOUND"})
|
|
||||||
conn.Close()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
// 加载历史到热存储
|
|
||||||
h.sessionMgr.LoadFromDB(c, conversationID)
|
|
||||||
sessionID = conversationID
|
|
||||||
} else {
|
|
||||||
// 创建新对话
|
|
||||||
sessionID, _ = h.sessionMgr.Create(c, claims.UserID, models.DefaultConfig())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 5. 进入正常 WS 处理循环
|
|
||||||
h.handleSession(conn, sessionID, claims.UserID)
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 5.3 前端连接方式
|
|
||||||
|
|
||||||
```typescript
|
|
||||||
// WebSocket 连接
|
|
||||||
const ws = new WebSocket(
|
|
||||||
`wss://${window.location.host}/ws?token=${accessToken}&conversation_id=${selectedConvId || ''}`
|
|
||||||
);
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 六、前端设计概要
|
|
||||||
|
|
||||||
### 6.1 页面路由
|
|
||||||
|
|
||||||
```
|
|
||||||
/ → 未登录重定向到 /login
|
|
||||||
/login → AuthPage(登录/注册表单)
|
|
||||||
/chat → 主界面(需登录)
|
|
||||||
/chat/:id → 主界面,自动加载指定对话
|
|
||||||
```
|
|
||||||
|
|
||||||
### 6.2 组件结构
|
|
||||||
|
|
||||||
```
|
|
||||||
App
|
|
||||||
├── AuthPage ← 新增:登录/注册
|
|
||||||
└── ChatLayout(需登录)
|
|
||||||
├── ConversationList ← 新增:侧边栏对话列表
|
|
||||||
│ ├── 对话项(标题、最后消息、时间)
|
|
||||||
│ ├── 新建对话按钮
|
|
||||||
│ └── 删除对话按钮
|
|
||||||
├── ChatPanel ← 现有,需适配多对话
|
|
||||||
├── VideoPreview ← 现有
|
|
||||||
├── MicManager ← 现有
|
|
||||||
└── ConfigPanel ← 现有
|
|
||||||
```
|
|
||||||
|
|
||||||
### 6.3 新增 Hook
|
|
||||||
|
|
||||||
```typescript
|
|
||||||
// useAuth — 认证状态管理
|
|
||||||
function useAuth() {
|
|
||||||
const [user, setUser] = useState<User | null>(null);
|
|
||||||
const [loading, setLoading] = useState(true);
|
|
||||||
|
|
||||||
const login = async (username: string, password: string) => { ... };
|
|
||||||
const register = async (username: string, password: string) => { ... };
|
|
||||||
const logout = async () => { ... };
|
|
||||||
const refreshToken = async () => { ... };
|
|
||||||
|
|
||||||
// 请求拦截器:自动附加 Authorization header
|
|
||||||
// 401 时自动尝试 refresh,失败则跳转登录
|
|
||||||
|
|
||||||
return { user, loading, login, register, logout };
|
|
||||||
}
|
|
||||||
|
|
||||||
// useConversations — 对话列表管理
|
|
||||||
function useConversations() {
|
|
||||||
const [conversations, setConversations] = useState<ConversationSummary[]>([]);
|
|
||||||
const [currentId, setCurrentId] = useState<string | null>(null);
|
|
||||||
|
|
||||||
const fetchList = async (page?: number) => { ... };
|
|
||||||
const createNew = async () => { ... };
|
|
||||||
const deleteConv = async (id: string) => { ... };
|
|
||||||
const renameConv = async (id: string, title: string) => { ... };
|
|
||||||
const selectConv = (id: string) => { setCurrentId(id); };
|
|
||||||
|
|
||||||
return { conversations, currentId, fetchList, createNew, deleteConv, renameConv, selectConv };
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 6.4 对话标题自动生成
|
|
||||||
|
|
||||||
```go
|
|
||||||
// 内部逻辑:首条 user 消息的前 20 个字符作为 title
|
|
||||||
func generateTitle(firstMessage string) string {
|
|
||||||
runes := []rune(firstMessage)
|
|
||||||
if len(runes) > 20 {
|
|
||||||
return string(runes[:20]) + "…"
|
|
||||||
}
|
|
||||||
return firstMessage
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
在 `AppendMessage` 时,如果 session 的 title 仍为 "新对话",自动更新为 `generateTitle(msg.Content)`。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 七、配置扩展
|
|
||||||
|
|
||||||
### 7.1 Go 配置结构体
|
|
||||||
|
|
||||||
```go
|
|
||||||
type Config struct {
|
|
||||||
App AppConfig `mapstructure:"app"`
|
|
||||||
Server ServerConfig `mapstructure:"server"`
|
|
||||||
Auth AuthConfig `mapstructure:"auth"` // 新增
|
|
||||||
Redis RedisConfig `mapstructure:"redis"`
|
|
||||||
AI AIConfig `mapstructure:"ai"`
|
|
||||||
Storage StorageConfig `mapstructure:"storage"`
|
|
||||||
Log LogConfig `mapstructure:"log"`
|
|
||||||
}
|
|
||||||
|
|
||||||
type AuthConfig struct {
|
|
||||||
JWTSecret string `mapstructure:"jwt_secret"` // 必须通过环境变量设置
|
|
||||||
AccessTTL int `mapstructure:"access_ttl"` // 分钟,默认 15
|
|
||||||
RefreshTTL int `mapstructure:"refresh_ttl"` // 分钟,默认 10080 (7天)
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 7.2 配置文件示例
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
# config.yaml
|
|
||||||
auth:
|
|
||||||
access_ttl: 15 # 分钟
|
|
||||||
refresh_ttl: 10080 # 7天
|
|
||||||
|
|
||||||
storage:
|
|
||||||
driver: "memory" # "memory" | "postgres"
|
|
||||||
dsn: ""
|
|
||||||
```
|
|
||||||
|
|
||||||
### 7.3 环境变量
|
|
||||||
|
|
||||||
| 配置项 | 环境变量 | 说明 |
|
|
||||||
|--------|---------|------|
|
|
||||||
| `auth.jwt_secret` | `CAMTALK_AUTH_JWT_SECRET` | **必须设置**,JWT 签名密钥 |
|
|
||||||
| `auth.access_ttl` | `CAMTALK_AUTH_ACCESS_TTL` | access_token 有效期(分钟) |
|
|
||||||
| `auth.refresh_ttl` | `CAMTALK_AUTH_REFRESH_TTL` | refresh_token 有效期(分钟) |
|
|
||||||
| `storage.driver` | `CAMTALK_STORAGE_DRIVER` | `"memory"` 或 `"postgres"` |
|
|
||||||
| `storage.dsn` | `CAMTALK_STORAGE_DSN` | PostgreSQL 连接串 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 八、实施阶段
|
|
||||||
|
|
||||||
### Phase 1:用户认证系统
|
|
||||||
|
|
||||||
- [ ] 数据库 schema 迁移脚本(users, refresh_tokens 表)
|
|
||||||
- [ ] `internal/auth/` 包:TokenManager, bcrypt 工具, JWT 中间件
|
|
||||||
- [ ] `internal/store/user.go`:UserRepository 接口 + PostgreSQL 实现
|
|
||||||
- [ ] REST API:`/api/auth/register`, `/api/auth/login`, `/api/auth/refresh`, `/api/auth/logout`
|
|
||||||
- [ ] 单元测试
|
|
||||||
|
|
||||||
### Phase 2:对话 CRUD + 消息持久化
|
|
||||||
|
|
||||||
- [ ] 数据库 schema 迁移脚本(sessions, messages 表改造)
|
|
||||||
- [ ] `internal/store/conversation.go`:ConversationRepository 接口 + PostgreSQL 实现
|
|
||||||
- [ ] Session Manager 扩展:Create 绑定 user_id, ListByUser, UpdateTitle
|
|
||||||
- [ ] REST API:`/api/conversations` CRUD + `/api/conversations/:id/messages`
|
|
||||||
- [ ] Write-through:AppendMessage 同时写 PostgreSQL
|
|
||||||
|
|
||||||
### Phase 3:对话历史恢复
|
|
||||||
|
|
||||||
- [ ] `sessionManager.LoadFromDB()` 实现
|
|
||||||
- [ ] 对话标题自动生成逻辑
|
|
||||||
- [ ] REST API:对话详情、历史消息查询(分页)
|
|
||||||
|
|
||||||
### Phase 4:前端集成
|
|
||||||
|
|
||||||
- [ ] `useAuth` hook + 请求拦截器(自动附加 token、自动 refresh)
|
|
||||||
- [ ] `AuthPage` 组件(登录/注册表单)
|
|
||||||
- [ ] `ConversationList` 组件
|
|
||||||
- [ ] `useConversations` hook
|
|
||||||
- [ ] 路由守卫:未登录重定向到 `/login`
|
|
||||||
- [ ] WebSocket 连接带 token + conversation_id
|
|
||||||
- [ ] `useVisionSession` 适配多对话切换
|
|
||||||
|
|
||||||
### Phase 5:配置与收尾
|
|
||||||
|
|
||||||
- [ ] 配置结构体扩展(AuthConfig)
|
|
||||||
- [ ] config.yaml 更新
|
|
||||||
- [ ] docker-compose 添加 PostgreSQL
|
|
||||||
- [ ] 集成测试
|
|
||||||
- [ ] 更新 `02-系统架构.md` 和 `03-接口文档.md`
|
|
||||||
425
docs/12-自定义情景.md
Normal file
425
docs/12-自定义情景.md
Normal file
@@ -0,0 +1,425 @@
|
|||||||
|
# 自建情景功能
|
||||||
|
|
||||||
|
## 概述
|
||||||
|
|
||||||
|
用户可以创建自己的情景,而不仅限于系统预置的 5 种情景。
|
||||||
|
|
||||||
|
**系统预置情景**(不可修改):
|
||||||
|
- 💬 自由对话
|
||||||
|
- 🎯 模拟面试官
|
||||||
|
- 📚 英语老师
|
||||||
|
- ⚔️ 辩论对手
|
||||||
|
- 🌐 同声翻译
|
||||||
|
|
||||||
|
**用户自建情景**(可增删改):
|
||||||
|
- 🎨 创意写作导师
|
||||||
|
- 🧘 心理咨询师
|
||||||
|
- 👨🍳 私人厨师
|
||||||
|
- 📖 历史学家
|
||||||
|
- ... (用户自由创建)
|
||||||
|
|
||||||
|
**用户旅程**:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. 用户点击"创建情景"按钮
|
||||||
|
↓
|
||||||
|
2. 弹出创建对话框
|
||||||
|
↓
|
||||||
|
3. 填写表单:
|
||||||
|
- 情景名称(必填)
|
||||||
|
- 情景图标(可选)
|
||||||
|
- 简短描述(可选)
|
||||||
|
- 角色 Prompt(必填,最少 10 字)
|
||||||
|
- 首句引导(可选)
|
||||||
|
↓
|
||||||
|
4. 点击"创建"
|
||||||
|
↓
|
||||||
|
5. 情景保存到数据库
|
||||||
|
↓
|
||||||
|
6. 情景出现在选择列表中
|
||||||
|
↓
|
||||||
|
7. 用户切换到自建情景
|
||||||
|
↓
|
||||||
|
8. AI 按照用户设定的 Prompt 扮演角色
|
||||||
|
```
|
||||||
|
|
||||||
|
**核心特性**:完整 CRUD 操作(创建/查看/编辑/删除),通过 `user_id` 实现用户数据完全隔离,Eino Graph 管线深度集成(动态加载自建情景 Prompt),中文/英文/日文全覆盖,Modal 对话框 + 图标选择器 + Prompt 编写指南,创建后立即可用无需刷新。
|
||||||
|
|
||||||
|
## 技术架构
|
||||||
|
|
||||||
|
### 数据流
|
||||||
|
|
||||||
|
**创建情景**:
|
||||||
|
|
||||||
|
```
|
||||||
|
用户填写表单 → POST /api/scenarios → Handler 验证
|
||||||
|
→ Repository.Create → PostgreSQL 插入 → 返回情景对象
|
||||||
|
```
|
||||||
|
|
||||||
|
**AI 对话使用自建情景**:
|
||||||
|
|
||||||
|
```
|
||||||
|
WebSocket 连接 → ServeWS 获取 userID
|
||||||
|
→ Eino Graph 初始化 → nodes_history 查询 user_scenarios
|
||||||
|
→ GetScenarioPrompt(customScenarios) → 构建 System Prompt
|
||||||
|
→ LLM 生成回复
|
||||||
|
```
|
||||||
|
|
||||||
|
### Eino 框架集成
|
||||||
|
|
||||||
|
**数据传递链路**:
|
||||||
|
|
||||||
|
```
|
||||||
|
JWT Token → userID
|
||||||
|
↓
|
||||||
|
Session.UserID
|
||||||
|
↓
|
||||||
|
PipelineInput.UserID
|
||||||
|
↓
|
||||||
|
PipelineState.UserID
|
||||||
|
↓
|
||||||
|
nodes_history.go: scenarioRepo.FindByUserID(userID)
|
||||||
|
↓
|
||||||
|
构建 customScenarios map[string]string
|
||||||
|
↓
|
||||||
|
llm.GetScenarioPrompt(scenarioID, language, customScenarios)
|
||||||
|
↓
|
||||||
|
LLM 使用自建情景 Prompt
|
||||||
|
```
|
||||||
|
|
||||||
|
**关键修改文件**:
|
||||||
|
|
||||||
|
| 文件 | 变更说明 |
|
||||||
|
|------|----------|
|
||||||
|
| `backend/internal/eino/state.go` | PipelineState 添加 `UserID` |
|
||||||
|
| `backend/internal/eino/types.go` | PipelineInput 添加 `UserID` |
|
||||||
|
| `backend/internal/eino/graph.go` | 接受 `scenarioRepo` 参数 |
|
||||||
|
| `backend/internal/eino/adapter.go` | 设置 UserID |
|
||||||
|
| `backend/internal/eino/nodes_history.go` | 查询自建情景 |
|
||||||
|
| `backend/internal/ws/handler.go` | 首句引导支持自建情景 |
|
||||||
|
|
||||||
|
## 数据模型
|
||||||
|
|
||||||
|
### 数据库表结构
|
||||||
|
|
||||||
|
**表名**: `user_scenarios`
|
||||||
|
|
||||||
|
```sql
|
||||||
|
CREATE TABLE user_scenarios (
|
||||||
|
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
name VARCHAR(50) NOT NULL,
|
||||||
|
icon VARCHAR(10) DEFAULT '✨',
|
||||||
|
description VARCHAR(100), -- 可选
|
||||||
|
prompt TEXT NOT NULL,
|
||||||
|
greeting VARCHAR(500), -- 可选
|
||||||
|
language VARCHAR(10) DEFAULT 'zh-CN',
|
||||||
|
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||||
|
updated_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||||
|
|
||||||
|
CONSTRAINT unique_user_scenario UNIQUE(user_id, name),
|
||||||
|
CONSTRAINT check_name_length CHECK (char_length(name) >= 2 AND char_length(name) <= 50),
|
||||||
|
CONSTRAINT check_description_length CHECK (description IS NULL OR char_length(description) <= 100),
|
||||||
|
CONSTRAINT check_prompt_length CHECK (char_length(prompt) >= 10 AND char_length(prompt) <= 2000),
|
||||||
|
CONSTRAINT check_greeting_length CHECK (greeting IS NULL OR char_length(greeting) <= 500)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX idx_user_scenarios_user_id ON user_scenarios(user_id);
|
||||||
|
CREATE INDEX idx_user_scenarios_created_at ON user_scenarios(created_at DESC);
|
||||||
|
```
|
||||||
|
|
||||||
|
**字段说明**:
|
||||||
|
|
||||||
|
| 字段 | 说明 |
|
||||||
|
|------|------|
|
||||||
|
| `id` | 情景唯一标识 |
|
||||||
|
| `user_id` | 所属用户,实现数据隔离 |
|
||||||
|
| `name` | 情景名称(2-50 字符) |
|
||||||
|
| `icon` | Emoji 图标(默认 ✨) |
|
||||||
|
| `description` | 简短描述(可选,最多 100 字符) |
|
||||||
|
| `prompt` | 角色 System Prompt(10-2000 字符) |
|
||||||
|
| `greeting` | 首句引导(可选,最多 500 字符) |
|
||||||
|
| `language` | 默认语言(zh-CN / en-US / ja-JP) |
|
||||||
|
|
||||||
|
### 后端数据模型
|
||||||
|
|
||||||
|
```go
|
||||||
|
// backend/internal/models/user_scenario.go
|
||||||
|
|
||||||
|
type UserScenario struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
UserID string `json:"user_id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Icon string `json:"icon"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Prompt string `json:"prompt"`
|
||||||
|
Greeting string `json:"greeting,omitempty"`
|
||||||
|
Language string `json:"language"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type CreateUserScenarioRequest struct {
|
||||||
|
Name string `json:"name" binding:"required,min=2,max=50"`
|
||||||
|
Icon string `json:"icon,omitempty"`
|
||||||
|
Description string `json:"description,omitempty" binding:"omitempty,max=100"`
|
||||||
|
Prompt string `json:"prompt" binding:"required,min=10,max=2000"`
|
||||||
|
Greeting string `json:"greeting,omitempty" binding:"omitempty,max=500"`
|
||||||
|
Language string `json:"language,omitempty"`
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 前端数据结构
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
// frontend/src/lib/api/scenarios.ts
|
||||||
|
|
||||||
|
export interface UserScenario {
|
||||||
|
id: string;
|
||||||
|
user_id: string;
|
||||||
|
name: string;
|
||||||
|
icon: string;
|
||||||
|
description: string;
|
||||||
|
prompt: string;
|
||||||
|
greeting?: string;
|
||||||
|
language: string;
|
||||||
|
created_at: string;
|
||||||
|
updated_at: string;
|
||||||
|
}
|
||||||
|
|
||||||
|
// frontend/src/hooks/useScenarios.ts
|
||||||
|
|
||||||
|
export interface ExtendedScenario {
|
||||||
|
id: string;
|
||||||
|
icon: string;
|
||||||
|
name: string;
|
||||||
|
nameKey?: string;
|
||||||
|
description?: string;
|
||||||
|
descKey?: string;
|
||||||
|
isCustom: boolean;
|
||||||
|
prompt?: string;
|
||||||
|
greeting?: string;
|
||||||
|
language?: string;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## REST API
|
||||||
|
|
||||||
|
### API 端点
|
||||||
|
|
||||||
|
| 方法 | 路径 | 说明 | 权限 |
|
||||||
|
|------|------|------|------|
|
||||||
|
| GET | `/api/scenarios` | 获取用户的所有自建情景 | 需登录 |
|
||||||
|
| POST | `/api/scenarios` | 创建新情景 | 需登录 |
|
||||||
|
| GET | `/api/scenarios/:id` | 获取单个情景详情 | 需登录 |
|
||||||
|
| PATCH | `/api/scenarios/:id` | 更新情景 | 需登录 |
|
||||||
|
| DELETE | `/api/scenarios/:id` | 删除情景 | 需登录 |
|
||||||
|
|
||||||
|
### API 示例
|
||||||
|
|
||||||
|
**创建情景**:
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /api/scenarios
|
||||||
|
Authorization: Bearer <access_token>
|
||||||
|
Content-Type: application/json
|
||||||
|
|
||||||
|
{
|
||||||
|
"name": "创意写作导师",
|
||||||
|
"icon": "✨",
|
||||||
|
"description": "帮助构思故事情节和写作技巧",
|
||||||
|
"prompt": "你是一位创意写作导师,帮助用户构思故事情节、人物设定和写作技巧...",
|
||||||
|
"greeting": "你好!我是你的创意写作导师。今天想聊聊什么故事创意呢?",
|
||||||
|
"language": "zh-CN"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
响应 201 Created:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "uuid-xxx",
|
||||||
|
"user_id": "uuid-user",
|
||||||
|
"name": "创意写作导师",
|
||||||
|
"icon": "✨"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**获取列表**:
|
||||||
|
|
||||||
|
```http
|
||||||
|
GET /api/scenarios
|
||||||
|
Authorization: Bearer <access_token>
|
||||||
|
```
|
||||||
|
|
||||||
|
响应 200 OK:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"scenarios": [],
|
||||||
|
"total": 3
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 前端实现
|
||||||
|
|
||||||
|
### 组件结构
|
||||||
|
|
||||||
|
```
|
||||||
|
frontend/src/
|
||||||
|
├── components/
|
||||||
|
│ ├── CreateScenarioModal/
|
||||||
|
│ │ └── index.tsx # 创建情景对话框
|
||||||
|
│ ├── EditScenarioModal/
|
||||||
|
│ │ └── index.tsx # 编辑情景对话框
|
||||||
|
│ └── ConfigPanel/
|
||||||
|
│ └── index.tsx # 设置面板(改造)
|
||||||
|
├── hooks/
|
||||||
|
│ └── useScenarios.ts # 情景管理 Hook
|
||||||
|
└── lib/
|
||||||
|
└── api/
|
||||||
|
└── scenarios.ts # API 调用封装
|
||||||
|
```
|
||||||
|
|
||||||
|
### 核心 Hook
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
// useScenarios.ts
|
||||||
|
|
||||||
|
export function useScenarios(token: string | null) {
|
||||||
|
const [allScenarios, setAllScenarios] = useState<ExtendedScenario[]>([]);
|
||||||
|
|
||||||
|
// 合并系统预置 + 用户自建
|
||||||
|
useEffect(() => {
|
||||||
|
const systemScenarios = scenarios.map(s => ({...s, isCustom: false}));
|
||||||
|
const customScenarios = customList.map(s => ({...s, isCustom: true}));
|
||||||
|
setAllScenarios([...systemScenarios, ...customScenarios]);
|
||||||
|
}, [customList]);
|
||||||
|
|
||||||
|
return {
|
||||||
|
allScenarios,
|
||||||
|
createScenario,
|
||||||
|
updateScenario,
|
||||||
|
deleteScenario,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 创建情景表单
|
||||||
|
|
||||||
|
**表单字段**:
|
||||||
|
|
||||||
|
- 名称(必填,2-50 字符)
|
||||||
|
- 图标(可选,24 个预设 emoji)
|
||||||
|
- 描述(可选,最多 100 字符)
|
||||||
|
- Prompt(必填,10-2000 字符)
|
||||||
|
- 首句引导(可选,最多 500 字符)
|
||||||
|
- 语言(可选,默认 zh-CN)
|
||||||
|
|
||||||
|
**表单验证**:
|
||||||
|
|
||||||
|
- 实时字符计数
|
||||||
|
- 长度限制提示
|
||||||
|
- 必填项高亮
|
||||||
|
|
||||||
|
## 使用指南
|
||||||
|
|
||||||
|
### 后端 API 测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. 注册用户
|
||||||
|
curl -X POST http://localhost:8080/api/auth/register \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"username":"testuser","password":"test12345"}'
|
||||||
|
|
||||||
|
# 2. 创建情景
|
||||||
|
TOKEN="<access_token>"
|
||||||
|
curl -X POST http://localhost:8080/api/scenarios \
|
||||||
|
-H "Authorization: Bearer $TOKEN" \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"name": "创意写作导师",
|
||||||
|
"icon": "✨",
|
||||||
|
"prompt": "你是一位创意写作导师...",
|
||||||
|
"language": "zh-CN"
|
||||||
|
}'
|
||||||
|
|
||||||
|
# 3. 获取列表
|
||||||
|
curl -X GET http://localhost:8080/api/scenarios \
|
||||||
|
-H "Authorization: Bearer $TOKEN"
|
||||||
|
|
||||||
|
# 4. 更新情景
|
||||||
|
curl -X PATCH http://localhost:8080/api/scenarios/<id> \
|
||||||
|
-H "Authorization: Bearer $TOKEN" \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"name":"高级写作导师"}'
|
||||||
|
|
||||||
|
# 5. 删除情景
|
||||||
|
curl -X DELETE http://localhost:8080/api/scenarios/<id> \
|
||||||
|
-H "Authorization: Bearer $TOKEN"
|
||||||
|
```
|
||||||
|
|
||||||
|
### 前端功能测试
|
||||||
|
|
||||||
|
1. 刷新浏览器(Cmd+Shift+R)
|
||||||
|
2. 登录账户
|
||||||
|
3. 打开设置面板(右上角齿轮)
|
||||||
|
4. 滚动到"我的情景"区域
|
||||||
|
5. 点击"+ 创建新情景"
|
||||||
|
6. 填写表单并提交
|
||||||
|
7. 验证列表中出现新情景
|
||||||
|
8. 切换到自建情景,验证首句引导
|
||||||
|
9. 发送消息,验证 AI 使用自建 Prompt
|
||||||
|
10. 编辑情景,验证数据预填充
|
||||||
|
11. 删除情景,验证二次确认
|
||||||
|
|
||||||
|
## 安全与限制
|
||||||
|
|
||||||
|
### 用户配额
|
||||||
|
|
||||||
|
```go
|
||||||
|
const MaxScenariosPerUser = 20 // 每个用户最多 20 个自建情景
|
||||||
|
```
|
||||||
|
|
||||||
|
### 权限控制
|
||||||
|
|
||||||
|
- 只能查看/编辑/删除自己的情景
|
||||||
|
- 系统预置情景不可编辑/删除
|
||||||
|
- 后端验证 `user_id` 匹配
|
||||||
|
|
||||||
|
### 数据验证
|
||||||
|
|
||||||
|
**后端**:
|
||||||
|
|
||||||
|
- 名称:2-50 字符
|
||||||
|
- 描述:可选,最多 100 字符
|
||||||
|
- Prompt:10-2000 字符
|
||||||
|
- 首句引导:可选,最多 500 字符
|
||||||
|
|
||||||
|
**前端**:
|
||||||
|
|
||||||
|
- 实时字符计数
|
||||||
|
- 超长提示
|
||||||
|
- 必填项高亮
|
||||||
|
|
||||||
|
## 未来优化方向
|
||||||
|
|
||||||
|
**V1.1**:
|
||||||
|
|
||||||
|
- Prompt 模板库
|
||||||
|
- 实时预览效果
|
||||||
|
- 导入导出功能
|
||||||
|
- 情景搜索和筛选
|
||||||
|
|
||||||
|
**V2.0**:
|
||||||
|
|
||||||
|
- 情景市场
|
||||||
|
- 情景分享链接
|
||||||
|
- AI 辅助优化 Prompt
|
||||||
|
- 协作编辑(团队情景)
|
||||||
|
|
||||||
|
## 参考资料
|
||||||
|
|
||||||
|
- [CLAUDE.md](../CLAUDE.md) — 项目开发指南
|
||||||
|
- [02-接口文档.md](./02-接口文档.md) — WebSocket 和 REST API
|
||||||
|
- [自建情景功能-权限隔离说明.md](./自建情景功能-权限隔离说明.md) — 安全设计
|
||||||
544
docs/13-日志追踪.md
Normal file
544
docs/13-日志追踪.md
Normal file
@@ -0,0 +1,544 @@
|
|||||||
|
# 日志追踪系统
|
||||||
|
|
||||||
|
## 概述
|
||||||
|
|
||||||
|
CamTalk 全链路日志追踪系统,通过统一的 trace ID 机制,将 REST API 和 WebSocket 两大入口的所有日志串联起来,实现分布式环境下的请求链路可观测性。
|
||||||
|
|
||||||
|
**核心目标**:
|
||||||
|
- 统一 trace ID 贯穿 REST/WebSocket 两大入口
|
||||||
|
- 所有日志自动附加 trace_id/request_id/session_id
|
||||||
|
- 保护用户隐私,敏感文本截断或降级
|
||||||
|
- 支持按 trace_id 快速定位完整请求链路
|
||||||
|
|
||||||
|
## Trace ID 作用域
|
||||||
|
|
||||||
|
| 标识 | 作用域 | 生成时机 | 用途 |
|
||||||
|
|-----|--------|---------|------|
|
||||||
|
| `trace_id` | **连接级**(整个 WebSocket 生命周期)<br/>**请求级**(单次 REST 请求) | REST: 中间件生成<br/>WebSocket: 升级时生成 | 关联同一连接/请求的所有日志 |
|
||||||
|
| `session_id` | 会话级(对话上下文存储) | ServeWS 时生成 | 标识会话存储 |
|
||||||
|
| `request_id` | 查询级(单次 WebSocket 查询) | 客户端每次查询传入 | 区分同一连接的不同查询 |
|
||||||
|
|
||||||
|
**WebSocket 场景示例**:用户打开页面建立 WebSocket,发起 3 次对话查询:
|
||||||
|
|
||||||
|
```
|
||||||
|
连接建立 trace_id=01J5AAA session_id=uuid-123
|
||||||
|
├─ 查询1 trace_id=01J5AAA request_id=req-001 (问天气)
|
||||||
|
├─ 查询2 trace_id=01J5AAA request_id=req-002 (问新闻)
|
||||||
|
└─ 查询3 trace_id=01J5AAA request_id=req-003 (问股票)
|
||||||
|
```
|
||||||
|
|
||||||
|
**REST 场景示例**:
|
||||||
|
|
||||||
|
```
|
||||||
|
POST /api/auth/login trace_id=01J5BBB request_id=01J5BBB
|
||||||
|
GET /api/conversations trace_id=01J5CCC request_id=01J5CCC
|
||||||
|
```
|
||||||
|
|
||||||
|
## 核心组件
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph TB
|
||||||
|
subgraph trace包["trace 包"]
|
||||||
|
ID["id.go<br/>ULID 生成器"]
|
||||||
|
CTX["context.go<br/>context key 管理"]
|
||||||
|
LOG["logger.go<br/>context-aware logger"]
|
||||||
|
MW["middleware.go<br/>Gin trace 中间件"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph logger包["logger 包"]
|
||||||
|
GINLOG["middleware.go<br/>Gin 请求日志"]
|
||||||
|
GINREC["GinRecovery<br/>panic 恢复"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph 入口层["入口层"]
|
||||||
|
REST["REST API<br/>trace 中间件注入"]
|
||||||
|
WS["WebSocket<br/>ServeWS 注入"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph 业务层["业务层"]
|
||||||
|
HANDLER["Handler"]
|
||||||
|
ADAPTER["Eino Adapter"]
|
||||||
|
NODES["Eino Nodes"]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph 存储层["存储层"]
|
||||||
|
PG["PostgreSQL<br/>session/user/message/scenario"]
|
||||||
|
REDIS["Redis<br/>session/cache/ratelimit"]
|
||||||
|
end
|
||||||
|
|
||||||
|
ID --> MW
|
||||||
|
CTX --> LOG
|
||||||
|
LOG --> HANDLER
|
||||||
|
LOG --> ADAPTER
|
||||||
|
LOG --> NODES
|
||||||
|
LOG --> PG
|
||||||
|
LOG --> REDIS
|
||||||
|
MW --> REST
|
||||||
|
GINLOG --> REST
|
||||||
|
WS --> LOG
|
||||||
|
```
|
||||||
|
|
||||||
|
### trace/id.go — ULID 生成器
|
||||||
|
|
||||||
|
使用 ULID(Universally Unique Lexicographically Sortable Identifier)作为 trace ID:
|
||||||
|
- 时间排序:前 48 位是毫秒时间戳,天然按时间排序
|
||||||
|
- 唯一性:后 80 位随机数,冲突概率极低
|
||||||
|
- 并发安全:使用 `crypto/rand` + `sync.Pool` 复用 entropy 对象
|
||||||
|
|
||||||
|
```go
|
||||||
|
package trace
|
||||||
|
|
||||||
|
import (
|
||||||
|
cryptorand "crypto/rand"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
"github.com/oklog/ulid/v2"
|
||||||
|
)
|
||||||
|
|
||||||
|
var entropyPool = sync.Pool{
|
||||||
|
New: func() interface{} {
|
||||||
|
return ulid.Monotonic(cryptorand.Reader, 0)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateTraceID 生成并发安全的 ULID trace ID
|
||||||
|
func GenerateTraceID() string {
|
||||||
|
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
|
||||||
|
defer entropyPool.Put(entropy)
|
||||||
|
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### trace/context.go — Context Key 管理
|
||||||
|
|
||||||
|
统一管理所有 trace 相关的 context key:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package trace
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
type traceIDKey struct{}
|
||||||
|
type requestIDKey struct{}
|
||||||
|
type sessionIDKey struct{}
|
||||||
|
|
||||||
|
// WithTraceID 将 trace ID 注入 context
|
||||||
|
func WithTraceID(ctx context.Context, traceID string) context.Context {
|
||||||
|
return context.WithValue(ctx, traceIDKey{}, traceID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetTraceID(ctx context.Context) string {
|
||||||
|
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// 类似定义 WithRequestID/GetRequestID 和 WithSessionID/GetSessionID
|
||||||
|
```
|
||||||
|
|
||||||
|
### trace/logger.go — Context-Aware Logger
|
||||||
|
|
||||||
|
自动从 context 提取 trace 字段并附加到日志:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package trace
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
|
||||||
|
func FromContext(ctx context.Context) *zap.SugaredLogger {
|
||||||
|
log := logger.Log
|
||||||
|
|
||||||
|
if traceID := GetTraceID(ctx); traceID != "" {
|
||||||
|
log = log.With("trace_id", traceID)
|
||||||
|
}
|
||||||
|
if requestID := GetRequestID(ctx); requestID != "" {
|
||||||
|
log = log.With("request_id", requestID)
|
||||||
|
}
|
||||||
|
if sessionID := GetSessionID(ctx); sessionID != "" {
|
||||||
|
log = log.With("session_id", sessionID)
|
||||||
|
}
|
||||||
|
|
||||||
|
return log
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**使用模式对比**:
|
||||||
|
|
||||||
|
```go
|
||||||
|
// Before: 手动传递字段
|
||||||
|
logger.Log.Infow("message", "session", sessionID, "request", requestID)
|
||||||
|
|
||||||
|
// After: 自动附加
|
||||||
|
trace.FromContext(ctx).Infow("message")
|
||||||
|
```
|
||||||
|
|
||||||
|
### trace/middleware.go — Gin Trace 中间件
|
||||||
|
|
||||||
|
为 REST 请求生成 trace ID 并注入 context:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package trace
|
||||||
|
|
||||||
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
|
||||||
|
func TraceMiddleware() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
traceID := GenerateTraceID()
|
||||||
|
ctx := WithTraceID(c.Request.Context(), traceID)
|
||||||
|
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
|
||||||
|
|
||||||
|
c.Request = c.Request.WithContext(ctx)
|
||||||
|
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
|
||||||
|
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### logger/middleware.go — 请求日志与 Panic 恢复
|
||||||
|
|
||||||
|
记录所有 HTTP 请求的 method/path/status/latency:
|
||||||
|
|
||||||
|
```go
|
||||||
|
package logger
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/hhs/camtalk/internal/trace"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GinLogger 记录每个 HTTP 请求的基础信息
|
||||||
|
func GinLogger() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
start := time.Now()
|
||||||
|
path := c.Request.URL.Path
|
||||||
|
|
||||||
|
c.Next()
|
||||||
|
|
||||||
|
latency := time.Since(start).Milliseconds()
|
||||||
|
status := c.Writer.Status()
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case status >= 500:
|
||||||
|
log.Errorw("request completed", "method", c.Request.Method,
|
||||||
|
"path", path, "status", status, "latency_ms", latency)
|
||||||
|
case status >= 400:
|
||||||
|
log.Warnw("request completed", "method", c.Request.Method,
|
||||||
|
"path", path, "status", status, "latency_ms", latency)
|
||||||
|
default:
|
||||||
|
log.Infow("request completed", "method", c.Request.Method,
|
||||||
|
"path", path, "status", status, "latency_ms", latency)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GinRecovery 自定义 panic 恢复中间件
|
||||||
|
func GinRecovery() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
defer func() {
|
||||||
|
if err := recover(); err != nil {
|
||||||
|
log := trace.FromContext(c.Request.Context())
|
||||||
|
log.Errorw("panic recovered", "error", err,
|
||||||
|
"path", c.Request.URL.Path, "method", c.Request.Method)
|
||||||
|
c.AbortWithStatus(500)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 中间件注册顺序
|
||||||
|
|
||||||
|
在 `cmd/server/main.go` 中,三层中间件按顺序注册:
|
||||||
|
|
||||||
|
```go
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
|
||||||
|
r.Use(logger.GinLogger()) // 第二层:记录请求
|
||||||
|
r.Use(logger.GinRecovery()) // 第三层:panic 恢复
|
||||||
|
```
|
||||||
|
|
||||||
|
## 日志输出示例
|
||||||
|
|
||||||
|
### REST 请求
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"level": "info",
|
||||||
|
"ts": 1718956800.123,
|
||||||
|
"msg": "login success",
|
||||||
|
"trace_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
|
||||||
|
"request_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
|
||||||
|
"username": "test_user"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### WebSocket 查询链路(含存储层)
|
||||||
|
|
||||||
|
```json
|
||||||
|
// 1. 查询接收
|
||||||
|
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query received"}
|
||||||
|
|
||||||
|
// 2. 会话加载(Redis)
|
||||||
|
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"redis session retrieved", "session_id":"abc-123"}
|
||||||
|
|
||||||
|
// 3. STT 完成
|
||||||
|
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"stt recognition completed", "text_len":45}
|
||||||
|
|
||||||
|
// 4. LLM 完成
|
||||||
|
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"llm generation completed", "tokens":150}
|
||||||
|
|
||||||
|
// 5. 消息持久化(PostgreSQL)
|
||||||
|
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"user", "tokens_used":45}
|
||||||
|
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"assistant", "tokens_used":150}
|
||||||
|
|
||||||
|
// 6. Pipeline 完成
|
||||||
|
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query processing completed", "latency_ms":2340}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 限流触发场景
|
||||||
|
|
||||||
|
```json
|
||||||
|
{"level":"warn", "trace_id":"01J5YYY", "msg":"rate limit triggered", "key":"ratelimit:user-456:query", "retry_after_sec":2.5}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 日志查询操作
|
||||||
|
|
||||||
|
### 按 trace_id 查询完整链路
|
||||||
|
|
||||||
|
**本地开发(文件日志)**:
|
||||||
|
```bash
|
||||||
|
# 查看完整链路
|
||||||
|
grep 'trace_id":"01J5XXX"' backend.log | jq .
|
||||||
|
|
||||||
|
# 查看链路时间线
|
||||||
|
grep 'trace_id":"01J5XXX"' backend.log | jq -r '[.ts, .msg] | @tsv'
|
||||||
|
```
|
||||||
|
|
||||||
|
**Grafana Loki**:
|
||||||
|
```logql
|
||||||
|
{app="camtalk-backend"}
|
||||||
|
|= "trace_id=01J5XXX"
|
||||||
|
| json
|
||||||
|
| line_format "{{.ts}} [{{.level}}] {{.msg}}"
|
||||||
|
```
|
||||||
|
|
||||||
|
### 查询慢请求(延迟 > 5s)
|
||||||
|
|
||||||
|
```logql
|
||||||
|
{app="camtalk-backend"}
|
||||||
|
| json
|
||||||
|
| msg="query processing completed"
|
||||||
|
| latency_ms > 5000
|
||||||
|
```
|
||||||
|
|
||||||
|
### 查询数据库错误
|
||||||
|
|
||||||
|
```logql
|
||||||
|
{app="camtalk-backend"}
|
||||||
|
| json
|
||||||
|
| level="error"
|
||||||
|
| msg=~".*failed"
|
||||||
|
| line_format "{{.trace_id}} {{.msg}} {{.error}}"
|
||||||
|
```
|
||||||
|
|
||||||
|
### 查询 Redis 降级事件
|
||||||
|
|
||||||
|
```logql
|
||||||
|
{app="camtalk-backend"}
|
||||||
|
| json
|
||||||
|
| level="warn"
|
||||||
|
| msg=~"redis.*failed"
|
||||||
|
```
|
||||||
|
|
||||||
|
### 查询错误率
|
||||||
|
|
||||||
|
```logql
|
||||||
|
sum(count_over_time({app="camtalk-backend"} | json | level="error" [5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
## 敏感内容处理规范
|
||||||
|
|
||||||
|
### 完全禁止记录
|
||||||
|
|
||||||
|
- 用户明文密码
|
||||||
|
- JWT token 完整内容(仅记录 "token_present: true")
|
||||||
|
- API Key 完整值(仅记录前 8 字符 + "...")
|
||||||
|
|
||||||
|
### 截断后记录(最多 50 字符)
|
||||||
|
|
||||||
|
- 用户输入文本 → `text_preview`
|
||||||
|
- LLM 生成文本 → `text_preview`
|
||||||
|
- STT 识别文本 → `text_preview`
|
||||||
|
|
||||||
|
**示例**:
|
||||||
|
```go
|
||||||
|
log.Debugw("stt recognition completed",
|
||||||
|
"text_len", len(text),
|
||||||
|
"text_preview", util.Truncate(text, 50))
|
||||||
|
```
|
||||||
|
|
||||||
|
### 仅记录长度/大小
|
||||||
|
|
||||||
|
- 图片数据 → `image_size_bytes`
|
||||||
|
- 音频数据 → `audio_size_bytes`
|
||||||
|
|
||||||
|
### 降级为 Debug 级别
|
||||||
|
|
||||||
|
所有包含用户文本预览的日志,生产环境默认不输出。
|
||||||
|
|
||||||
|
## 日志级别使用准则
|
||||||
|
|
||||||
|
| 场景 | 级别 | 示例 |
|
||||||
|
|-----|------|-----|
|
||||||
|
| 请求生命周期里程碑 | Info | `"query received"`, `"pipeline completed"` |
|
||||||
|
| 中间步骤详情 | Debug | `"stt recognition completed"`, `"history assembled"` |
|
||||||
|
| 敏感内容相关 | Debug | 所有包含用户文本的日志 |
|
||||||
|
| 预期内的失败 | Warn | `"login failed"`, `"rate limited"` |
|
||||||
|
| 系统错误 | Error | `"database query failed"`, `"tts synthesis failed"` |
|
||||||
|
| 严重故障 | Error + stack | `"panic recovered"` |
|
||||||
|
|
||||||
|
## 存储层日志实现
|
||||||
|
|
||||||
|
### PostgreSQL Repository 层
|
||||||
|
|
||||||
|
所有数据库操作统一使用 `trace.FromContext(ctx)` 记录日志:
|
||||||
|
|
||||||
|
**已实现文件**:
|
||||||
|
- `backend/internal/store/session_pg.go` — 会话 CRUD
|
||||||
|
- `backend/internal/store/user_pg.go` — 用户与 refresh token 操作
|
||||||
|
- `backend/internal/store/message_pg.go` — 对话消息存储
|
||||||
|
- `backend/internal/store/user_scenario_repository.go` — 用户自定义情景
|
||||||
|
|
||||||
|
**日志策略**:
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
_, err := r.pool.Exec(ctx, ...)
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("save session failed", "session_id", s.ID, "error", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("session saved", "session_id", s.ID, "user_id", s.UserID)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**NotFound 处理**:预期内的空结果不记录错误:
|
||||||
|
|
||||||
|
```go
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
return nil, ErrSessionNotFound // 不记录日志
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("find session failed", "session_id", id, "error", err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Redis 服务层
|
||||||
|
|
||||||
|
**已实现文件**:
|
||||||
|
- `backend/internal/session/redis.go` — RedisManager(会话存储)
|
||||||
|
- `backend/internal/store/cached_user.go` — CachedUserRepository(用户缓存装饰器)
|
||||||
|
- `backend/internal/ratelimit/redis_bucket.go` — RedisLimiter(令牌桶限流器)
|
||||||
|
|
||||||
|
**会话存储日志**(`redis.go`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
|
||||||
|
return nil, fmt.Errorf("redis get session: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(vals) == 0 {
|
||||||
|
return nil, ErrSessionNotFound // 不记录日志
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debugw("redis session retrieved", "session_id", sessionID)
|
||||||
|
return session, nil
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**缓存降级日志**(`cached_user.go`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
if _, err := pipe.Exec(ctx); err != nil {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
log.Warnw("redis cache write failed for refresh token", "error", err)
|
||||||
|
// 降级:DB 已写入成功,Redis 失败不影响正确性
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**限流触发日志**(`redis_bucket.go`):
|
||||||
|
|
||||||
|
```go
|
||||||
|
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||||
|
log := trace.FromContext(ctx)
|
||||||
|
|
||||||
|
result, err := l.script.Run(ctx, ...).Result()
|
||||||
|
if err != nil {
|
||||||
|
log.Errorw("rate limit check failed", "key", key, "error", err)
|
||||||
|
return true, 0 // fail-open 策略
|
||||||
|
}
|
||||||
|
|
||||||
|
if allowed == 0 {
|
||||||
|
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
|
||||||
|
return false, retryAfter
|
||||||
|
}
|
||||||
|
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**级别选择原则**:
|
||||||
|
- **Error**:Redis 连接失败、Lua 脚本执行失败(影响功能)
|
||||||
|
- **Warn**:缓存写入失败(可降级)、限流触发(预期内异常)
|
||||||
|
- **Debug**:正常操作完成(避免 Info 级别噪音)
|
||||||
|
|
||||||
|
## 编码规范
|
||||||
|
|
||||||
|
1. **日志语言**:统一使用英文
|
||||||
|
2. **结构化**:始终使用 `Infow`/`Errorw`/`Warnw`/`Debugw`
|
||||||
|
3. **Context 传递**:使用 `trace.FromContext(ctx)` 而非直接引用 `logger.Log`
|
||||||
|
4. **敏感内容**:禁止在 Info 及以上级别记录用户文本原文
|
||||||
|
5. **错误日志**:采用"调用方记录"原则,底层函数 return wrapped error
|
||||||
|
6. **级别约定**:
|
||||||
|
- `Debug`:内部状态跟踪、开发调试信息(数据库/缓存成功操作)
|
||||||
|
- `Info`:请求/连接生命周期、关键操作里程碑
|
||||||
|
- `Warn`:可降级异常(Redis 故障、限流触发)
|
||||||
|
- `Error`:影响用户的操作失败(数据库错误、Redis 连接失败)
|
||||||
|
- `Fatal`:仅启动阶段不可恢复错误
|
||||||
|
7. **预期内的空结果**:`pgx.ErrNoRows`、`redis.Nil` 等不记录错误日志
|
||||||
|
|
||||||
|
## 性能考量
|
||||||
|
|
||||||
|
### FromContext 开销
|
||||||
|
|
||||||
|
- 有 trace_id:~200-300 ns/op
|
||||||
|
- 无 trace_id:~10-20 ns/op(仅返回全局 logger)
|
||||||
|
- 1000 QPS 场景额外开销约 0.2ms,可接受
|
||||||
|
|
||||||
|
### ULID 生成吞吐量
|
||||||
|
|
||||||
|
- 单线程:~500k ops/s
|
||||||
|
- 并发 8 线程:~2M ops/s
|
||||||
|
|
||||||
|
**验收标准**:1000 QPS 下,trace 系统开销 < 1% CPU,< 0.5ms P99 延迟。
|
||||||
@@ -1,221 +0,0 @@
|
|||||||
# CamTalk 后端完善计划
|
|
||||||
|
|
||||||
## Context
|
|
||||||
|
|
||||||
后端当前是一个骨架:`main.go` 启动 Gin 服务器,`ws/handler.go` 实现了 WebSocket 连接生命周期和消息分发,`models/models.go` 定义了所有协议消息类型,`config/config.go` 实现了 Viper 配置加载。但所有业务逻辑都是 TODO 桩——没有 Session Manager、没有 AI 服务客户端、没有编排层、没有日志/错误工具、没有测试。前端已基本完成,正在等待后端提供真实的 AI 管道。
|
|
||||||
|
|
||||||
**目标**:按设计文档(`docs/03-接口文档.md` 为最高依据)逐步填充所有业务模块,使端到端的 STT → LLM → TTS 流式管道可用。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 分阶段实施
|
|
||||||
|
|
||||||
### Phase 1:基础设施(logger、errors、config 接入、graceful shutdown)
|
|
||||||
|
|
||||||
**目标**:为后续模块提供日志、错误码、配置等基础能力,替换 `main.go` 中的硬编码值。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 1.1 | 实现 Zap 日志封装 | `internal/logger/logger.go` | 提供 `Init(level, format)` 和全局 `*zap.SugaredLogger`,替换所有 `log.Printf` |
|
|
||||||
| 1.2 | 实现错误码常量 + WS 错误发送工具 | `internal/errors/codes.go` | 10 个错误码常量 + `SendWSError(client, code, requestID, err)` |
|
|
||||||
| 1.3 | main.go 接入 config.Load() | `cmd/server/main.go` | 用 `cfg.Server.Host:Port` 替换硬编码 `:8080`,初始化 logger |
|
|
||||||
| 1.4 | 添加 graceful shutdown | `cmd/server/main.go` | `signal.NotifyContext` + `http.Server.Shutdown`,10s drain |
|
|
||||||
| 1.5 | 添加 .gitignore | `backend/.gitignore` | 排除 `server` 二进制、`.env`、`tmp/` |
|
|
||||||
|
|
||||||
> **CORS**:不在此处实现,生产环境由 Nginx 反向代理统一处理跨域。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 2:Session Manager
|
|
||||||
|
|
||||||
**目标**:实现会话生命周期管理,让 WS handler 能追踪会话、存储对话历史。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 2.1 | 定义 SessionManager 接口 | `internal/session/manager.go` | 方法:`Create`, `Get`, `UpdateConfig`, `GetHistory`, `AppendMessage`, `SetActiveRequest`, `ClearActiveRequest`, `Touch`, `Destroy` |
|
|
||||||
| 2.2 | 实现内存版 SessionManager | `internal/session/memory.go` | `sync.RWMutex` + `map[string]*sessionEntry`,TTL 30 分钟,历史上限 20 条 |
|
|
||||||
| 2.3 | 实现 Redis 版 SessionManager | `internal/session/redis.go` | `session:{id}:meta` Hash + `session:{id}:history` List,TTL 刷新,选配 |
|
|
||||||
| 2.4 | 编写 Session Manager 测试 | `internal/session/memory_test.go` | 覆盖 Create/Get/Expire/Destroy/AppendMessage/History 上限 |
|
|
||||||
| 2.5 | WS handler 接入 SessionManager | `internal/ws/handler.go` | `ServeWS` 接收 `session.Manager` 参数;`connected` 消息后创建会话;`query` 时 Touch + SetActiveRequest;`config` 时 UpdateConfig;断开时不销毁(自然过期) |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 3:AI 服务层接口 + 实现
|
|
||||||
|
|
||||||
**目标**:定义并实现三个 AI 服务客户端,每个服务一个独立包。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| **3a. STT** | | | |
|
|
||||||
| 3.1 | STT 接口定义 | `internal/ai/stt/stt.go` | `Service` 接口:`Recognize(ctx, audio []byte, opts Options) (string, error)`。`Options`: Encoding, SampleRate, Language |
|
|
||||||
| 3.2 | Deepgram 实现 | `internal/ai/stt/deepgram.go` | WebSocket 连接 `wss://api.deepgram.com/v1/listen`,发送 PCM 音频,接收转录结果,5s 超时 |
|
|
||||||
| 3.3 | STT 测试(mock) | `internal/ai/stt/deepgram_test.go` | httptest/WebSocket mock,验证连接、发送、超时 |
|
|
||||||
| **3b. LLM** | | | |
|
|
||||||
| 3.4 | LLM 接口定义 | `internal/ai/llm/llm.go` | `Service` 接口:`ChatStream(ctx, req Request) (<-chan Chunk, error)`。`Request`: Image, Text, History, Language。`Chunk`: Delta, Done, TokensUsed, Model |
|
|
||||||
| 3.5 | OpenAI 实现 | `internal/ai/llm/openai.go` | `POST /v1/chat/completions` + `stream: true`,SSE 解析,10s 超时,image 以 `data:image/jpeg;base64,...` 传入 |
|
|
||||||
| 3.6 | System Prompt 定义 | `internal/ai/llm/prompt.go` | 中文视觉助手提示词,根据 Language/DetailLevel 动态构建 |
|
|
||||||
| 3.7 | LLM 测试(mock) | `internal/ai/llm/openai_test.go` | httptest mock SSE 流,验证流式解析、超时、错误处理 |
|
|
||||||
| **3c. TTS** | | | |
|
|
||||||
| 3.8 | TTS 接口定义 | `internal/ai/tts/tts.go` | `Service` 接口:`SynthesizeStream(ctx, textStream <-chan string, opts Options) (<-chan Chunk, error)`。`Chunk`: Audio []byte, IsLast |
|
|
||||||
| 3.9 | OpenAI 实现 | `internal/ai/tts/openai.go` | `POST /v1/audio/speech` 模型 `tts-1`,逐句发送,返回 MP3 流,5s/句超时 |
|
|
||||||
| 3.10 | TTS 测试(mock) | `internal/ai/tts/openai_test.go` | httptest mock,验证逐句合成、超时 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 4:AI Orchestrator(核心编排)
|
|
||||||
|
|
||||||
**目标**:实现 STT → LLM → TTS 流式并行管道,这是后端最关键的业务逻辑。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 4.1 | Orchestrator 接口 | `internal/orchestrator/orchestrator.go` | `ProcessQuery(ctx, sessionID, req, history, sender)` — 接收查询并执行管道 |
|
|
||||||
| 4.2 | Sender 接口 | `internal/orchestrator/sender.go` | 抽象 WS 推送:`SendSTTResult`, `SendLLMChunk`, `SendLLMDone`, `SendTTSAudio`, `SendError`,便于测试 |
|
|
||||||
| 4.3 | 管道实现 | `internal/orchestrator/pipeline.go` | ① `stt.Recognize()` → 发送 `stt_result` ② `llm.ChatStream()` 并行消费 token → 发送 `llm_chunk` + 句子切分 → channel ③ `tts.SynthesizeStream()` 从 channel 读取 → 发送 `tts_audio` ④ 流结束 → 发送 `llm_done` |
|
|
||||||
| 4.4 | 句子切分器 | `internal/orchestrator/splitter.go` | 按 `。!?\n.!?` 切分,buffer size 4 channel |
|
|
||||||
| 4.5 | 错误降级 | 同上文件 | STT 失败→STT_ERROR+abort;LLM 超时→LLM_TIMEOUT;TTS 失败→静默跳过 |
|
|
||||||
| 4.6 | Interrupt 支持 | 同上文件 | context cancel 触发所有流中止 |
|
|
||||||
| 4.7 | Orchestrator 测试 | `internal/orchestrator/pipeline_test.go` | mock 三个 AI service + mock sender,验证完整流程、中断、错误降级 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 5:WS Handler 完整接入
|
|
||||||
|
|
||||||
**目标**:将 Session Manager + Orchestrator 串入 WebSocket handler,实现端到端消息处理。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 5.1 | Client 扩展 | `internal/ws/handler.go` | 添加 `session.Manager`、`orchestrator.Orchestrator`、`context.CancelFunc`(用于 interrupt) |
|
|
||||||
| 5.2 | query 处理 | 同上 | 解码 audio Base64 → `stt.Recognize` 的输入;Touch 会话;设置 active request;启动 `orchestrator.ProcessQuery` goroutine |
|
|
||||||
| 5.3 | config 处理 | 同上 | 调用 `session.UpdateConfig()` |
|
|
||||||
| 5.4 | interrupt 处理 | 同上 | 查找 active request 的 cancel func,调用 `cancel()`,ClearActiveRequest |
|
|
||||||
| 5.5 | Disconnect 处理 | 同上 | 取消当前活跃请求(如有),不销毁会话 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 6:REST API 补全
|
|
||||||
|
|
||||||
**目标**:补全设计文档中的 REST 端点。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 6.1 | Session 路由 | `internal/api/session.go` | `POST /api/sessions` 创建会话,`DELETE /api/sessions/:id` 销毁会话 |
|
|
||||||
| 6.2 | Health 更新 | `cmd/server/main.go` | 从 SessionManager 获取 `active_sessions` 真实值 |
|
|
||||||
| 6.3 | 路由注册 | `cmd/server/main.go` | 统一注册 REST + WS 路由,注入依赖 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 7:Rate Limiter + Model Router(可选/MVP 后)
|
|
||||||
|
|
||||||
**目标**:防止滥用 + 智能模型选择,MVP 可简化或跳过。
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 7.1 | 令牌桶 Rate Limiter | `internal/middleware/ratelimit.go` | `golang.org/x/time/rate` 或自实现,按 session ID 限流 |
|
|
||||||
| 7.2 | Rate Limiter 中间件 | `internal/middleware/ratelimit.go` | 在 WS query 路径上检查,超限返回 `RATE_LIMITED` |
|
|
||||||
| 7.3 | Model Router | `internal/ai/router.go` | 规则引擎:简单识别→GPT-4o-mini,深度分析→GPT-4o,暂不实现 o1 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Phase 8:集成测试 + 文档同步
|
|
||||||
|
|
||||||
| # | 任务 | 文件 | 说明 |
|
|
||||||
|---|------|------|------|
|
|
||||||
| 8.1 | WS 集成测试 | `internal/ws/handler_test.go` | 启动 Gin test server + gorilla websocket client,验证完整 query→stt_result→llm_chunk→llm_done→tts_audio 流程 |
|
|
||||||
| 8.2 | 文档同步 | `docs/03-接口文档.md` | 代码实现与文档有偏差时更新文档 |
|
|
||||||
| 8.3 | go.sum 清理 | `backend/` | `go mod tidy` 清理无用依赖 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 关键文件清单
|
|
||||||
|
|
||||||
```
|
|
||||||
backend/
|
|
||||||
cmd/server/main.go ← Phase 1.3, 1.4, 1.5, 6.2, 6.3
|
|
||||||
internal/
|
|
||||||
config/config.go ← 已完成,Phase 1.3 接入
|
|
||||||
logger/logger.go ← Phase 1.1(新建)
|
|
||||||
errors/codes.go ← Phase 1.2(新建)
|
|
||||||
models/models.go ← 已完成,可能小幅扩展
|
|
||||||
session/
|
|
||||||
manager.go ← Phase 2.1(新建)
|
|
||||||
memory.go ← Phase 2.2(新建)
|
|
||||||
redis.go ← Phase 2.3(新建)
|
|
||||||
memory_test.go ← Phase 2.4(新建)
|
|
||||||
ai/
|
|
||||||
stt/
|
|
||||||
stt.go ← Phase 3.1(新建)
|
|
||||||
deepgram.go ← Phase 3.2(新建)
|
|
||||||
deepgram_test.go ← Phase 3.3(新建)
|
|
||||||
llm/
|
|
||||||
llm.go ← Phase 3.4(新建)
|
|
||||||
openai.go ← Phase 3.5(新建)
|
|
||||||
prompt.go ← Phase 3.6(新建)
|
|
||||||
openai_test.go ← Phase 3.7(新建)
|
|
||||||
tts/
|
|
||||||
tts.go ← Phase 3.8(新建)
|
|
||||||
openai.go ← Phase 3.9(新建)
|
|
||||||
openai_test.go ← Phase 3.10(新建)
|
|
||||||
router.go ← Phase 7.3(新建)
|
|
||||||
orchestrator/
|
|
||||||
orchestrator.go ← Phase 4.1(新建)
|
|
||||||
sender.go ← Phase 4.2(新建)
|
|
||||||
pipeline.go ← Phase 4.3, 4.4, 4.5, 4.6(新建)
|
|
||||||
pipeline_test.go ← Phase 4.7(新建)
|
|
||||||
api/
|
|
||||||
session.go ← Phase 6.1(新建)
|
|
||||||
middleware/
|
|
||||||
ratelimit.go ← Phase 7.1, 7.2(新建)
|
|
||||||
ws/
|
|
||||||
handler.go ← Phase 5.1-5.5(修改)
|
|
||||||
handler_test.go ← Phase 8.1(新建)
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 新增依赖
|
|
||||||
|
|
||||||
| 包 | 用途 | Phase |
|
|
||||||
|----|------|-------|
|
|
||||||
| `go.uber.org/zap` | 结构化日志 | 1 |
|
|
||||||
| `github.com/redis/go-redis/v9` | Redis 客户端 | 2.3 |
|
|
||||||
| `github.com/gorilla/websocket` | 已有,Deepgram WS 也复用 | 3.2 |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 执行顺序与依赖关系
|
|
||||||
|
|
||||||
```
|
|
||||||
Phase 1 (基础设施)
|
|
||||||
↓
|
|
||||||
Phase 2 (Session Manager)
|
|
||||||
↓
|
|
||||||
Phase 3 (AI 服务层) ← 可与 Phase 2 并行开发
|
|
||||||
↓
|
|
||||||
Phase 4 (Orchestrator) ← 依赖 Phase 2 + 3
|
|
||||||
↓
|
|
||||||
Phase 5 (WS Handler 接入) ← 依赖 Phase 4
|
|
||||||
↓
|
|
||||||
Phase 6 (REST API) ← 依赖 Phase 2
|
|
||||||
↓
|
|
||||||
Phase 7 (Rate Limiter + Router) ← 独立,可推后
|
|
||||||
↓
|
|
||||||
Phase 8 (集成测试 + 文档)
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 验证方案
|
|
||||||
|
|
||||||
1. **单元测试**:每个模块独立测试,mock 外部依赖(AI API、Redis)
|
|
||||||
2. **集成测试**:`httptest` 启动 Gin server,用 gorilla/websocket 客户端模拟完整 query 流程
|
|
||||||
3. **端到端手动测试**:启动后端 → 打开前端 → 摄像头+麦克风对话 → 验证 stt_result / llm_chunk / tts_audio 消息流
|
|
||||||
4. **go vet + go test ./...** 通过
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## 设计文档参考
|
|
||||||
|
|
||||||
- 接口规范(最高优先级):`docs/03-接口文档.md`
|
|
||||||
- 系统架构:`docs/02-系统架构.md`
|
|
||||||
- 技术选型:`docs/04-技术选型.md`
|
|
||||||
- 成本控制:`docs/08-成本控制.md`
|
|
||||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user