Compare commits
261 Commits
44f4f3ce1a
...
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 | |||
| 943f36e9ec | |||
| 083ede8b6d | |||
| 14242eb896 | |||
| 6db5bef519 | |||
| 1ddd75d34f | |||
| 7ea6c0d4f2 | |||
| 97357013bb | |||
| 057a226610 | |||
| b60e513317 | |||
| 638a387550 | |||
| 9dc9025d10 | |||
| 92e47718bc | |||
| 6ce36e3ed8 | |||
| 6032608e49 | |||
| 668da26e23 | |||
| f22243a2d2 | |||
| eba0607f94 | |||
| 1f19435951 | |||
| 74cd6c7d2d | |||
| 6b99336941 | |||
| 12eec6fc0e | |||
| 0ba2f9807c | |||
| 91fdc007f7 | |||
| 443f613fbd | |||
| 35372438a5 | |||
| 97484735cf | |||
| 2f05b5fa2b | |||
| 213aa67e9f | |||
| 5eeb592bf0 | |||
| 46cf35967e | |||
| d5ee6deeb9 | |||
| 31dfabd882 | |||
| bab67a8749 | |||
| 5d70bc8862 | |||
| c0279e94c7 | |||
| c68657bf30 | |||
| 5e89ab01bd | |||
| 60fc3ac5d1 | |||
| 0a0eaf7eb9 |
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,16 +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(冷数据持久化层)
|
||||||
pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN)
|
dsn := cfg.Storage.Persistence.DSN
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
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)
|
||||||
}
|
}
|
||||||
@@ -70,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",
|
||||||
@@ -112,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":
|
||||||
@@ -125,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(
|
||||||
@@ -136,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")
|
||||||
@@ -156,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=
|
|
||||||
|
|||||||
@@ -15,10 +15,11 @@ type Service interface {
|
|||||||
|
|
||||||
// Request 推理请求。
|
// Request 推理请求。
|
||||||
type Request struct {
|
type Request struct {
|
||||||
Image []byte // JPEG 图片(已从 Base64 解码)
|
Image []byte // JPEG 图片(已从 Base64 解码)
|
||||||
Text string // 用户语音识别后的文本
|
Text string // 用户语音识别后的文本
|
||||||
History []models.Message // 最近 N 轮对话历史
|
History []models.Message // 最近 N 轮对话历史
|
||||||
Language string // 语言,如 "zh-CN"
|
Language string // 语言,如 "zh-CN"
|
||||||
|
SystemPrompt string // 情景自定义 system prompt(非空时覆盖默认 prompt)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Chunk 流式推理的一个增量片段。
|
// Chunk 流式推理的一个增量片段。
|
||||||
|
|||||||
@@ -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, "")}},
|
|
||||||
})
|
|
||||||
|
|
||||||
// 历史消息
|
|
||||||
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,10 +2,26 @@ package llm
|
|||||||
|
|
||||||
import "strings"
|
import "strings"
|
||||||
|
|
||||||
// BuildSystemPrompt 根据语言和细节级别构建系统提示词。
|
// BuildSystemPrompt 根据语言、细节级别和情景 prompt 构建系统提示词。
|
||||||
func BuildSystemPrompt(language, detailLevel string) string {
|
// scenarioPrompt 非空时,覆盖默认视觉助手 prompt。
|
||||||
|
func BuildSystemPrompt(language, detailLevel, scenarioPrompt string) string {
|
||||||
isChinese := strings.HasPrefix(language, "zh")
|
isChinese := strings.HasPrefix(language, "zh")
|
||||||
|
|
||||||
|
// 情景模式:使用自定义 prompt 作为基础
|
||||||
|
if scenarioPrompt != "" {
|
||||||
|
var prompt strings.Builder
|
||||||
|
prompt.WriteString(scenarioPrompt)
|
||||||
|
if detailLevel == "high" {
|
||||||
|
if isChinese {
|
||||||
|
prompt.WriteString(" 请在涉及视觉内容时提供更详细的描述,包括颜色、位置、数量等细节。")
|
||||||
|
} else {
|
||||||
|
prompt.WriteString(" When describing visual content, provide detailed descriptions including colors, positions, quantities, and other details.")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return prompt.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// 默认模式:视觉助手
|
||||||
var prompt strings.Builder
|
var prompt strings.Builder
|
||||||
if isChinese {
|
if isChinese {
|
||||||
prompt.WriteString("你是一个视觉助手。用户通过摄像头看到一个场景,并用语音向你提问。请用简洁自然的中文回答。如果涉及视觉描述,先说\"我看到……\"。回答控制在3-5句话以内,除非用户要求详细说明。")
|
prompt.WriteString("你是一个视觉助手。用户通过摄像头看到一个场景,并用语音向你提问。请用简洁自然的中文回答。如果涉及视觉描述,先说\"我看到……\"。回答控制在3-5句话以内,除非用户要求详细说明。")
|
||||||
|
|||||||
148
backend/internal/ai/llm/scenarios.go
Normal file
148
backend/internal/ai/llm/scenarios.go
Normal file
@@ -0,0 +1,148 @@
|
|||||||
|
package llm
|
||||||
|
|
||||||
|
import "strings"
|
||||||
|
|
||||||
|
// scenarioPrompt 定义单个情景的多语言 system prompt 和首句引导。
|
||||||
|
type scenarioPrompt struct {
|
||||||
|
ZH string
|
||||||
|
EN string
|
||||||
|
JA string
|
||||||
|
GreetingZH string // 首句引导(中文)
|
||||||
|
GreetingEN string // 首句引导(英文)
|
||||||
|
GreetingJA string // 首句引导(日文)
|
||||||
|
}
|
||||||
|
|
||||||
|
// scenarioPrompts 预置情景 → prompt 映射表。
|
||||||
|
// key 为情景 ID(与前端 Scenario.id 对齐)。
|
||||||
|
var scenarioPrompts = map[string]scenarioPrompt{
|
||||||
|
"interviewer": {
|
||||||
|
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.",
|
||||||
|
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": {
|
||||||
|
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.",
|
||||||
|
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": {
|
||||||
|
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.",
|
||||||
|
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": {
|
||||||
|
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.",
|
||||||
|
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。
|
||||||
|
// 支持系统预置情景和用户自建情景。
|
||||||
|
// customScenarios: 用户自建情景映射表(scenarioID → prompt),可为 nil
|
||||||
|
// 返回空字符串表示无此情景(使用默认 prompt)。
|
||||||
|
func GetScenarioPrompt(scenarioID, language string, customScenarios map[string]string) string {
|
||||||
|
if scenarioID == "" || scenarioID == "free_chat" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1. 优先查找系统预置情景
|
||||||
|
if p, ok := scenarioPrompts[scenarioID]; ok {
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(language, "zh"):
|
||||||
|
return p.GreetingZH
|
||||||
|
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,89 +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("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
|
||||||
|
}
|
||||||
@@ -17,6 +17,7 @@ type SessionConfig struct {
|
|||||||
TTSEnabled bool `json:"tts_enabled"`
|
TTSEnabled bool `json:"tts_enabled"`
|
||||||
DetailLevel string `json:"detail_level"` // "low" | "high"
|
DetailLevel string `json:"detail_level"` // "low" | "high"
|
||||||
Language string `json:"language"`
|
Language string `json:"language"`
|
||||||
|
Scenario string `json:"scenario"` // 情景 ID,如 "free_chat"、"interviewer"
|
||||||
}
|
}
|
||||||
|
|
||||||
// DefaultSessionTitle 默认会话标题。
|
// DefaultSessionTitle 默认会话标题。
|
||||||
@@ -24,7 +25,7 @@ const DefaultSessionTitle = "新对话"
|
|||||||
|
|
||||||
// DefaultConfig 默认会话配置。
|
// DefaultConfig 默认会话配置。
|
||||||
func DefaultConfig() SessionConfig {
|
func DefaultConfig() SessionConfig {
|
||||||
return SessionConfig{TTSEnabled: true, DetailLevel: "low", Language: "zh-CN"}
|
return SessionConfig{TTSEnabled: true, DetailLevel: "low", Language: "zh-CN", Scenario: "free_chat"}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SessionConfigPatch 会话配置增量更新(指针字段表示"未传则不更新")。
|
// SessionConfigPatch 会话配置增量更新(指针字段表示"未传则不更新")。
|
||||||
@@ -32,6 +33,7 @@ type SessionConfigPatch struct {
|
|||||||
TTSEnabled *bool `json:"tts_enabled,omitempty"`
|
TTSEnabled *bool `json:"tts_enabled,omitempty"`
|
||||||
DetailLevel *string `json:"detail_level,omitempty"`
|
DetailLevel *string `json:"detail_level,omitempty"`
|
||||||
Language *string `json:"language,omitempty"`
|
Language *string `json:"language,omitempty"`
|
||||||
|
Scenario *string `json:"scenario,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply 将 patch 中的非 nil 字段覆盖到 cfg。
|
// Apply 将 patch 中的非 nil 字段覆盖到 cfg。
|
||||||
@@ -45,6 +47,9 @@ func (p SessionConfigPatch) Apply(cfg *SessionConfig) {
|
|||||||
if p.Language != nil {
|
if p.Language != nil {
|
||||||
cfg.Language = *p.Language
|
cfg.Language = *p.Language
|
||||||
}
|
}
|
||||||
|
if p.Scenario != nil {
|
||||||
|
cfg.Scenario = *p.Scenario
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// User 用户。
|
// User 用户。
|
||||||
@@ -81,6 +86,7 @@ type WsConfig struct {
|
|||||||
TTSEnabled *bool `json:"tts_enabled,omitempty"`
|
TTSEnabled *bool `json:"tts_enabled,omitempty"`
|
||||||
DetailLevel *string `json:"detail_level,omitempty"`
|
DetailLevel *string `json:"detail_level,omitempty"`
|
||||||
Language *string `json:"language,omitempty"`
|
Language *string `json:"language,omitempty"`
|
||||||
|
Scenario *string `json:"scenario,omitempty"`
|
||||||
} `json:"payload"`
|
} `json:"payload"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
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,402 +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)
|
|
||||||
llmReq := llm.Request{
|
|
||||||
Image: image,
|
|
||||||
Text: userText,
|
|
||||||
History: history,
|
|
||||||
Language: 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)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
@@ -276,15 +310,72 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
|||||||
TTSEnabled: msg.Payload.TTSEnabled,
|
TTSEnabled: msg.Payload.TTSEnabled,
|
||||||
DetailLevel: msg.Payload.DetailLevel,
|
DetailLevel: msg.Payload.DetailLevel,
|
||||||
Language: msg.Payload.Language,
|
Language: msg.Payload.Language,
|
||||||
|
Scenario: msg.Payload.Scenario,
|
||||||
}
|
}
|
||||||
if err := client.sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
|
if err := client.sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
|
||||||
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)
|
||||||
@@ -312,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 '情景最后更新时间';
|
||||||
42
deploy.sh
42
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -69,4 +93,4 @@ case "${1:-}" in
|
|||||||
logs) shift; cmd_logs "$@" ;;
|
logs) shift; cmd_logs "$@" ;;
|
||||||
status) cmd_status ;;
|
status) cmd_status ;;
|
||||||
*) usage; exit 1 ;;
|
*) usage; exit 1 ;;
|
||||||
esac
|
esac
|
||||||
@@ -17,34 +17,70 @@ 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
|
||||||
|
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.xuanyuan.me/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
|
||||||
- ./backend/migrations:/docker-entrypoint-initdb.d
|
- ./backend/migrations:/docker-entrypoint-initdb.d
|
||||||
networks:
|
networks:
|
||||||
- camtalk-net
|
- camtalk-net
|
||||||
|
healthcheck:
|
||||||
|
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
|
||||||
|
timeout: 3s
|
||||||
|
retries: 10
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
|
|
||||||
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) — 安全设计
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user