feat: v2 版本 #206

Merged
huanghaosheng merged 201 commits from v2 into main 2026-06-22 16:01:28 +08:00
227 changed files with 45666 additions and 7141 deletions

View File

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

View File

@@ -2,20 +2,34 @@ name: Deploy
on:
push:
branches: [main]
branches: [main, v2]
jobs:
deploy:
runs-on: aliyun
steps:
- name: Checkout
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
- name: Deploy
run: |
chmod +x deploy.sh
./deploy.sh build
./deploy.sh restart
sed -i 's/dl-cdn.alpinelinux.org/mirrors.aliyun.com/g' /etc/apk/repositories
apk add --no-cache rsync docker-cli docker-cli-compose
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

3
.gitignore vendored
View File

@@ -24,3 +24,6 @@ Thumbs.db
# ---- Obsidian ----
.obsidian/
.claudian/
修改过程笔记/
学习复盘/

123
CLAUDE.md
View File

@@ -1,106 +1,73 @@
# CLAUDE.md
本文件为 Claude Code (claude.ai/code) 在本仓库中工作时提供指引。
CamTalk — 多模态实时 AI 视觉对话助手(摄像头 + 麦克风 + 视觉 + 语音 AI
## 项目概述
CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。项目目前处于设计文档阶段,源代码正在逐步构建。
> **文档优先原则:** 执行任何开发任务前,先读取 `docs/` 下的相关设计文档(架构、接口、技术选型等),以文档为最高依据。代码实现应与文档一致;若有偏差,优先更新文档(尤其是接口文档)。
> **文档优先原则:** 开发前先读 `docs/` 设计文档,以文档为准;若代码与文档不一致,优先更新文档(尤其接口文档)。详细设计见 `docs/01-13` 系列文档。**注意**`docs/Eino/` 框架文档内容庞大(~75 个文件),仅在需要了解 Eino Graph/节点/Callback 等框架细节时才读取。
## 架构
三层系统:
三层系统:前端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()`
2. **Go 网关**Gin, gorilla/websocket, Viper, Zap—— WebSocket 服务器、会话管理、AI 编排。每个 WebSocket 连接一个 goroutine。
3. **云端 AI 服务** —— 通过 OpenAI 兼容接口可灵活切换。默认GPT-4oLLM、DeepgramSTT、OpenAI TTS。仅通过 Go 网关访问,浏览器不直连。
**AI 编排流水线**Eino Graph 7 节点 DAG`STT → History → ChatModel → Msg2Str → Splitter → TTS → Done`。LLM token 通过 Callback 实时推送TTS 逐句并行合成。
**关键模式**LLM 文本流和 TTS 音频流并行推送给客户端,以最小化感知延迟
**会话存储**TieredManagerL1 Memory → L2 Redis → L3 PostgreSQL 三级存储30 分钟 TTLRedis 故障自动降级
**存储**MVP 阶段使用进程内存(`MemoryManager`Redis 实现已就绪可通过配置切换PostgreSQL 为规划中。Repository 接口模式(`HistoryRepository``UsageRepository`MVP 用内存实现
**鉴权**JWT 双 token 轮转Access 120min + Refresh 7d重放攻击检测DB hash 校验Redis 缓存装饰器
## 技术栈
| 层级 | 技术 |
|------|------|
| 前端 | React 18, TypeScript, Vite, @ricky0123/vad-web |
| 后端 | Go, Gin, gorilla/websocket, Viper, Zap |
| LLM | GPT-4o默认通过 OpenAI 兼容接口可切换) |
| STT | Deepgram默认 / MiMo ASR |
| TTS | OpenAI TTS默认 / MiMo TTS |
前端React 18 + TypeScript + ViteVAD@ricky0123/vad-webONNX Runtime
后端Go 1.25+, Gin, WebSocket, Viper, Zap, CloudWeGo Eino Graph
AIDashScope qwen3-vl-plus, MiMo ASR/TTS可切换 Deepgram/OpenAI TTS
存储PostgreSQL 15 + Redis 7
## 构建与运行命令
## 快速启动
```bash
# 前端
cd frontend && npm install
npm run dev # Vite 开发服务器
npm run build # 生产构建
npm run lint # ESLint 检查
npm run test # Vitest 测试
# 后端
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 ./... # 静态分析
# 前端npm run devVite代理 /ws 和 /api 到 :8080
# 后端go run ./cmd/server监听 :8080
# 生产:./deploy.sh up4 容器frontend/backend/postgres/redis
```
基础设施MVP 使用进程内存管理会话状态。Redis 已实现可通过配置切换PostgreSQL 为规划中。
核心环境变量(`.env.example``CAMTALK_AI_LLM_API_KEY`, `CAMTALK_AI_STT_API_KEY`, `CAMTALK_AUTH_JWT_SECRET`, `CAMTALK_STORAGE_DSN`
## WebSocket 协议
配置优先级:环境变量 > `config.{APP_ENV}.yaml` > `config.yaml`
环境切换:`APP_ENV=dev|prod`dev 默认prod 启用限流 + 严格 CORS
端点:`ws://localhost:8080/ws`
## 协议与 API
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。完整契约见 `docs/03-接口文档.md`
**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`),订阅模式,自动重连
**客户端 → 服务端**`query`(图像 Base64 + 音频 Base64`config``interrupt``ping`
**服务端 → 客户端**`connected``stt_result``llm_chunk``llm_done``tts_audio``error``pong`
**REST API**`/api/auth/*`(注册/登录/刷新/登出),`/api/conversations/*`CRUD + 消息分页),`/api/health`
**心跳**客户端每 30 秒 ping服务端 60 秒无 ping 断开连接。
**重连**:指数退避 + 抖动 —— 1s, 2s, 4s, 8s… 最大 30s。
**错误码**`INVALID_MESSAGE`, `SESSION_NOT_FOUND`, `RATE_LIMITED`, `IMAGE_TOO_LARGE`, `LLM_TIMEOUT`, `STT/TTS/LLM_ERROR`, `INVALID_TOKEN`, 等
## REST API辅助
## 关键文件路径
- `GET /api/health` — 健康检查(版本、运行时间、活跃会话数)
- `POST /api/sessions` — 创建会话可选MVP 在 WS 连接时自动创建)
- `DELETE /api/sessions/{id}` — 销毁会话
**后端核心**
- `backend/internal/eino/` — Graph 定义、节点、Callback、Adapter、State
- `backend/internal/session/tiered.go` — 三级会话存储
- `backend/internal/store/` — Repository 实现PG + 内存 + Redis 缓存)
- `backend/internal/ws/handler.go` — WebSocket 连接管理
- `backend/migrations/` — SQL 迁移文件
## 错误码
`INVALID_MESSAGE``SESSION_NOT_FOUND``RATE_LIMITED``IMAGE_TOO_LARGE``AUDIO_TOO_SHORT``LLM_TIMEOUT``LLM_ERROR``STT_ERROR``TTS_ERROR``INTERNAL_ERROR`
## 前端组件结构
| 组件 | 职责 |
|------|------|
| `CameraManager` | 摄像头流采集 |
| `MicManager` | 麦克风音频采集 |
| `EdgeProcessor` | VAD + 关键帧检测Canvas 像素比较) |
| `WebSocketManager` | WebSocket 连接生命周期管理 |
| `ChatPanel` | 消息展示 |
| `VideoPreview` | 摄像头画面预览 |
## 后端模块结构
| 模块 | 职责 |
|------|------|
| WebSocket Handler | 连接管理、单播消息推送 |
| Session Manager | 会话状态、对话历史Memory/Redis30 分钟 TTL |
| AI Orchestrator | STT→LLM→TTS 流式并行管道编排 |
| AI Service Layer | AI 服务抽象层STT/LLM/TTS 多 provider |
| REST API | 健康检查、会话管理Gin 路由) |
| Models | 数据模型定义 |
| Model Router | 按请求选择 AI 模型(规划中) |
| Rate Limiter | 按用户的令牌桶速率限制(规划中) |
**前端核心**
- `frontend/src/hooks/useVisionSession.ts` — 核心会话 Hook~500 行)
- `frontend/src/lib/websocket.ts` — WebSocket 客户端单例
- `frontend/src/lib/auth.tsx` — JWT 自动刷新 + AuthProvider
- `frontend/src/lib/api.ts` — REST 客户端401 拦截 + token 刷新)
- `frontend/src/lib/ttsPlayer.ts` — 流式 TTS 音频播放队列
- `frontend/vite.config.ts` — VAD 模型文件自动复制 + 代理配置
## 编码规范
- **Go**遵循标准 Go 规范。所有 AI 调用使用 `context.Context` 做取消/超时。并发 map 访问使用 `sync.RWMutex`。结构体标签用 `json:"snake_case"`
- **TypeScript**:严格模式。所有数据模型用接口定义。WebSocket 消息类型用可辨识联合类型(`type` 字段)
- **提交信息**Conventional Commits 格式,描述用中文。示例:`feat: 添加 WebSocket 连接管理``fix: 修复心跳超时判断``docs: 更新接口文档`
- **禁止自动 push**:除非用户明确要求。
- **文档优先**:实现功能前先读取 `docs/` 下的相关设计文档。实现与文档不一致时,优先更新 `docs/` 下的接口文档。
- **Go**标准规范,`context.Context` 超时控制,`sync.RWMutex` 并发保护,`json:"snake_case"` 标签,编译期接口检查 `var _ Interface = (*Impl)(nil)`
- **TypeScript**:严格模式,接口定义数据模型,WebSocket 消息用可辨识联合类型(`type` 字段区分
- **CORS**:禁止后端代码/配置文件配置 CORS统一由代理层处理开发环境 Vite proxy生产环境 Nginx
- **提交信息**Conventional Commits中文描述`feat: 添加 WebSocket 心跳`
- **禁止自动 push**:除非用户明确要求
- **文档优先**:开发前先读 `docs/` 设计文档,代码与文档不一致时优先更新文档

598
README.md
View File

@@ -1,167 +1,561 @@
# CamTalk
多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。
<div align="center">
- **路演视频**[哔哩哔哩弹幕网——七牛云第四批议题1](https://www.bilibili.com/video/BV1dDJK6cE5S/)
- **线上体验**http://8.161.227.145:9000
**多模态实时 AI 视觉对话助手**
> ⚠️ **注意**:由于线上地址使用 HTTP 协议,浏览器默认禁止在非 HTTPS 环境下调用摄像头和麦克风。需要按以下步骤配置 Chrome 浏览器:
用户通过摄像头和麦克风与 AI 交互AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应
[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
[![Go Version](https://img.shields.io/badge/Go-1.25+-00ADD8?logo=go)](https://go.dev/)
[![React](https://img.shields.io/badge/React-18-61DAFB?logo=react)](https://react.dev/)
[![TypeScript](https://img.shields.io/badge/TypeScript-5-3178C6?logo=typescript)](https://www.typescriptlang.org/)
[路演视频](https://www.bilibili.com/video/BV1dDJK6cE5S/) • [在线体验](http://8.161.227.145:9000) • [文档](docs/README.md)
</div>
---
> ⚠️ **在线体验提示**:由于演示环境使用 HTTP 协议,需配置 Chrome 允许非 HTTPS 下访问摄像头/麦克风:
>
> 1. 在浏览器地址栏中输入 `chrome://flags/#unsafely-treat-insecure-origin-as-secure`,回车
> 2. 将 **Insecure origins treated as secure** 选项设置为 **Enabled**(已启用)
> 3. 在输入框中输入 `http://8.161.227.145:9000` 地址
> 4. 点击右下角弹出的 **Relaunch** 按钮,自动重启浏览器
>
> 重启后即可在该 HTTP 地址下正常调用摄像头和麦克风。
> 1. 访问 `chrome://flags/#unsafely-treat-insecure-origin-as-secure`
> 2. 启用该选项,并在输入框填入 `http://8.161.227.145:9000`
> 3. 点击 **Relaunch** 重启浏览器
![](docs/pictures/1.png)
![CamTalk 界面截图](docs/pictures/1.png)
## 架构
## ✨ 核心特性
三层系统,前端做轻量预处理,后端做智能编排,云端 AI 服务按需调用:
- 🎥 **多模态理解**:摄像头视觉 + 麦克风语音双输入AI 理解完整场景
- 🗣️ **自然对话**:基于 VAD 的端到端语音交互,低延迟流式响应
- 🚀 **实时推送**LLM 文本流 + TTS 音频流并行推送,感知延迟 < 0.5 秒
- 🎭 **情景模式**:自由对话、面试官、英语老师等多场景支持
- 💾 **对话历史**:自动保存会话,支持搜索、重命名、删除、时间分组
- 🔐 **安全认证**JWT 双 token 轮转 + Refresh Token Rotation 防重放
- 📊 **三级存储**Memory → Redis → PostgreSQL 自动降级,保障可靠性
- 🌐 **国际化**:支持中文、英文、日文界面
## 🏗️ 系统架构
CamTalk 采用**三层架构**:前端轻量预处理 → Go 网关智能编排 → 云端 AI 按需调用
```mermaid
graph TB
subgraph client[浏览器客户端]
A1[媒体采集]
A2[VAD 语音检测]
A3[关键帧检测]
A4[UI 渲染]
subgraph Browser["🌐 浏览器客户端"]
UI["React UI 渲染"]
VAD["VAD 语音检测"]
Media["媒体采集"]
end
subgraph gateway[Go 网关 :8080]
B1[WebSocket Handler]
B2[Session Manager]
B3[AI Orchestrator]
B4[REST API]
subgraph Gateway["⚙️ Go 网关 (Eino Graph)"]
WS["WebSocket Handler"]
Auth["JWT 认证"]
Session["会话管理 (三级存储)"]
Orch["AI 编排器 (7节点DAG)"]
end
subgraph cloud[云端 AI 服务]
C1[STT 语音识别]
C2[LLM 多模态推理]
C3[TTS 语音合成]
subgraph AI["☁️ 云端 AI 服务"]
STT["STT (MiMo/Deepgram)"]
LLM["LLM (qwen3-vl-plus)"]
TTS["TTS (MiMo/OpenAI)"]
end
client <-->|WebSocket| gateway
gateway <-->|HTTP| cloud
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
| 层级 | 技术 |
|------|------|
| 前端 | 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 |
```
START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END
```
## 项目结构
**核心优势**
- **流式处理**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>声明式 DAGStream 模式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/
├── frontend/ # 浏览器客户端
├── frontend/ # 🌐 浏览器客户端
│ └── src/
│ ├── components/ # UI 组件
│ │ ├── LandingPage/ # 登录着陆页 + LoginModal
│ │ ├── CameraManager/ # 摄像头流采集
│ │ ├── MicManager/ # 麦克风音频采集
│ │ ├── EdgeProcessor/ # VAD + 关键帧检测
│ │ ├── WebSocketManager/ # WS 连接管理
│ │ ├── ChatPanel/ # 消息展示
│ │ ── VideoPreview/ # 摄像头画面预览
│ │ ├── ConfigPanel/ # 配置面板
│ │ └── Toast/ # 通知提示
│ │ ├── MicManager/ # 麦克风音频采集 + VAD
│ │ ├── WebSocketManager/ # WS 连接生命周期
│ │ ├── ChatPanel/ # 消息展示 + 流式回复
│ │ ├── SessionSidebar/ # 对话历史侧边栏
│ │ ── ConfigPanel/ # 配置面板(主题/TTS/语言/场景)
│ ├── hooks/ # 自定义 Hooks
│ │ ├── useVisionSession.ts # 核心会话 Hook
│ │ ├── useVisionSession.ts # 核心会话 Hook (~500 行)
│ │ ├── useSessionList.ts # 对话列表管理
│ │ └── useObservationMode.ts # 观察模式
│ ├── lib/ # 工具库
│ │ ├── websocket.ts # WebSocket 连接管理
│ │ ├── audio.ts # 音频编码
│ │ ├── ttsPlayer.ts # TTS 播放器
│ │ ── sampling.ts # 采样策略
│ │ ├── websocket.ts # WebSocket 单例(心跳/重连/订阅)
│ │ ├── api.ts # REST 客户端401拦截+刷新)
│ │ ├── auth.tsx # AuthProviderJWT 自动刷新)
│ │ ── ttsPlayer.ts # TTS 流式播放队列
│ │ └── i18n/ # 国际化zh-CN/en-US/ja-JP
│ └── types/ # TypeScript 类型定义
├── backend/ # Go 网关
│ ├── cmd/server/ # 入口
├── backend/ # ⚙️ Go 网关
│ ├── cmd/server/ # 服务入口main.go
│ └── 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 服务抽象层
│ │ ├── llm/ # LLM 服务OpenAI 兼容)
│ │ ├── stt/ # STT 服务Deepgram/MiMo
│ │ └── tts/ # TTS 服务OpenAI/MiMo
│ ├── orchestrator/ # AI 编排器STT→LLM→TTS 管道)
│ ├── session/ # 会话管理Memory/Redis
│ │ ├── llm/ # LLM 提示词与场景
│ │ ├── stt/ # STT 服务(MiMo/Deepgram
│ │ └── tts/ # TTS 服务(MiMo/OpenAI
│ ├── ws/ # WebSocket Handler
│ ├── api/ # REST API
│ ├── config/ # 配置管理
── models/ # 数据模型
│ ├── errors/ # 错误码
│ └── logger/ # 日志
├── docs/ # 设计文档
└── CLAUDE.md # Claude Code 指引
│ ├── api/ # REST APIAuth/Conversation
│ ├── config/ # 配置管理Viper
── logger/ # 日志Zap + Trace ID
├── migrations/ # 📊 数据库迁移(嵌入式 SQL
├── docs/ # 📚 设计文档
│ ├── 01-架构设计.md
│ ├── 02-接口文档.md
│ ├── 08-Eino框架与编排设计.md
│ ├── 10-鉴权体系.md
│ └── 13-日志追踪.md
├── deploy.sh # 🐳 部署脚本Docker Compose
├── docker-compose.yml # 容器编排配置
└── CLAUDE.md # 🤖 Claude Code 开发指引
```
## 快速开始
## 🚀 快速开始
### 前置条件
- Node.js >= 18
- Go >= 1.24
- **Node.js** >= 18
- **Go** >= 1.25
- **PostgreSQL** >= 15可选 Docker
- **Redis** >= 7可选用于缓存加速
### 前端
### 本地开发
#### 1. 克隆项目
```bash
cd frontend
npm install
npm run dev # Vite 开发服务器 http://localhost:5173
git clone https://github.com/yourusername/CamTalk.git
cd CamTalk
```
### 后端
#### 2. 配置环境变量
```bash
# 复制环境变量模板
cp backend/.env.example backend/.env
# 编辑 .env 文件,填入以下必需配置:
# - CAMTALK_AUTH_JWT_SECRET使用 openssl rand -hex 32 生成)
# - CAMTALK_STORAGE_DSNPostgreSQL 连接字符串)
# - CAMTALK_AI_LLM_API_KEYDashScope API Key
# - CAMTALK_AI_STT_API_KEYMiMo/Deepgram API Key
# - CAMTALK_AI_TTS_API_KEYMiMo/OpenAI API Key
```
#### 3. 启动后端
```bash
cd backend
# 安装依赖
go mod download
go run ./cmd/server # 启动网关 :8080
```
### 配置
# 运行数据库迁移(自动创建表)
go run ./cmd/server migrate
后端配置文件位于 `backend/config.yaml`,支持环境变量覆盖(前缀 `CAMTALK_`)。
```bash
# 最小启动(需要至少一个 AI 服务的 API Key
cd backend
CAMTALK_AI_LLM_API_KEY=sk-xxx \
CAMTALK_AI_STT_API_KEY=xxx \
# 启动服务(监听 :8080
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`
**服务端 → 客户端**`connected``stt_result``llm_chunk``llm_done``tts_audio``error``pong`
#### 5. 访问应用
完整协议见 [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 逐 tokenCallback |
| `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) | 项目目标与核心挑战 |
| [02-系统架构](docs/02-系统架构.md) | 三层架构、技术栈、部署方案 |
| [03-接口文档](docs/03-接口文档.md) | WebSocket 协议、REST API、配置管理 |
| [04-技术选型](docs/04-技术选型.md) | AI 服务栈、持久化层、前端边缘处理选型 |
| [05-用户故事](docs/05-用户故事.md) | 用户场景与优先级 |
| [06-语音交互](docs/06-语音交互.md) | VAD → STT → LLM → TTS 全链路 |
| [07-视觉理解](docs/07-视觉理解.md) | 帧采样、关键帧检测、多模态输入 |
| [08-成本控制](docs/08-成本控制.md) | 采样策略、端云协同、模型分级 |
| [01-架构设计](docs/01-架构设计.md) | 三层架构、技术栈、数据库设计、部署方案 |
| [02-接口文档](docs/02-接口文档.md) | WebSocket 协议、REST API、AI 服务层、编排器、配置管理 |
| [08-Eino框架与编排设计](docs/08-Eino框架与编排设计.md) | Eino Graph 7 节点 DAG、节点实现、流式处理、Callback AOP |
| [10-鉴权体系](docs/10-鉴权体系.md) | JWT 双 token 轮转、Refresh Token Rotation、密码安全、中间件 |
| [11-令牌桶限流](docs/11-令牌桶限流.md) | 限流算法、配置策略、生产环境保护 |
| [13-日志追踪](docs/13-日志追踪.md) | Zap 日志、Trace ID 全链路追踪、日志级别 |
## 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
View 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
View File

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

View File

@@ -8,11 +8,19 @@ ENV GOPROXY=https://goproxy.cn,https://goproxy.io,direct
# 先复制依赖清单,利用 Docker 缓存层
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 . .
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
@@ -21,9 +29,9 @@ RUN apk add --no-cache ca-certificates tzdata
WORKDIR /app
# 复制二进制和配置
# 复制二进制和配置文件(敏感配置通过 docker-compose env_file 注入覆盖)
COPY --from=builder /camtalk .
COPY config.yaml .
COPY config/ ./config/
EXPOSE 8080

View File

@@ -10,17 +10,20 @@ import (
"time"
"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/auth"
"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"
eino "github.com/hhs/camtalk/internal/eino"
"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/store"
"github.com/hhs/camtalk/internal/trace"
"github.com/hhs/camtalk/internal/ws"
migrations "github.com/hhs/camtalk/migrations"
)
@@ -32,8 +35,8 @@ var Version string
var startTime = time.Now()
func main() {
// 加载配置
cfg, err := config.Load()
// 加载配置(工作目录用于定位 .env 和 config.yaml
cfg, err := config.Load(".")
if err != nil {
panic("failed to load config: " + err.Error())
}
@@ -47,20 +50,27 @@ func main() {
"addr", cfg.Server.Addr(),
)
// 初始化存储层(条件初始化 PostgreSQL
// 初始化存储层(三级存储架构L1 内存 → L2 Redis → L3 PostgreSQL
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var userRepo store.UserRepository
var msgRepo store.MessageRepository
var sessRepo store.SessionRepository
var pool *pgxpool.Pool // 数据库连接池
if cfg.Storage.Driver == "postgres" {
if cfg.Storage.DSN == "" {
logger.Log.Fatalw("storage.dsn is required when storage.driver is postgres",
// L3: PostgreSQL冷数据持久化层
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")
}
pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN)
var err error
pool, err = store.NewPostgresPool(ctx, dsn)
if err != nil {
logger.Log.Fatalw("failed to connect to postgres", "error", err)
}
@@ -74,27 +84,76 @@ func main() {
userRepo = store.NewPgUserRepository(pool)
msgRepo = store.NewPgMessageRepository(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 {
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 sessionOpts []session.Option
if msgRepo != nil {
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
if cfg.Storage.Redis.Enabled {
// L1 + L2 + L3 三级存储
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 服务
logger.Log.Infow("initializing AI services",
@@ -116,9 +175,6 @@ func main() {
sttService = stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, cfg.AI.STT.Timeout, logger.Log)
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
switch strings.ToLower(cfg.AI.TTS.Provider) {
case "mimo", "xiaomi":
@@ -129,8 +185,16 @@ func main() {
logger.Log.Infow("TTS service initialized", "provider", "openai", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "speed", cfg.AI.TTS.Speed)
}
// 初始化 Orchestrator
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr, cfg)
// 初始化 Eino Graph + Orchestrator
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(
@@ -140,13 +204,32 @@ func main() {
)
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 模式
if cfg.App.Env == "prod" {
gin.SetMode(gin.ReleaseMode)
}
r := gin.New()
r.Use(gin.Recovery())
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
r.Use(trace.GinLogger()) // 第二层:记录请求
r.Use(trace.GinRecovery()) // 第三层panic 恢复
// REST API
apiGroup := r.Group("/api")
@@ -160,14 +243,29 @@ func main() {
// Auth REST 端点
authHandler := api.NewAuthHandler(authService, tokenMgr)
authHandler.RegisterRoutes(apiGroup)
authHandler.RegisterRoutes(apiGroup, limiter)
// Conversation REST 端点
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
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
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr))
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter, userScenarioRepo))
// HTTP Server
srv := &http.Server{

View File

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

View File

@@ -0,0 +1,68 @@
# CamTalk 开发环境配置
# 通过 APP_ENV=dev 加载此文件,覆盖 config.yaml 中的配置
server:
host: "0.0.0.0"
port: 8080
heartbeat_interval: 30
heartbeat_timeout: 60
allowed_origins: [] # 开发环境允许所有来源
session:
ttl: 30 # 开发环境会话较短,方便测试过期逻辑
max_history: 20
ai:
stt:
provider: mimo # 与生产环境一致
model: mimo-v2.5-asr
endpoint: "https://api.xiaomimimo.com/v1"
timeout: 10 # 开发环境超时较长,方便调试
http_client_timeout: 30
llm:
provider: dashscope # 与生产环境一致
model: qwen3-vl-plus
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
timeout: 60 # 开发环境 LLM 超时较长
http_client_timeout: 120
tts:
provider: mimo # 与生产环境一致
model: mimo-v2.5-tts
voice: mimo_default
speed: 1.0
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
timeout: 10
http_client_timeout: 30
output_format: mp3
sample_rate: 24000
storage:
redis:
enabled: true # 开发环境启用 Redis测试三级存储
persistence:
enabled: true # 开发环境启用持久化
redis:
addr: "localhost:6379" # 本地 Redis
password: ""
db: 0
auth:
access_ttl: 120 # 开发环境 Access Token 2 小时,方便调试
refresh_ttl: 10080 # 7 天
ratelimit:
enabled: false # 开发环境关闭限流,方便测试
query:
capacity: 10
rate: 0.2
login:
capacity: 5
rate: 0.1
register:
capacity: 3
rate: 0.05
log:
level: debug # 开发环境 debug 日志
format: console # 控制台格式,易读

View File

@@ -0,0 +1,71 @@
# CamTalk 生产环境配置
# 通过 APP_ENV=prod 加载此文件,覆盖 config.yaml 中的配置
server:
host: "0.0.0.0"
port: 8080
read_timeout: 30
write_timeout: 30
shutdown_timeout: 15 # 生产环境优雅关闭时间稍长
heartbeat_interval: 30
heartbeat_timeout: 60
session:
ttl: 60 # 生产环境会话 1 小时
max_history: 20
ai:
stt:
provider: mimo # 生产环境推荐 MiMo性价比高
model: mimo-v2.5-asr
endpoint: "https://api.xiaomimimo.com/v1"
timeout: 5 # 生产环境严格超时控制
http_client_timeout: 30
llm:
provider: dashscope # 生产环境推荐通义千问,稳定性好
model: qwen3-vl-plus
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
timeout: 30
http_client_timeout: 60
tts:
provider: mimo # 生产环境推荐 MiMo TTS
model: mimo-v2.5-tts
voice: mimo_default
speed: 1.0
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
timeout: 5
http_client_timeout: 30
output_format: mp3
sample_rate: 24000
storage:
redis:
enabled: true # 生产环境必须启用 Redis
persistence:
enabled: true # 生产环境必须启用持久化
driver: postgres
redis:
addr: "redis:6379" # Docker Compose 内部服务名
password: "" # 密码通过 CAMTALK_REDIS_PASSWORD 环境变量设置
db: 0
auth:
access_ttl: 120 # 生产环境 Access Token 2 小时
refresh_ttl: 10080 # Refresh Token 7 天
ratelimit:
enabled: true # 生产环境启用限流
query:
capacity: 10 # 允许突发 10 个请求
rate: 0.2 # 每 5 秒恢复 1 个令牌
login:
capacity: 5 # 防暴力破解
rate: 0.1 # 每 10 秒恢复 1 次
register:
capacity: 3 # 防批量注册
rate: 0.05 # 每 20 秒恢复 1 次
log:
level: info # 生产环境 info 级别
format: json # JSON 格式,便于日志收集和分析

View File

@@ -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 # 是否启用 RedisL2 热数据层)
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

View File

@@ -3,23 +3,36 @@ module github.com/hhs/camtalk
go 1.25.0
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/golang-jwt/jwt/v5 v5.3.1
github.com/google/uuid v1.6.0
github.com/gorilla/websocket v1.5.3
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/spf13/viper v1.21.0
github.com/stretchr/testify v1.11.1
go.uber.org/zap v1.28.0
golang.org/x/crypto v0.31.0
)
require (
github.com/bytedance/sonic v1.11.6 // indirect
github.com/bytedance/sonic/loader v0.1.1 // indirect
github.com/bahlo/generic-list-go v0.2.0 // 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/cloudwego/base64x v0.1.4 // indirect
github.com/cloudwego/iasm v0.2.0 // indirect
github.com/cloudwego/base64x v0.1.6 // indirect
github.com/cloudwego/eino-ext/libs/acl/openai v0.1.17 // 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/gabriel-vasile/mimetype v1.4.3 // 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-viper/mapstructure/v2 v2.4.0 // 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/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/cpuid/v2 v2.2.10 // 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/meguminnnnnnnnn/go-openai v0.1.2 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // 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/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.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/spf13/afero v1.15.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/twitchyliquid64/golang-asm v0.15.1 // 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/multierr v1.10.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/arch v0.8.0 // indirect
golang.org/x/crypto v0.23.0 // indirect
golang.org/x/arch v0.11.0 // indirect
golang.org/x/exp v0.0.0-20230713183714-613f0c0eb8a1 // indirect
golang.org/x/net v0.25.0 // indirect
golang.org/x/sync v0.17.0 // indirect
golang.org/x/sys v0.30.0 // indirect

View File

@@ -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/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
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/bytedance/sonic v1.11.6/go.mod h1:LysEHSvpvDySVdC2f87zGWf6CIKJcAvqab1ZaiQtds4=
github.com/bytedance/sonic/loader v0.1.1 h1:c+e5Pt1k/cy5wMveRDyk2X4B9hF4g7an8N3zCYjJFNM=
github.com/bytedance/sonic/loader v0.1.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU=
github.com/buger/jsonparser v1.1.1 h1:2PnMjfWD7wBILjqQbt530v576A/cAbQvEW9gGIpYMUs=
github.com/buger/jsonparser v1.1.1/go.mod h1:6RYKKt7H4d4+iWqouImQ9R2FZql3VbhNgx27UK13J/0=
github.com/bugsnag/bugsnag-go v1.4.0/go.mod h1:2oa8nejYd4cQ/b0hMIopN0lCRxU0bueqREvZLWFrtK8=
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/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudwego/base64x v0.1.4 h1:jwCgWpFanWmN8xoIUHa2rtzmkd5J2plF/dnLS6Xd/0Y=
github.com/cloudwego/base64x v0.1.4/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w=
github.com/cloudwego/iasm v0.2.0 h1:1KNIy1I1H9hNNFEEH3DVnI4UujN+1zjpuk6gwHLTssg=
github.com/cloudwego/iasm v0.2.0/go.mod h1:8rXZaNYT2n95jn+zTI1sDr+IgcD2GVs0nlbbQPiEFhY=
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
github.com/cloudwego/eino v0.9.9 h1:x63hvRif6ANPh9YEPoTIrp1potEeoLQFAjOclKaX/Kg=
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.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
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/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/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/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/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/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/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4=
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/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
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/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/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
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/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/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/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
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/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/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/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/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/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/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/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/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-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg=
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/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/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/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/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/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/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/go.mod h1:3n1Cwaq1E1/1lhQhtRK2ts/ZwZEhjcQeJQ1RuC6Q/8U=
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/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY=
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.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/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.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.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/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
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/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
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/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/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
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/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/go.mod h1:20+QtiLqy0Nd6FdQB9TLXag12DsQkrbs3htMFfDN80Y=
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.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
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.8.0 h1:3wRIsP3pM4yUptoR96otTUOXI367OS0+c9eeRi9doIc=
golang.org/x/arch v0.8.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
golang.org/x/crypto v0.23.0 h1:dIJU/v2J8Mdglj/8rJ6UUOM3Zc9zLZxVZwwxMooUSAI=
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
golang.org/x/arch v0.11.0 h1:KXV8WWKCXm6tRpLirl2szsO5j/oOODwZf4hATmGVNs4=
golang.org/x/arch v0.11.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4=
golang.org/x/crypto v0.31.0 h1:ihbySMvVjLAeSH1IbfcRTkD/iNscyz8rGzjF/E5hV6U=
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/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/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.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
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/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
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 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
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.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
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=

View File

@@ -1,239 +0,0 @@
package llm
import (
"bufio"
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"go.uber.org/zap"
)
// OpenAIService 基于 OpenAI Chat Completions API 的 LLM 实现。
type OpenAIService struct {
apiKey string
model string
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
client *http.Client
}
// NewOpenAIService 创建 OpenAI LLM 服务。
// model、endpoint 由 config 层保证非空。
func NewOpenAIService(apiKey, model, endpoint string, timeoutSec, httpClientTimeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 10 * time.Second
}
httpClientTimeout := time.Duration(httpClientTimeoutSec) * time.Second
if httpClientTimeout <= 0 {
httpClientTimeout = 60 * time.Second
}
return &OpenAIService{
apiKey: apiKey,
model: model,
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: httpClientTimeout},
}
}
// --- OpenAI API 请求/响应结构 ---
type chatRequest struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages"`
Stream bool `json:"stream"`
}
type chatMessage struct {
Role string `json:"role"`
Content []contentPart `json:"content"`
}
type contentPart struct {
Type string `json:"type"`
Text string `json:"text"`
ImageURL *imageURL `json:"image_url,omitempty"`
}
type imageURL struct {
URL string `json:"url"`
}
// streamDelta SSE 流式响应的单个 delta。
type streamDelta struct {
Choices []struct {
Delta struct {
Content string `json:"content"`
} `json:"delta"`
FinishReason *string `json:"finish_reason"`
} `json:"choices"`
Usage *struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Model string `json:"model"`
}
// ChatStream 实现 llm.Service。调用 OpenAI Chat Completions API 流式推理。
func (o *OpenAIService) ChatStream(ctx context.Context, req Request) (<-chan Chunk, error) {
// 构建请求
messages := o.buildMessages(req)
body := chatRequest{
Model: o.model,
Messages: messages,
Stream: true,
}
payload, err := json.Marshal(body)
if err != nil {
return nil, fmt.Errorf("llm: marshal request: %w", err)
}
if err != nil {
return nil, fmt.Errorf("llm: marshal request: %w", err)
}
// 创建带超时的 context
ctx, cancel := context.WithTimeout(ctx, o.timeout)
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/chat/completions", bytes.NewReader(payload))
if err != nil {
cancel()
return nil, fmt.Errorf("llm: create request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+o.apiKey)
resp, err := o.client.Do(httpReq)
if err != nil {
cancel()
return nil, fmt.Errorf("llm: send request: %w", err)
}
if resp.StatusCode != http.StatusOK {
cancel()
bodyBytes, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return nil, fmt.Errorf("llm: api error (status %d): %s", resp.StatusCode, string(bodyBytes))
}
// 启动 goroutine 解析 SSE 流
ch := make(chan Chunk, 64)
go func() {
defer close(ch)
defer cancel()
defer resp.Body.Close()
o.parseSSEStream(resp.Body, ch)
}()
return ch, nil
}
// parseSSEStream 解析 SSE 流,将 delta 发送到 channel。
func (o *OpenAIService) parseSSEStream(body io.Reader, ch chan<- Chunk) {
scanner := bufio.NewScanner(body)
scanner.Buffer(make([]byte, 0, 64*1024), 256*1024)
var fullText strings.Builder
var lastModel string
for scanner.Scan() {
line := scanner.Text()
// SSE 格式data: {...}
if !strings.HasPrefix(line, "data: ") {
continue
}
data := strings.TrimPrefix(line, "data: ")
if data == "[DONE]" {
// 流结束,发送最终 chunk
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
return
}
var delta streamDelta
if err := json.Unmarshal([]byte(data), &delta); err != nil {
o.logger.Warnw("llm: unmarshal delta failed", "error", err, "data", data)
continue
}
if delta.Model != "" {
lastModel = delta.Model
}
// 提取增量文本
if len(delta.Choices) > 0 {
content := delta.Choices[0].Delta.Content
if content != "" {
fullText.WriteString(content)
ch <- Chunk{Delta: content, Done: false, Model: lastModel}
}
// 某些模型在最后一个 choice 中携带 usage
if delta.Choices[0].FinishReason != nil && delta.Usage != nil {
ch <- Chunk{
Delta: "",
Done: true,
Model: lastModel,
TokensUsed: &TokenUsage{
Prompt: delta.Usage.PromptTokens,
Completion: delta.Usage.CompletionTokens,
Total: delta.Usage.TotalTokens,
},
}
return
}
}
}
// scanner 结束但没收到 [DONE]
if err := scanner.Err(); err != nil {
o.logger.Warnw("llm: scan error", "error", err)
}
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
}
// buildMessages 构建 OpenAI Chat API 的 messages 数组。
func (o *OpenAIService) buildMessages(req Request) []chatMessage {
var messages []chatMessage
// System prompt情景覆盖优先
messages = append(messages, chatMessage{
Role: "system",
Content: []contentPart{{Type: "text", Text: BuildSystemPrompt(req.Language, "", req.SystemPrompt)}},
})
// 历史消息
for _, msg := range req.History {
messages = append(messages, chatMessage{
Role: msg.Role,
Content: []contentPart{{Type: "text", Text: msg.Content}},
})
}
// 当前用户消息(图像 + 文本)
var parts []contentPart
if len(req.Image) > 0 {
b64 := base64.StdEncoding.EncodeToString(req.Image)
parts = append(parts, contentPart{
Type: "image_url",
ImageURL: &imageURL{URL: "data:image/jpeg;base64," + b64},
})
}
parts = append(parts, contentPart{Type: "text", Text: req.Text})
messages = append(messages, chatMessage{Role: "user", Content: parts})
return messages
}

View File

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

View File

@@ -2,54 +2,147 @@ package llm
import "strings"
// scenarioPrompt 定义单个情景的多语言 system prompt。
// scenarioPrompt 定义单个情景的多语言 system prompt 和首句引导
type scenarioPrompt struct {
ZH string
EN string
JA string
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句话。",
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句话。",
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) string {
func GetScenarioPrompt(scenarioID, language string, customScenarios map[string]string) string {
if scenarioID == "" || scenarioID == "free_chat" {
return ""
}
p, ok := scenarioPrompts[scenarioID]
if !ok {
// 1. 优先查找系统预置情景
if p, ok := scenarioPrompts[scenarioID]; ok {
switch {
case strings.HasPrefix(language, "zh"):
return p.ZH
case strings.HasPrefix(language, "ja"):
return p.JA
default:
return p.EN
}
}
// 2. 查找用户自建情景
if customScenarios != nil {
if customPrompt, ok := customScenarios[scenarioID]; ok {
return customPrompt
}
}
// 3. 默认空字符串
return ""
}
// GetScenarioGreeting 根据情景 ID 和语言获取对应的首句引导。
// 支持系统预置情景和用户自建情景。
// customGreetings: 用户自建情景的首句引导映射表scenarioID → greeting可为 nil
// 返回空字符串表示无此情景或不需要引导(自由对话)。
func GetScenarioGreeting(scenarioID, language string, customGreetings map[string]string) string {
if scenarioID == "" || scenarioID == "free_chat" {
return ""
}
switch {
case strings.HasPrefix(language, "zh"):
return p.ZH
case strings.HasPrefix(language, "ja"):
return p.JA
default:
return p.EN
// 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 ""
}

View File

@@ -11,6 +11,8 @@ import (
"strings"
"time"
"github.com/hhs/camtalk/internal/trace"
"github.com/hhs/camtalk/internal/util"
"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)
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
}

View File

@@ -9,6 +9,8 @@ import (
"net/http"
"time"
"github.com/hhs/camtalk/internal/trace"
"github.com/hhs/camtalk/internal/util"
"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)
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
}

View File

@@ -8,6 +8,8 @@ import (
"github.com/hhs/camtalk/internal/auth"
apperr "github.com/hhs/camtalk/internal/errors"
"github.com/hhs/camtalk/internal/ratelimit"
"github.com/hhs/camtalk/internal/trace"
)
// AuthHandler 提供认证相关的 REST 端点。
@@ -25,11 +27,26 @@ func NewAuthHandler(authService auth.Service, tokenMgr *auth.TokenManager) *Auth
}
// RegisterRoutes 注册认证相关路由到给定的路由组。
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup) {
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup, limiter ratelimit.Limiter) {
authGroup := rg.Group("/auth")
{
authGroup.POST("/register", h.Register)
authGroup.POST("/login", h.Login)
// 注册和登录端点添加限流中间件(按 IP 限流)
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("/logout", auth.AuthMiddleware(h.tokenMgr), h.Logout)
}
@@ -37,6 +54,9 @@ func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup) {
// Register POST /api/auth/register — 用户注册。
func (h *AuthHandler) Register(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
clientIP := c.ClientIP()
var req auth.RegisterRequest
if err := c.ShouldBindJSON(&req); err != nil {
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)
if err != nil {
log.Warnw("register failed",
"username", req.Username,
"client_ip", clientIP,
"error", err)
handleAuthError(c, err)
return
}
log.Infow("register success",
"username", req.Username,
"client_ip", clientIP)
c.JSON(http.StatusCreated, resp)
}
// Login POST /api/auth/login — 用户登录。
func (h *AuthHandler) Login(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
clientIP := c.ClientIP()
var req auth.LoginRequest
if err := c.ShouldBindJSON(&req); err != nil {
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)
if err != nil {
log.Warnw("login failed",
"username", req.Username,
"client_ip", clientIP,
"error", err)
handleAuthError(c, err)
return
}
log.Infow("login success",
"username", req.Username,
"client_ip", clientIP)
c.JSON(http.StatusOK, resp)
}
// Refresh POST /api/auth/refresh — 刷新令牌。
func (h *AuthHandler) Refresh(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
var req auth.RefreshRequest
if err := c.ShouldBindJSON(&req); err != nil {
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)
if err != nil {
log.Warnw("token refresh failed",
"client_ip", c.ClientIP(),
"error", err)
handleAuthError(c, err)
return
}
log.Infow("token refresh success",
"client_ip", c.ClientIP())
c.JSON(http.StatusOK, resp)
}
// Logout POST /api/auth/logout — 登出(需要认证)。
func (h *AuthHandler) Logout(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
userID := c.GetString(auth.ContextKeyUserID)
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 {
log.Errorw("logout failed",
"user_id", userID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to logout",
@@ -150,6 +198,8 @@ func (h *AuthHandler) Logout(c *gin.Context) {
return
}
log.Infow("logout success",
"user_id", userID)
c.JSON(http.StatusOK, gin.H{
"message": "logged out successfully",
})

View File

@@ -47,7 +47,7 @@ func newTestRouter(svc auth.Service) *gin.Engine {
r := gin.New()
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
h := api.NewAuthHandler(svc, tm)
h.RegisterRoutes(r.Group("/api"))
h.RegisterRoutes(r.Group("/api"), nil) // 测试时不启用限流
return r
}
@@ -57,7 +57,7 @@ func newTestRouterWithToken(svc auth.Service) (*gin.Engine, *auth.TokenManager)
r := gin.New()
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
h := api.NewAuthHandler(svc, tm)
h.RegisterRoutes(r.Group("/api"))
h.RegisterRoutes(r.Group("/api"), nil) // 测试时不启用限流
return r, tm
}

View File

@@ -13,6 +13,7 @@ import (
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store"
"github.com/hhs/camtalk/internal/trace"
)
// ConversationHandler 提供对话相关的 REST 端点。
@@ -47,6 +48,7 @@ func (h *ConversationHandler) RegisterRoutes(rg *gin.RouterGroup) {
// List GET /api/conversations — 获取当前用户的对话列表。
func (h *ConversationHandler) List(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
userID := c.GetString(auth.ContextKeyUserID)
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)
if err != nil {
log.Errorw("list conversations failed",
"user_id", userID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to list conversations",
@@ -83,6 +88,7 @@ type CreateConversationRequest struct {
// Create POST /api/conversations — 创建新对话。
func (h *ConversationHandler) Create(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
userID := c.GetString(auth.ContextKeyUserID)
var req CreateConversationRequest
@@ -95,6 +101,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
sessionID, err := h.sessionMgr.Create(c.Request.Context(), userID, cfg)
if err != nil {
log.Errorw("create conversation failed",
"user_id", userID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"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)
if err != nil {
log.Errorw("retrieve created conversation failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to retrieve created conversation",
@@ -111,6 +123,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
return
}
log.Infow("conversation created",
"conversation_id", sess.ID,
"user_id", userID)
c.JSON(http.StatusCreated, gin.H{
"id": sess.ID,
"title": sess.Title,
@@ -144,6 +159,7 @@ type UpdateTitleRequest struct {
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
@@ -176,6 +192,9 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
})
return
}
log.Errorw("update title failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to update title",
@@ -190,6 +209,7 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
// Delete DELETE /api/conversations/:id — 删除对话。
func (h *ConversationHandler) Delete(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
@@ -205,6 +225,9 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
})
return
}
log.Errorw("delete conversation failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to delete conversation",
@@ -221,6 +244,7 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
// - limit: 返回消息数量上限,默认 50
// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
func (h *ConversationHandler) GetMessages(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
@@ -239,6 +263,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
if h.msgRepo != nil {
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
if err != nil {
log.Errorw("get messages failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to get messages",
@@ -246,6 +273,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
return
}
count, _ := h.msgRepo.GetMessageCount(c.Request.Context(), sessionID)
if messages == nil {
messages = []store.StoredMessage{}
}
c.JSON(http.StatusOK, gin.H{
"messages": messages,
"total": count,
@@ -263,6 +293,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
})
return
}
log.Errorw("get messages failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to get messages",
@@ -284,6 +317,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
}
messages := allMessages[start:]
if messages == nil {
messages = []models.Message{}
}
c.JSON(http.StatusOK, gin.H{
"messages": messages,
"total": total,

View File

@@ -8,6 +8,7 @@ import (
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/trace"
)
// SessionHandler 提供会话相关的 REST 端点。
@@ -27,6 +28,8 @@ type CreateSessionRequest struct {
// CreateSession POST /api/sessions — 创建新会话。
func (h *SessionHandler) CreateSession(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
var req CreateSessionRequest
// 请求体可选,解析失败不报错(使用默认配置)
_ = c.ShouldBindJSON(&req)
@@ -38,6 +41,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
sessionID, err := h.sessionMgr.Create(c.Request.Context(), "", cfg)
if err != nil {
log.Errorw("create session failed",
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": "INTERNAL_ERROR",
"message": "failed to create session",
@@ -48,6 +53,9 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
// 获取创建后的会话以返回 created_at
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
if err != nil {
log.Errorw("retrieve created session failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": "INTERNAL_ERROR",
"message": "failed to retrieve created session",
@@ -55,6 +63,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
return
}
log.Infow("session created",
"session_id", sess.ID)
c.JSON(http.StatusCreated, gin.H{
"session_id": sess.ID,
"created_at": sess.CreatedAt,
@@ -63,6 +73,7 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
// DestroySession DELETE /api/sessions/:id — 销毁会话。
func (h *SessionHandler) DestroySession(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
err := h.sessionMgr.Destroy(c.Request.Context(), sessionID)
@@ -74,6 +85,9 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
})
return
}
log.Errorw("destroy session failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": "INTERNAL_ERROR",
"message": "failed to destroy session",
@@ -81,6 +95,8 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
return
}
log.Infow("session destroyed",
"session_id", sessionID)
c.Status(http.StatusNoContent)
}

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

View File

@@ -15,10 +15,17 @@ var (
ErrInvalidToken = errors.New("invalid or expired token")
)
// 令牌类型常量。
const (
TokenTypeAccess = "access"
TokenTypeRefresh = "refresh"
)
// Claims JWT 声明。
type Claims struct {
UserID string `json:"user_id"`
Username string `json:"username"`
UserID string `json:"user_id"`
Username string `json:"username"`
TokenType string `json:"token_type"`
jwt.RegisteredClaims
}
@@ -45,8 +52,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
// access token
accessClaims := &Claims{
UserID: userID,
Username: username,
UserID: userID,
Username: username,
TokenType: TokenTypeAccess,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
IssuedAt: jwt.NewNumericDate(now),
@@ -62,8 +70,9 @@ func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh s
// refresh token含唯一 token_id 用于 DB 关联)
tokenID := uuid.New().String()
refreshClaims := &Claims{
UserID: userID,
Username: username,
UserID: userID,
Username: username,
TokenType: TokenTypeRefresh,
RegisteredClaims: jwt.RegisteredClaims{
ID: tokenID,
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。
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。
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。

View File

@@ -103,6 +103,44 @@ func TestValidateRefresh_ExpiredToken(t *testing.T) {
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) {
accessTTL := 15 * time.Minute
refreshTTL := 7 * 24 * time.Hour

View File

@@ -5,6 +5,8 @@ import (
"strings"
"github.com/gin-gonic/gin"
"github.com/hhs/camtalk/internal/trace"
)
// contextKey 用于在 Gin context 中存储 Claims 的 key。
@@ -17,8 +19,13 @@ const (
// 校验成功后将 user_id 和 username 写入 Gin Context。
func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
return func(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
authHeader := c.GetHeader("Authorization")
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{
"code": "INVALID_TOKEN",
"message": "missing authorization header",
@@ -29,6 +36,10 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
// 提取 Bearer token
parts := strings.SplitN(authHeader, " ", 2)
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{
"code": "INVALID_TOKEN",
"message": "invalid authorization format",
@@ -38,6 +49,11 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
claims, err := tokenMgr.ValidateAccess(parts[1])
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{
"code": "INVALID_TOKEN",
"message": "invalid or expired token",

View File

@@ -166,6 +166,9 @@ func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthRes
userID, err := s.userRepo.FindRefreshToken(ctx, tokenHash)
if err != nil {
if errors.Is(err, store.ErrRefreshTokenNotFound) {
// JWT 校验已通过但 DB 中不存在 → token 已被 rotation 删除,属于复用行为
// 吊销该用户全部 refresh token强制所有设备重新登录
_ = s.userRepo.DeleteUserRefreshTokens(ctx, claims.UserID)
return nil, ErrRefreshTokenUsed
}
return nil, err

View File

@@ -188,3 +188,48 @@ func TestLogout_Success(t *testing.T) {
})
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 失败已间接验证
}

View File

@@ -2,22 +2,23 @@ package config
import (
"fmt"
"os"
"strings"
"path/filepath"
"github.com/joho/godotenv"
"github.com/spf13/viper"
)
// Config 应用配置。
type Config struct {
App AppConfig `mapstructure:"app"`
Server ServerConfig `mapstructure:"server"`
Session SessionConfig `mapstructure:"session"`
Redis RedisConfig `mapstructure:"redis"`
AI AIConfig `mapstructure:"ai"`
Storage StorageConfig `mapstructure:"storage"`
Log LogConfig `mapstructure:"log"`
Auth AuthConfig `mapstructure:"auth"`
App AppConfig `mapstructure:"app"`
Server ServerConfig `mapstructure:"server"`
Session SessionConfig `mapstructure:"session"`
Redis RedisConfig `mapstructure:"redis"`
AI AIConfig `mapstructure:"ai"`
Storage StorageConfig `mapstructure:"storage"`
Log LogConfig `mapstructure:"log"`
Auth AuthConfig `mapstructure:"auth"`
RateLimit RateLimitConfig `mapstructure:"ratelimit"`
}
// SessionConfig 会话管理配置。
@@ -91,10 +92,23 @@ type TTSConfig struct {
}
type StorageConfig struct {
Redis RedisStorageConfig `mapstructure:"redis"`
Persistence PersistenceConfig `mapstructure:"persistence"`
// Deprecated: 使用 Redis 和 Persistence 替代
Driver string `mapstructure:"driver"`
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 {
Level string `mapstructure:"level"`
Format string `mapstructure:"format"`
@@ -107,90 +121,152 @@ type AuthConfig struct {
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 100807天
}
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml
func Load() (*Config, error) {
// RateLimitConfig 限流配置
type RateLimitConfig struct {
Enabled bool `mapstructure:"enabled"`
Query BucketConfig `mapstructure:"query"`
Login BucketConfig `mapstructure:"login"`
Register BucketConfig `mapstructure:"register"`
}
// BucketConfig 令牌桶配置。
type BucketConfig struct {
Capacity int `mapstructure:"capacity"` // 桶容量(突发上限)
Rate float64 `mapstructure:"rate"` // 每秒填充令牌数
}
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml > 默认值。
// 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.SetConfigName("config")
v.SetConfigType("yaml")
v.AddConfigPath(".")
v.AddConfigPath("./config")
v.AddConfigPath("./backend")
v.AddConfigPath("..") // 兼容从 backend/cmd/ 启动
v.AddConfigPath("../..") // 兼容从 backend/cmd/server/ 启动
v.AddConfigPath(filepath.Join(workDir, "config")) // 配置文件在 config/ 目录下
v.AddConfigPath(workDir) // 兼容旧路径
// 默认值
v.SetDefault("app.env", "dev")
v.SetDefault("app.version", "dev")
v.SetDefault("server.host", "0.0.0.0")
v.SetDefault("server.port", 8080)
v.SetDefault("server.read_timeout", 30)
v.SetDefault("server.write_timeout", 30)
v.SetDefault("server.heartbeat_interval", 30)
v.SetDefault("server.heartbeat_timeout", 60)
v.SetDefault("server.shutdown_timeout", 10)
v.SetDefault("session.ttl", 30)
v.SetDefault("session.max_history", 20)
v.SetDefault("redis.addr", "localhost:6379")
v.SetDefault("redis.db", 0)
v.SetDefault("ai.stt.provider", "deepgram")
v.SetDefault("ai.stt.model", "nova-2")
v.SetDefault("ai.stt.endpoint", "wss://api.deepgram.com/v1/listen")
v.SetDefault("ai.stt.timeout", 5)
v.SetDefault("ai.stt.http_client_timeout", 30)
v.SetDefault("ai.llm.provider", "openai")
v.SetDefault("ai.llm.model", "gpt-4o")
v.SetDefault("ai.llm.endpoint", "https://api.openai.com/v1")
v.SetDefault("ai.llm.timeout", 10)
v.SetDefault("ai.llm.http_client_timeout", 60)
v.SetDefault("ai.tts.provider", "openai")
v.SetDefault("ai.tts.model", "tts-1")
v.SetDefault("ai.tts.voice", "mimo_default")
v.SetDefault("ai.tts.speed", 1.0)
v.SetDefault("ai.tts.endpoint", "https://api.openai.com/v1")
v.SetDefault("ai.tts.timeout", 5)
v.SetDefault("ai.tts.http_client_timeout", 30)
v.SetDefault("ai.tts.output_format", "mp3")
v.SetDefault("ai.tts.sample_rate", 24000)
v.SetDefault("storage.driver", "memory")
v.SetDefault("storage.dsn", "")
v.SetDefault("log.level", "info")
v.SetDefault("log.format", "console")
v.SetDefault("auth.access_ttl", 15)
v.SetDefault("auth.refresh_ttl", 10080)
// 2. 设置默认值(与 config.yaml 保持一致,仅作为兜底)
setDefaults(v)
// 读取基础配置文件
_ = v.ReadInConfig() // 文件不存在不报错
// 根据 APP_ENV 覆盖
env := os.Getenv("APP_ENV")
if env == "" {
env = v.GetString("app.env")
// 3. 读取 config.yaml
if err := v.ReadInConfig(); err != nil {
return nil, fmt.Errorf("config: read config.yaml: %w", err)
}
// 4. 合并环境专属配置 config.{env}.yaml可选
env := v.GetString("app.env")
if env != "" {
v.SetConfigName("config." + env)
_ = v.MergeInConfig()
_ = v.MergeInConfig() // 文件不存在也不报错
}
// 环境变量覆盖
v.SetEnvPrefix("CAMTALK")
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
v.AutomaticEnv()
// 5. 显式绑定敏感信息环境变量(不用 AutomaticEnv避免隐式映射
bindEnvVars(v)
var cfg Config
if err := v.Unmarshal(&cfg); err != nil {
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 nil, fmt.Errorf("config: unmarshal: %w", err)
}
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")
}

View 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. 调用 GraphStream 模式 + 运行时 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
}

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

View 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,
}
}

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

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

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

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

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

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

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

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

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

View File

@@ -14,13 +14,11 @@ type Orchestrator interface {
// ctx 用于整体超时和中断控制。
// sessionID 用于会话管理和历史获取。
// req 包含图像和音频数据。
// history 是最近的对话历史。
// sender 用于向客户端推送消息。
ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender Sender,
) error
}

View File

@@ -1,403 +0,0 @@
package orchestrator
import (
"context"
"encoding/base64"
"strings"
"sync"
"time"
"unicode/utf8"
"github.com/hhs/camtalk/internal/ai/llm"
"github.com/hhs/camtalk/internal/ai/stt"
"github.com/hhs/camtalk/internal/ai/tts"
"github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/session"
)
// Pipeline 实现 Orchestrator 接口,管理 STT → LLM → TTS 流式管道。
type Pipeline struct {
sttService stt.Service
llmService llm.Service
ttsService tts.Service
sessionMgr session.Manager
model string // LLM 模型名,用于 llm_done 上报
ttsVoice string // TTS 音色
ttsSpeed float64 // TTS 语速
ttsOutputFmt string // TTS 输出格式
ttsSampleRate int // TTS 输出采样率
}
// New 创建 Pipeline 实例。
func New(
sttService stt.Service,
llmService llm.Service,
ttsService tts.Service,
sessionMgr session.Manager,
cfg *config.Config,
) *Pipeline {
return &Pipeline{
sttService: sttService,
llmService: llmService,
ttsService: ttsService,
sessionMgr: sessionMgr,
model: cfg.AI.LLM.Model,
ttsVoice: cfg.AI.TTS.Voice,
ttsSpeed: cfg.AI.TTS.Speed,
ttsOutputFmt: cfg.AI.TTS.OutputFormat,
ttsSampleRate: cfg.AI.TTS.SampleRate,
}
}
// ProcessQuery 实现 Orchestrator 接口。
func (p *Pipeline) ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender Sender,
) error {
log := logger.Log
startTime := time.Now()
// 解码音频数据(文本输入模式可跳过)
var audio []byte
if req.Text == "" && req.Audio != "" {
var err error
audio, err = base64.StdEncoding.DecodeString(req.Audio)
if err != nil {
log.Errorw("音频解码失败", "error", err)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "INVALID_MESSAGE",
Message: "音频数据解码失败",
})
return err
}
}
// 解码图片数据(可选)
var image []byte
if req.Image != "" {
var err error
image, err = base64.StdEncoding.DecodeString(req.Image)
if err != nil {
log.Errorw("图片解码失败", "error", err)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "INVALID_MESSAGE",
Message: "图片数据解码失败",
})
return err
}
}
// 设置活跃请求
if err := p.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil {
log.Errorw("设置活跃请求失败", "error", err)
}
defer p.sessionMgr.ClearActiveRequest(ctx, sessionID)
// 获取会话配置
sess, err := p.sessionMgr.Get(ctx, sessionID)
if err != nil {
log.Errorw("获取会话失败", "error", err)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "SESSION_NOT_FOUND",
Message: "会话不存在",
})
return err
}
// Step 1: 获取用户文本(语音识别或直接使用输入文本)
var userText string
if req.Text != "" {
// 文本输入模式:跳过 STT直接使用用户输入的文本
log.Infow("使用文本输入", "request_id", req.RequestID, "text", req.Text)
userText = req.Text
// 发送 stt_result 以保持前端消息流一致性
if err := sender.SendSTTResult(models.WsSTTResult{
Type: "stt_result",
RequestID: req.RequestID,
Text: userText,
IsFinal: true,
}); err != nil {
log.Errorw("发送 STT 结果失败", "error", err)
}
} else {
// 语音模式:执行 STT 语音识别
log.Infow("开始语音识别", "request_id", req.RequestID, "audio_bytes", len(audio))
sttResult, err := p.sttService.Recognize(ctx, audio, stt.Options{
Encoding: "pcm_s16le",
SampleRate: 16000,
Language: sess.Config.Language,
})
if err != nil {
log.Errorw("语音识别失败", "error", err, "audio_bytes", len(audio))
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "STT_ERROR",
Message: "语音识别失败: " + err.Error(),
})
return err
}
userText = sttResult
// STT 返回空文本:未识别到语音,发送结果后直接返回(不调 LLM
if strings.TrimSpace(userText) == "" {
log.Infow("语音识别结果为空", "request_id", req.RequestID)
userText = "(未识别到语音)"
if err := sender.SendSTTResult(models.WsSTTResult{
Type: "stt_result",
RequestID: req.RequestID,
Text: userText,
IsFinal: true,
}); err != nil {
log.Errorw("发送 STT 结果失败", "error", err)
}
// 发送空的 llm_done 以结束本轮处理
latency := time.Since(startTime).Milliseconds()
_ = sender.SendLLMDone(models.WsLLMDone{
Type: "llm_done",
RequestID: req.RequestID,
FullText: "",
Model: p.model,
LatencyMs: latency,
})
return nil
}
// 发送 STT 结果
if err := sender.SendSTTResult(models.WsSTTResult{
Type: "stt_result",
RequestID: req.RequestID,
Text: userText,
IsFinal: true,
}); err != nil {
log.Errorw("发送 STT 结果失败", "error", err)
}
}
// 追加用户消息到历史
p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "user",
Content: userText,
})
// Step 2+3: LLM 流式推理 + TTS 并行合成
log.Infow("开始 LLM 推理", "request_id", req.RequestID, "scenario", sess.Config.Scenario)
llmReq := llm.Request{
Image: image,
Text: userText,
History: history,
Language: sess.Config.Language,
SystemPrompt: llm.GetScenarioPrompt(sess.Config.Scenario, sess.Config.Language),
}
llmStream, err := p.llmService.ChatStream(ctx, llmReq)
if err != nil {
log.Errorw("LLM 流式推理启动失败", "error", err)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "LLM_ERROR",
Message: "LLM 推理失败",
})
return err
}
// 创建句子切分器
sentenceCh := make(chan string, 4)
splitter := NewSplitter(sentenceCh)
// 并行LLM 消费 + TTS 合成
var wg sync.WaitGroup
var fullText string
var ttsErr error
// goroutine 1: 消费 LLM token + 句子切分
var tokenUsage *llm.TokenUsage
wg.Add(1)
go func() {
defer wg.Done()
defer close(sentenceCh)
fullText, tokenUsage = p.consumeLLMStream(ctx, llmStream, req.RequestID, sender, splitter)
}()
// goroutine 2: TTS 合成(如果启用)
if sess.Config.TTSEnabled {
wg.Add(1)
go func() {
defer wg.Done()
log.Infow("开始 TTS 合成", "request_id", req.RequestID)
ttsErr = p.synthesizeTTS(ctx, sentenceCh, req.RequestID, sender)
}()
} else {
// 如果 TTS 未启用,需要消费 sentenceCh 防止阻塞
go func() {
for range sentenceCh {
}
}()
}
// 等待所有 goroutine 完成
wg.Wait()
// TTS 失败静默跳过
if ttsErr != nil {
log.Warnw("TTS 合成失败(已跳过)", "error", ttsErr)
}
// 追加助手消息到历史
p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "assistant",
Content: fullText,
})
// 发送 llm_done
latency := time.Since(startTime).Milliseconds()
done := models.WsLLMDone{
Type: "llm_done",
RequestID: req.RequestID,
FullText: fullText,
Model: p.model,
LatencyMs: latency,
}
if tokenUsage != nil {
done.TokensUsed = struct {
Prompt int `json:"prompt"`
Completion int `json:"completion"`
Total int `json:"total"`
}{
Prompt: tokenUsage.Prompt,
Completion: tokenUsage.Completion,
Total: tokenUsage.Total,
}
}
if err := sender.SendLLMDone(done); err != nil {
log.Errorw("发送 llm_done 失败", "error", err)
}
log.Infow("查询处理完成",
"request_id", req.RequestID,
"latency_ms", latency,
"text_length", utf8.RuneCountInString(fullText),
)
return nil
}
// consumeLLMStream 消费 LLM 流式输出,发送 llm_chunk 并进行句子切分。
// 返回完整文本和 token 用量。
func (p *Pipeline) consumeLLMStream(
ctx context.Context,
stream <-chan llm.Chunk,
requestID string,
sender Sender,
splitter *Splitter,
) (string, *llm.TokenUsage) {
log := logger.Log
var fullText strings.Builder
var tokenUsage *llm.TokenUsage
for chunk := range stream {
// 检查上下文是否已取消
select {
case <-ctx.Done():
log.Infow("LLM 流被中断", "request_id", requestID)
return fullText.String(), tokenUsage
default:
}
if chunk.Done {
// 流结束,记录 token 用量
if chunk.TokensUsed != nil {
tokenUsage = chunk.TokensUsed
log.Infow("LLM 用量统计",
"request_id", requestID,
"prompt_tokens", tokenUsage.Prompt,
"completion_tokens", tokenUsage.Completion,
"total_tokens", tokenUsage.Total,
)
}
break
}
// 累积全文
fullText.WriteString(chunk.Delta)
// 发送 llm_chunk
if err := sender.SendLLMChunk(models.WsLLMChunk{
Type: "llm_chunk",
RequestID: requestID,
Delta: chunk.Delta,
Role: "assistant",
}); err != nil {
log.Errorw("发送 llm_chunk 失败", "error", err)
}
// 句子切分
splitter.Feed(chunk.Delta)
}
// 刷新切分器中的剩余文本
splitter.Flush()
return fullText.String(), tokenUsage
}
// synthesizeTTS 从句子 channel 读取文本,进行 TTS 合成并发送音频。
func (p *Pipeline) synthesizeTTS(
ctx context.Context,
sentenceCh <-chan string,
requestID string,
sender Sender,
) error {
log := logger.Log
ttsStream, err := p.ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{
Voice: p.ttsVoice,
Speed: p.ttsSpeed,
OutputFmt: p.ttsOutputFmt,
SampleRate: p.ttsSampleRate,
})
if err != nil {
log.Errorw("TTS 合成启动失败", "error", err)
return err
}
// 消费 TTS 音频流
for chunk := range ttsStream {
// 检查上下文是否已取消
select {
case <-ctx.Done():
log.Infow("TTS 流被中断", "request_id", requestID)
return ctx.Err()
default:
}
// Base64 编码音频数据
audioBase64 := base64.StdEncoding.EncodeToString(chunk.Audio)
if err := sender.SendTTSAudio(models.WsTTSAudio{
Type: "tts_audio",
RequestID: requestID,
Audio: audioBase64,
MimeType: "audio/mp3",
IsLast: chunk.IsLast,
Final: chunk.Final,
}); err != nil {
log.Errorw("发送 tts_audio 失败", "error", err)
}
}
return nil
}

View File

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

View File

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

View File

@@ -0,0 +1,172 @@
package ratelimit
import (
"context"
"sync"
"time"
"github.com/hhs/camtalk/internal/config"
)
// TokenBucket 内存令牌桶,适用于单实例部署。
type TokenBucket struct {
capacity int // 桶容量
rate float64 // 每秒填充令牌数
tokens float64 // 当前令牌数
lastRefill time.Time // 上次填充时间
mu sync.Mutex
}
// newTokenBucket 创建令牌桶。
func newTokenBucket(capacity int, rate float64) *TokenBucket {
return &TokenBucket{
capacity: capacity,
rate: rate,
tokens: float64(capacity), // 初始满桶
lastRefill: time.Now(),
}
}
// allow 尝试消耗一个令牌。
func (b *TokenBucket) allow() (bool, time.Duration) {
b.mu.Lock()
defer b.mu.Unlock()
now := time.Now()
elapsed := now.Sub(b.lastRefill).Seconds()
// 补充令牌
newTokens := elapsed * b.rate
b.tokens = min(float64(b.capacity), b.tokens+newTokens)
b.lastRefill = now
// 尝试消耗一个令牌
if b.tokens >= 1 {
b.tokens -= 1
return true, 0
}
// 计算需要等待的时间
if b.rate == 0 {
// rate=0 时永远无法补充令牌
return false, 24 * time.Hour // 返回一个很大的值
}
retryAfter := time.Duration((1-b.tokens)/b.rate*1000) * time.Millisecond
return false, retryAfter
}
// MemoryLimiter 管理多个用户的令牌桶。
type MemoryLimiter struct {
buckets map[string]*TokenBucket
config config.RateLimitConfig
mu sync.RWMutex
stopOnce sync.Once
done chan struct{}
}
// NewMemoryLimiter 创建内存限流器。
func NewMemoryLimiter(cfg config.RateLimitConfig) *MemoryLimiter {
limiter := &MemoryLimiter{
buckets: make(map[string]*TokenBucket),
config: cfg,
done: make(chan struct{}),
}
// 启动后台清理 goroutine
go limiter.cleanup()
return limiter
}
// Allow 实现 Limiter 接口。
func (l *MemoryLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
bucket := l.getOrCreateBucket(key)
return bucket.allow()
}
// Stop 实现 Limiter 接口。
func (l *MemoryLimiter) Stop() {
l.stopOnce.Do(func() {
close(l.done)
})
}
// getOrCreateBucket 获取或创建令牌桶。
func (l *MemoryLimiter) getOrCreateBucket(key string) *TokenBucket {
// 先尝试读锁
l.mu.RLock()
bucket, exists := l.buckets[key]
l.mu.RUnlock()
if exists {
return bucket
}
// 需要创建新桶,升级为写锁
l.mu.Lock()
defer l.mu.Unlock()
// 双重检查(可能其他 goroutine 已创建)
bucket, exists = l.buckets[key]
if exists {
return bucket
}
// 根据 key 确定配置(简化版:假设 key 格式为 "userID:action"
cfg := l.getBucketConfig(key)
bucket = newTokenBucket(cfg.Capacity, cfg.Rate)
l.buckets[key] = bucket
return bucket
}
// getBucketConfig 根据 key 获取桶配置。
func (l *MemoryLimiter) getBucketConfig(key string) config.BucketConfig {
// 简化实现:从 key 后缀判断动作类型
// 实际使用时调用方会传递正确的 key
// 默认使用 query 配置
return l.config.Query
}
// cleanup 定期清理不活跃的桶。
func (l *MemoryLimiter) cleanup() {
ticker := time.NewTicker(10 * time.Minute)
defer ticker.Stop()
for {
select {
case <-ticker.C:
l.removeInactiveBuckets()
case <-l.done:
return
}
}
}
// removeInactiveBuckets 移除超过 10 分钟无活动的桶。
func (l *MemoryLimiter) removeInactiveBuckets() {
l.mu.Lock()
defer l.mu.Unlock()
now := time.Now()
for key, bucket := range l.buckets {
bucket.mu.Lock()
inactive := now.Sub(bucket.lastRefill) > 10*time.Minute
bucket.mu.Unlock()
if inactive {
delete(l.buckets, key)
}
}
}
// min 返回两个 float64 中的较小值。
func min(a, b float64) float64 {
if a < b {
return a
}
return b
}
// 编译期接口检查
var _ Limiter = (*MemoryLimiter)(nil)

View File

@@ -0,0 +1,203 @@
package ratelimit
import (
"context"
"sync"
"testing"
"time"
"github.com/hhs/camtalk/internal/config"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestTokenBucket_Allow_FirstRequest(t *testing.T) {
bucket := newTokenBucket(5, 0.2)
allowed, retryAfter := bucket.allow()
assert.True(t, allowed)
assert.Equal(t, time.Duration(0), retryAfter)
}
func TestTokenBucket_Allow_ConsumeUntilEmpty(t *testing.T) {
bucket := newTokenBucket(3, 0.2)
// 连续消耗 3 个令牌
for i := 0; i < 3; i++ {
allowed, _ := bucket.allow()
assert.True(t, allowed, "request %d should be allowed", i+1)
}
// 第 4 个请求应被拒绝
allowed, retryAfter := bucket.allow()
assert.False(t, allowed)
assert.Greater(t, retryAfter, time.Duration(0))
}
func TestTokenBucket_Allow_RetryAfterCorrect(t *testing.T) {
bucket := newTokenBucket(1, 1.0) // 每秒 1 个令牌
// 消耗唯一的令牌
allowed, _ := bucket.allow()
require.True(t, allowed)
// 立即再次请求应被拒绝
allowed, retryAfter := bucket.allow()
assert.False(t, allowed)
// retryAfter 应约为 1 秒(允许一定误差)
assert.InDelta(t, 1000, retryAfter.Milliseconds(), 100)
}
func TestTokenBucket_Allow_RefillAfterWait(t *testing.T) {
bucket := newTokenBucket(2, 10.0) // 每秒 10 个令牌(每 100ms 一个)
// 消耗 2 个令牌
bucket.allow()
bucket.allow()
// 等待 150ms应补充至少 1 个令牌
time.Sleep(150 * time.Millisecond)
allowed, _ := bucket.allow()
assert.True(t, allowed)
}
func TestTokenBucket_Allow_CapacityLimit(t *testing.T) {
bucket := newTokenBucket(3, 1.0)
// 等待足够长时间让桶"溢出"
time.Sleep(100 * time.Millisecond)
// 但最多只能消耗 capacity 个令牌
for i := 0; i < 3; i++ {
allowed, _ := bucket.allow()
assert.True(t, allowed, "request %d should be allowed", i+1)
}
// 第 4 个应被拒绝
allowed, _ := bucket.allow()
assert.False(t, allowed)
}
func TestTokenBucket_Allow_ConcurrentSafe(t *testing.T) {
bucket := newTokenBucket(100, 10.0)
var wg sync.WaitGroup
successCount := 0
var mu sync.Mutex
// 100 个并发请求
for i := 0; i < 100; i++ {
wg.Add(1)
go func() {
defer wg.Done()
allowed, _ := bucket.allow()
if allowed {
mu.Lock()
successCount++
mu.Unlock()
}
}()
}
wg.Wait()
// 应该正好 100 个成功(桶容量为 100
assert.Equal(t, 100, successCount)
}
func TestTokenBucket_Allow_ZeroCapacity(t *testing.T) {
bucket := newTokenBucket(0, 1.0)
allowed, retryAfter := bucket.allow()
assert.False(t, allowed)
assert.Greater(t, retryAfter, time.Duration(0))
}
func TestTokenBucket_Allow_ZeroRate(t *testing.T) {
bucket := newTokenBucket(1, 0.0)
// 第一个通过
allowed, _ := bucket.allow()
assert.True(t, allowed)
// 第二个被拒绝,且 retryAfter 应为无限大(实际上会很大)
allowed, retryAfter := bucket.allow()
assert.False(t, allowed)
// rate=0 时retryAfter 理论上无限大,实际会是一个很大的值
assert.Greater(t, retryAfter, 1*time.Hour)
}
func TestMemoryLimiter_Allow_DifferentKeys(t *testing.T) {
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
}
limiter := NewMemoryLimiter(cfg)
defer limiter.Stop()
ctx := context.Background()
// user1 消耗 2 个令牌
allowed, _ := limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
allowed, _ = limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
// user1 第 3 个被拒绝
allowed, _ = limiter.Allow(ctx, "user1:query")
assert.False(t, allowed)
// user2 应该不受影响
allowed, _ = limiter.Allow(ctx, "user2:query")
assert.True(t, allowed)
}
func TestMemoryLimiter_Cleanup(t *testing.T) {
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
}
limiter := NewMemoryLimiter(cfg)
defer limiter.Stop()
ctx := context.Background()
// 创建一个桶
limiter.Allow(ctx, "user1:query")
// 验证桶已创建
limiter.mu.RLock()
initialCount := len(limiter.buckets)
limiter.mu.RUnlock()
assert.Equal(t, 1, initialCount)
// 手动触发清理(模拟 10 分钟后)
limiter.mu.Lock()
for _, bucket := range limiter.buckets {
bucket.mu.Lock()
bucket.lastRefill = time.Now().Add(-11 * time.Minute)
bucket.mu.Unlock()
}
limiter.mu.Unlock()
limiter.removeInactiveBuckets()
// 验证桶已被清理
limiter.mu.RLock()
finalCount := len(limiter.buckets)
limiter.mu.RUnlock()
assert.Equal(t, 0, finalCount)
}
func TestMemoryLimiter_Stop(t *testing.T) {
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
}
limiter := NewMemoryLimiter(cfg)
// 多次调用 Stop 不应 panic
limiter.Stop()
limiter.Stop()
}

View File

@@ -0,0 +1,17 @@
package ratelimit
import (
"context"
"time"
)
// Limiter 速率限制器接口。
type Limiter interface {
// Allow 判断 key 是否允许执行一次操作。
// key 通常为 "userID:action" 格式。
// 返回 (allowed, retryAfter)。retryAfter 表示需要等待的时间。
Allow(ctx context.Context, key string) (bool, time.Duration)
// Stop 停止限流器,清理资源(如后台 goroutine
Stop()
}

View File

@@ -0,0 +1,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()
}
}

View File

@@ -0,0 +1,196 @@
package ratelimit
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// mockLimiter 用于测试的 mock 限流器。
type mockLimiter struct {
allowFunc func(ctx context.Context, key string) (bool, time.Duration)
}
func (m *mockLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
if m.allowFunc != nil {
return m.allowFunc(ctx, key)
}
return true, 0
}
func (m *mockLimiter) Stop() {}
// 编译期接口检查
var _ Limiter = (*mockLimiter)(nil)
func TestMiddleware_Allow(t *testing.T) {
gin.SetMode(gin.TestMode)
limiter := &mockLimiter{
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
return true, 0
},
}
router := gin.New()
router.Use(Middleware(limiter, func(c *gin.Context) string {
return "user1:test"
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
err := json.Unmarshal(w.Body.Bytes(), &resp)
require.NoError(t, err)
assert.Equal(t, "ok", resp["status"])
}
func TestMiddleware_Deny(t *testing.T) {
gin.SetMode(gin.TestMode)
limiter := &mockLimiter{
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
return false, 5 * time.Second
},
}
router := gin.New()
router.Use(Middleware(limiter, func(c *gin.Context) string {
return "user1:test"
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
// 验证返回 429
assert.Equal(t, http.StatusTooManyRequests, w.Code)
// 验证 Retry-After header
assert.Equal(t, "5", w.Header().Get("Retry-After"))
// 验证响应体
var resp map[string]interface{}
err := json.Unmarshal(w.Body.Bytes(), &resp)
require.NoError(t, err)
assert.Equal(t, "RATE_LIMITED", resp["code"])
assert.Contains(t, resp["message"], "retry after")
}
func TestMiddleware_NilLimiter(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(Middleware(nil, func(c *gin.Context) string {
return "user1:test"
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
// nil limiter 应该放行
assert.Equal(t, http.StatusOK, w.Code)
}
func TestMiddleware_EmptyKey(t *testing.T) {
gin.SetMode(gin.TestMode)
limiter := &mockLimiter{
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
// 不应该被调用
t.Error("Allow should not be called with empty key")
return false, 0
},
}
router := gin.New()
router.Use(Middleware(limiter, func(c *gin.Context) string {
return "" // 返回空 key
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
// 空 key 应该放行
assert.Equal(t, http.StatusOK, w.Code)
}
func TestMiddleware_KeyFunc(t *testing.T) {
gin.SetMode(gin.TestMode)
var capturedKey string
limiter := &mockLimiter{
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
capturedKey = key
return true, 0
},
}
router := gin.New()
router.Use(Middleware(limiter, func(c *gin.Context) string {
// 从 query 参数提取 user_id
userID := c.Query("user_id")
return userID + ":test"
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test?user_id=user123", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "user123:test", capturedKey)
}
func TestMiddleware_RetryAfterRounding(t *testing.T) {
gin.SetMode(gin.TestMode)
limiter := &mockLimiter{
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
return false, 2500 * time.Millisecond // 2.5 秒
},
}
router := gin.New()
router.Use(Middleware(limiter, func(c *gin.Context) string {
return "user1:test"
}))
router.GET("/test", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"status": "ok"})
})
req := httptest.NewRequest(http.MethodGet, "/test", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusTooManyRequests, w.Code)
// 2.5 秒向上取整为 3 秒
assert.Equal(t, "3", w.Header().Get("Retry-After"))
}

View File

@@ -0,0 +1,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] = ttlkey 过期时间,秒)
local key = KEYS[1]
local capacity = tonumber(ARGV[1])
local rate = tonumber(ARGV[2])
local now = tonumber(ARGV[3])
local ttl = tonumber(ARGV[4])
local data = redis.call('HMGET', key, 'tokens', 'last_refill')
local tokens = tonumber(data[1]) or capacity
local last_refill = tonumber(data[2]) or now
-- 计算新令牌
local elapsed = math.max(0, now - last_refill)
tokens = math.min(capacity, tokens + elapsed * rate)
local allowed = 0
local retry_after = 0
if tokens >= 1 then
tokens = tokens - 1
allowed = 1
else
if rate == 0 then
retry_after = 86400 -- 24小时
else
retry_after = (1 - tokens) / rate
end
end
-- 回写状态
redis.call('HMSET', key, 'tokens', tokens, 'last_refill', now)
redis.call('EXPIRE', key, ttl)
return {allowed, tostring(retry_after)}
`
// RedisLimiter Redis 令牌桶限流器。
type RedisLimiter struct {
client *redis.Client
config config.RateLimitConfig
script *redis.Script
}
// NewRedisLimiter 创建 Redis 限流器。
func NewRedisLimiter(client *redis.Client, cfg config.RateLimitConfig) *RedisLimiter {
return &RedisLimiter{
client: client,
config: cfg,
script: redis.NewScript(luaScript),
}
}
// Allow 实现 Limiter 接口。
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
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)

View File

@@ -0,0 +1,228 @@
package ratelimit
import (
"context"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/hhs/camtalk/internal/config"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// setupMiniRedis 创建一个内存 Redis 实例用于测试。
func setupMiniRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) {
mr, err := miniredis.Run()
require.NoError(t, err)
client := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
})
t.Cleanup(func() {
client.Close()
mr.Close()
})
return mr, client
}
func TestRedisLimiter_Allow_FirstRequest(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 5, Rate: 0.2},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
assert.Equal(t, time.Duration(0), retryAfter)
}
func TestRedisLimiter_Allow_ConsumeUntilEmpty(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 3, Rate: 0.2},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
key := "user1:query"
// 连续消耗 3 个令牌
for i := 0; i < 3; i++ {
allowed, _ := limiter.Allow(ctx, key)
assert.True(t, allowed, "request %d should be allowed", i+1)
}
// 第 4 个请求应被拒绝
allowed, retryAfter := limiter.Allow(ctx, key)
assert.False(t, allowed)
assert.Greater(t, retryAfter, time.Duration(0))
}
func TestRedisLimiter_Allow_DifferentKeys(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
// user1 消耗 2 个令牌
allowed, _ := limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
allowed, _ = limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
// user1 第 3 个被拒绝
allowed, _ = limiter.Allow(ctx, "user1:query")
assert.False(t, allowed)
// user2 应该不受影响
allowed, _ = limiter.Allow(ctx, "user2:query")
assert.True(t, allowed)
}
func TestRedisLimiter_Allow_RefillAfterWait(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 2, Rate: 10.0}, // 每秒 10 个令牌
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
key := "user1:query"
// 消耗 2 个令牌
limiter.Allow(ctx, key)
limiter.Allow(ctx, key)
// 真实等待 150msLua 脚本使用系统时间)
time.Sleep(150 * time.Millisecond)
// 应该补充了至少 1 个令牌
allowed, _ := limiter.Allow(ctx, key)
assert.True(t, allowed)
}
func TestRedisLimiter_Allow_CapacityLimit(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 3, Rate: 1.0},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
key := "user1:query"
// 真实等待让桶"溢出"
time.Sleep(100 * time.Millisecond)
// 但最多只能消耗 capacity 个令牌
for i := 0; i < 3; i++ {
allowed, _ := limiter.Allow(ctx, key)
assert.True(t, allowed, "request %d should be allowed", i+1)
}
// 第 4 个应被拒绝
allowed, _ := limiter.Allow(ctx, key)
assert.False(t, allowed)
}
func TestRedisLimiter_Allow_ZeroRate(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 1, Rate: 0.0},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
key := "user1:query"
// 第一个通过
allowed, _ := limiter.Allow(ctx, key)
assert.True(t, allowed)
// 第二个被拒绝retryAfter 应该很大
allowed, retryAfter := limiter.Allow(ctx, key)
assert.False(t, allowed)
assert.Greater(t, retryAfter, 1*time.Hour)
}
func TestRedisLimiter_Allow_KeyTTL(t *testing.T) {
mr, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 5, Rate: 1.0},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
key := "user1:query"
// 第一次请求
limiter.Allow(ctx, key)
// 验证 key 已设置 TTL
ttl := mr.TTL(key)
assert.Greater(t, ttl, time.Duration(0))
assert.LessOrEqual(t, ttl, 600*time.Second)
}
func TestRedisLimiter_Allow_FailOpen(t *testing.T) {
mr, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
}
limiter := NewRedisLimiter(client, cfg)
ctx := context.Background()
// 关闭 Redis 模拟故障
mr.Close()
// 应该 fail-open允许请求
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
assert.True(t, allowed)
assert.Equal(t, time.Duration(0), retryAfter)
}
func TestRedisLimiter_Stop(t *testing.T) {
_, client := setupMiniRedis(t)
cfg := config.RateLimitConfig{
Enabled: true,
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
}
limiter := NewRedisLimiter(client, cfg)
// Stop 应该不会 panic即使多次调用
limiter.Stop()
limiter.Stop()
}
func TestFormatKey(t *testing.T) {
key := FormatKey("user123", "query")
assert.Equal(t, "ratelimit:user123:query", key)
}

View File

@@ -18,6 +18,7 @@ type ConversationSummary struct {
Title string `json:"title"`
LastMessage string `json:"last_message"`
MessageCount int `json:"message_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}

View File

@@ -142,11 +142,11 @@ func (m *MemoryManager) Create(ctx context.Context, userID string, config models
}
m.mu.Unlock()
// Write-Through异步写 PG
// Write-Through异步写 PG(使用 Background context避免 HTTP 请求结束后 context 被取消)
if m.sessRepo != nil {
go func() {
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,
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
}); err != nil {
@@ -204,11 +204,11 @@ func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patc
cfg := entry.session.Config
m.mu.Unlock()
// Write-Through异步更新 PG
// Write-Through异步更新 PG(使用 Background context
if m.sessRepo != nil {
go func() {
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)
}
}()
@@ -233,10 +233,10 @@ func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title
entry.lastActive = time.Now()
m.mu.Unlock()
// Write-Through异步更新 PG
// Write-Through异步更新 PG(使用 Background context
if m.sessRepo != nil {
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)
}
}()
@@ -271,6 +271,7 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
list = append(list, ConversationSummary{
ID: rec.ID,
Title: rec.Title,
CreatedAt: rec.CreatedAt,
UpdatedAt: rec.UpdatedAt,
})
sessionIDs = append(sessionIDs, rec.ID)
@@ -320,6 +321,7 @@ func (m *MemoryManager) listByUserFromMemory(ctx context.Context, userID string,
summary := ConversationSummary{
ID: entry.session.ID,
Title: entry.session.Title,
CreatedAt: entry.session.CreatedAt,
UpdatedAt: entry.lastActive,
}
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)
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
titleUpdated := false
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
entry.session.Title = generateTitle(msg.Content)
titleUpdated = true
}
// 超过上限时裁剪,保留最新的 maxHistory 条
@@ -406,13 +410,31 @@ func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg m
now := time.Now()
entry.lastActive = now
entry.session.UpdatedAt = now
// 复制标题(释放锁后安全使用)
persistTitle := entry.session.Title
m.mu.Unlock()
// Write-Through异步写冷存储,不阻塞调用方
// Write-Through消息同步写入 PostgreSQL保证调用顺序 = 插入顺序,
// 避免用户消息和 AI 消息的异步 goroutine 执行顺序不确定导致排序错乱)
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() {
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
if titleUpdated {
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)
m.mu.Unlock()
// Write-Through异步删除 PG
// Write-Through异步删除 PG(使用 Background context
if m.sessRepo != nil {
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)
}
}()

View File

@@ -10,8 +10,9 @@ import (
"github.com/google/uuid"
"github.com/redis/go-redis/v9"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/trace"
"github.com/hhs/camtalk/internal/util"
)
// 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}
}
// 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 histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
func userSessKey(id string) string { return fmt.Sprintf("user:%s:sessions", id) }
// Create 创建新会话。userID 为空表示匿名会话。
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()
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)
}
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
}
@@ -86,8 +98,11 @@ const placeholderHistoryMark = "__placeholder__"
// Get 获取会话。
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
log := trace.FromContext(ctx)
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
if err != nil {
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
return nil, fmt.Errorf("redis get session: %w", err)
}
if len(vals) == 0 {
@@ -105,6 +120,7 @@ func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Sessi
sess.Config.DetailLevel = vals["config.detail_level"]
sess.Config.Language = vals["config.language"]
log.Debugw("redis session retrieved", "session_id", sessionID)
return sess, nil
}
@@ -140,7 +156,9 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
// 刷新 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
}
@@ -160,7 +178,9 @@ func (m *RedisManager) UpdateTitle(ctx context.Context, sessionID string, title
}
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
}
@@ -275,7 +295,11 @@ func (m *RedisManager) GetHistory(ctx context.Context, sessionID string, limit i
}
var msg models.Message
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
}
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)
}
logger.Log.Debugw("redis session destroyed", "session", sessionID)
log := trace.FromContext(ctx)
log.Debugw("redis session destroyed", "session_id", sessionID)
return nil
}

View 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内存→ L2Redis→ L3PostgreSQL
//
// 读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()
}

View 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
}
// 写 RedisSET + 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先查 Redismiss 时查 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
}
// 回填 RedisSET + SADDTTL 使用保守默认值
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)
}

View File

@@ -8,6 +8,7 @@ import (
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/trace"
)
// 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 {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx,
`INSERT INTO messages (session_id, role, content, tokens_used) VALUES ($1, $2, $3, $4)`,
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) {
log := trace.FromContext(ctx)
if limit <= 0 {
limit = 50
}
@@ -56,6 +67,7 @@ func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string,
)
}
if err != nil {
log.Errorw("get messages failed", "session_id", sessionID, "error", 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]
}
log.Debugw("messages retrieved", "session_id", sessionID, "count", len(rows))
return rows, nil
}
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...)
if err != nil {
log.Errorw("query messages failed", "error", err)
return nil, err
}
defer pgxRows.Close()
var messages []StoredMessage
messages := make([]StoredMessage, 0)
for pgxRows.Next() {
var m StoredMessage
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
}
messages = append(messages, m)
}
if err := pgxRows.Err(); err != nil {
log.Errorw("iterate message rows failed", "error", err)
return nil, err
}
return messages, nil
}
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
log := trace.FromContext(ctx)
var m StoredMessage
err := r.pool.QueryRow(ctx,
`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
}
if err != nil {
log.Errorw("get last message failed", "session_id", sessionID, "error", err)
return nil, err
}
log.Debugw("last message retrieved", "session_id", sessionID, "message_id", m.ID)
return &m, nil
}
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
log := trace.FromContext(ctx)
var count int
err := r.pool.QueryRow(ctx,
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
sessionID,
).Scan(&count)
if err != nil {
log.Errorw("get message count failed", "session_id", sessionID, "error", err)
return 0, err
}
log.Debugw("message count retrieved", "session_id", sessionID, "count", count)
return count, nil
}
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
log := trace.FromContext(ctx)
if len(sessionIDs) == 0 {
return map[string]SessionMessageStats{}, nil
}
@@ -143,6 +173,7 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
sessionIDs,
)
if err != nil {
log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err)
return nil, err
}
defer rows.Close()
@@ -152,12 +183,16 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
var sid string
var stats SessionMessageStats
if err := rows.Scan(&sid, &stats.MessageCount, &stats.LastMessage); err != nil {
log.Errorw("scan message stats row failed", "error", err)
return nil, err
}
result[sid] = stats
}
if err := rows.Err(); err != nil {
log.Errorw("iterate message stats rows failed", "error", err)
return nil, err
}
log.Debugw("session message stats retrieved", "session_count", len(sessionIDs), "result_count", len(result))
return result, nil
}

View File

@@ -6,6 +6,8 @@ import (
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/trace"
)
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
@@ -19,6 +21,8 @@ func NewPgSessionRepository(pool *pgxpool.Pool) *PgSessionRepository {
}
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx,
`INSERT INTO sessions (id, user_id, title, config, created_at, updated_at)
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`,
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) {
log := trace.FromContext(ctx)
var s SessionRecord
err := r.pool.QueryRow(ctx,
`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
}
if err != nil {
log.Errorw("find session failed", "session_id", id, "error", err)
return nil, err
}
log.Debugw("session found", "session_id", id)
return &s, nil
}
func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error) {
log := trace.FromContext(ctx)
if page <= 0 {
page = 1
}
@@ -60,6 +77,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
if err := r.pool.QueryRow(ctx,
`SELECT COUNT(*) FROM sessions WHERE user_id = $1`, userID,
).Scan(&total); err != nil {
log.Errorw("count user sessions failed", "user_id", userID, "error", err)
return nil, 0, err
}
@@ -73,6 +91,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
userID, size, offset,
)
if err != nil {
log.Errorw("find user sessions failed", "user_id", userID, "error", err)
return nil, 0, err
}
defer rows.Close()
@@ -81,66 +100,90 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
for rows.Next() {
var s SessionRecord
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
}
list = append(list, s)
}
if err := rows.Err(); err != nil {
log.Errorw("iterate session rows failed", "user_id", userID, "error", err)
return nil, 0, err
}
log.Debugw("user sessions found", "user_id", userID, "count", len(list), "total", total)
return list, total, nil
}
func (r *PgSessionRepository) UpdateTitle(ctx context.Context, id string, title string) error {
log := trace.FromContext(ctx)
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET title = $2, updated_at = NOW() WHERE id = $1`,
id, title,
)
if err != nil {
log.Errorw("update session title failed", "session_id", id, "error", err)
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
log.Debugw("session title updated", "session_id", id)
return nil
}
func (r *PgSessionRepository) UpdateConfig(ctx context.Context, id string, configJSON []byte) error {
log := trace.FromContext(ctx)
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET config = $2, updated_at = NOW() WHERE id = $1`,
id, configJSON,
)
if err != nil {
log.Errorw("update session config failed", "session_id", id, "error", err)
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
log.Debugw("session config updated", "session_id", id)
return nil
}
func (r *PgSessionRepository) Touch(ctx context.Context, id string) error {
log := trace.FromContext(ctx)
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET updated_at = NOW() WHERE id = $1`, id,
)
if err != nil {
log.Errorw("touch session failed", "session_id", id, "error", err)
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
log.Debugw("session touched", "session_id", id)
return nil
}
func (r *PgSessionRepository) Delete(ctx context.Context, id string) error {
log := trace.FromContext(ctx)
tag, err := r.pool.Exec(ctx,
`DELETE FROM sessions WHERE id = $1`, id,
)
if err != nil {
log.Errorw("delete session failed", "session_id", id, "error", err)
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
log.Debugw("session deleted", "session_id", id)
return nil
}

View File

@@ -7,6 +7,8 @@ import (
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/trace"
)
// 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) {
log := trace.FromContext(ctx)
var id string
err := r.pool.QueryRow(ctx,
`INSERT INTO users (username, password_hash) VALUES ($1, $2) RETURNING id`,
username, passwordHash,
).Scan(&id)
if err != nil {
log.Errorw("create user failed", "username", username, "error", err)
return "", err
}
log.Debugw("user created", "user_id", id, "username", username)
return id, nil
}
func (r *PgUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
log := trace.FromContext(ctx)
var u User
err := r.pool.QueryRow(ctx,
`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
}
if err != nil {
log.Errorw("find user by username failed", "username", username, "error", err)
return nil, err
}
log.Debugw("user found by username", "user_id", u.ID, "username", username)
return &u, nil
}
func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
log := trace.FromContext(ctx)
var u User
err := r.pool.QueryRow(ctx,
`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
}
if err != nil {
log.Errorw("find user by id failed", "user_id", id, "error", err)
return nil, err
}
log.Debugw("user found by id", "user_id", id)
return &u, nil
}
func (r *PgUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx,
`INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`,
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) {
log := trace.FromContext(ctx)
var userID string
err := r.pool.QueryRow(ctx,
`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
}
if err != nil {
log.Errorw("find refresh token failed", "error", err)
return "", err
}
log.Debugw("refresh token found", "user_id", userID)
return userID, nil
}
func (r *PgUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx,
`DELETE FROM refresh_tokens WHERE token_hash = $1`,
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 {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx,
`DELETE FROM refresh_tokens WHERE user_id = $1`,
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
}

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

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

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

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

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

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

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

View File

@@ -0,0 +1,9 @@
package util
// Truncate 截断字符串到指定长度,超出部分用 "..." 替换
func Truncate(s string, maxLen int) string {
if len(s) <= maxLen {
return s
}
return s[:maxLen] + "..."
}

View File

@@ -3,6 +3,7 @@ package ws
import (
"context"
"encoding/json"
"fmt"
"net/http"
"sync"
"time"
@@ -10,13 +11,16 @@ import (
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/hhs/camtalk/internal/ai/llm"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/errors"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/orchestrator"
"github.com/hhs/camtalk/internal/ratelimit"
"github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store"
"github.com/hhs/camtalk/internal/trace"
)
// newUpgrader 根据配置创建 WebSocket upgrader。
@@ -40,12 +44,12 @@ func newUpgrader(cfg *config.Config) websocket.Upgrader {
// Client 代表一个 WebSocket 客户端连接。
type Client struct {
conn *websocket.Conn
sessionID string
sessionMgr session.Manager
orchestrator orchestrator.Orchestrator
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
mu sync.Mutex
conn *websocket.Conn
sessionID string
sessionMgr session.Manager
orchestrator orchestrator.Orchestrator
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
mu sync.Mutex
}
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
@@ -92,21 +96,19 @@ func (w *WSClient) SendError(err models.WsError) error {
}
// 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)
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
version := cfg.App.Version
maxHistory := cfg.Session.MaxHistory
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,
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 错误) ---
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)
if err != nil {
logger.Log.Errorw("websocket upgrade failed", "error", err)
log := trace.FromContext(ctx)
log.Errorw("websocket upgrade failed", "error", err)
return
}
defer conn.Close()
@@ -143,13 +156,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
var sessionID string
if 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 {
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
if err != nil {
logger.Log.Errorw("create session failed", "error", err)
log := trace.FromContext(ctx)
log.Errorw("create session failed", "error", err)
return
}
ctx = trace.WithSessionID(ctx, sessionID)
}
client := &Client{
@@ -166,7 +183,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
SessionID: sessionID,
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()
@@ -184,7 +202,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
select {
case <-ticker.C:
if time.Since(lastPong) > heartbeatTimeout {
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
log := trace.FromContext(ctx)
log.Warnw("heartbeat timeout")
conn.Close()
return
}
@@ -199,7 +218,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_, message, err := conn.ReadMessage()
if err != nil {
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
}
@@ -224,23 +244,36 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
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
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 {
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
ctx, cancel := context.WithCancel(context.Background())
processCtx, cancel := context.WithCancel(queryCtx)
client.mu.Lock()
client.cancelFuncs[msg.RequestID] = cancel
client.mu.Unlock()
@@ -260,8 +293,9 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
}()
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, history, sender); err != nil {
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil {
log := trace.FromContext(processCtx)
log.Errorw("process query failed", "error", err)
}
}()
@@ -282,10 +316,66 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
errors.SendWSError(client, errors.CodeInternalError, "", err)
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":
logger.Log.Infow("interrupt received", "session", sessionID)
log := trace.FromContext(ctx)
log.Infow("interrupt received")
// 获取活跃请求 ID 并取消
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
@@ -313,12 +403,14 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
// 取消所有活跃请求
client.mu.Lock()
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()
}
client.cancelFuncs = make(map[string]context.CancelFunc)
client.mu.Unlock()
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
logger.Log.Infow("client disconnected", "session", sessionID)
log = trace.FromContext(ctx)
log.Infow("client disconnected")
}

View File

@@ -8,11 +8,11 @@ import (
"testing"
"time"
"context"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"context"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/config"
@@ -48,7 +48,6 @@ func (m *MockOrchestrator) ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender orchestrator.Sender,
) error {
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},
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)
@@ -222,9 +221,9 @@ func TestWS_QueryFullFlow(t *testing.T) {
imageB64 := base64.StdEncoding.EncodeToString([]byte("fake-image-data"))
mock := &MockOrchestrator{
STTResult: "你好,世界",
LLMDeltas: []string{"你好", ",世界!"},
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
STTResult: "你好,世界",
LLMDeltas: []string{"你好", ",世界!"},
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
}
srv, wsURL := setupTestServer(t, mock)
@@ -333,7 +332,7 @@ func TestWS_UnknownMessageType(t *testing.T) {
err := conn.WriteJSON(map[string]string{"type": "unknown_type"})
require.NoError(t, err)
errMsg := readJSON(t, conn)
errMsg := readJSON(t, conn)
assert.Equal(t, "error", errMsg["type"])
assert.Equal(t, "INVALID_MESSAGE", errMsg["code"])
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},
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)
return srv, tokenMgr, sessionMgr
@@ -643,7 +642,7 @@ func TestWS_AuthExpiredToken(t *testing.T) {
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
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)
defer srv.Close()

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

View File

@@ -0,0 +1,36 @@
-- 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(10) DEFAULT '',
description VARCHAR(100) NOT NULL,
prompt TEXT NOT NULL,
greeting VARCHAR(200),
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 (char_length(description) >= 5 AND char_length(description) <= 100),
CONSTRAINT check_prompt_length CHECK (char_length(prompt) >= 50 AND char_length(prompt) <= 2000)
);
-- 为用户 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 '用户自建情景表';
COMMENT ON COLUMN user_scenarios.id IS '情景唯一标识';
COMMENT ON COLUMN user_scenarios.user_id IS '所属用户 ID外键关联 users 表';
COMMENT ON COLUMN user_scenarios.name IS '情景名称,如"创意写作导师"';
COMMENT ON COLUMN user_scenarios.icon IS 'Emoji 图标,如"🎨"';
COMMENT ON COLUMN user_scenarios.description IS '简短描述,显示在情景卡片上';
COMMENT ON COLUMN user_scenarios.prompt IS '角色 System Prompt定义 AI 行为';
COMMENT ON COLUMN user_scenarios.greeting IS '首句引导,可选';
COMMENT ON COLUMN user_scenarios.language IS '默认语言,如 zh-CN、en-US';

View File

@@ -4,30 +4,52 @@ set -euo pipefail
PROJECT_DIR="$(cd "$(dirname "$0")" && pwd)"
cd "$PROJECT_DIR"
# .env 固定路径act_runner 容器已挂载 /opt/camtalk
ENV_FILE="/opt/camtalk/.env"
# 颜色输出
GREEN='\033[0;32m'
NC='\033[0m'
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() {
info "构建 Docker 镜像..."
# 启用 BuildKit 加速构建
DOCKER_BUILDKIT=1 docker compose build --parallel
DOCKER_BUILDKIT=1 $DC build --parallel
info "构建完成"
}
cmd_up() {
info "启动服务..."
docker compose up -d
$DC up -d
info "服务已启动"
info "前端: http://8.161.227.145:9000"
info "健康检查: http://8.161.227.145:9000/api/health"
PUBLIC_IP=$(curl -s --connect-timeout 3 https://ifconfig.me 2>/dev/null || \
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() {
info "停止服务..."
docker compose down
$DC down
info "服务已停止"
}
@@ -38,11 +60,11 @@ cmd_restart() {
}
cmd_logs() {
docker compose logs -f "${@}"
$DC logs -f "${@}"
}
cmd_status() {
docker compose ps
$DC ps
}
usage() {
@@ -58,6 +80,8 @@ CamTalk 部署脚本
restart 重启服务
logs 查看日志(可加服务名,如: $0 logs backend
status 查看服务状态
.env 路径: $ENV_FILE
EOF
}
@@ -69,4 +93,4 @@ case "${1:-}" in
logs) shift; cmd_logs "$@" ;;
status) cmd_status ;;
*) usage; exit 1 ;;
esac
esac

View File

@@ -17,25 +17,31 @@ services:
context: ./backend
dockerfile: Dockerfile
container_name: camtalk-backend
env_file:
- /opt/camtalk/.env
environment:
- APP_ENV=production
- CAMTALK_STORAGE_DRIVER=postgres
- CAMTALK_STORAGE_DSN=postgres://camtalk:camtalk123@postgres:5432/camtalk?sslmode=disable
- CAMTALK_AUTH_JWT_SECRET=78uWBBAF8XEQEotKDlrnlnd4y8i4WN3E4zXmNmC8BYQ=
# 运行环境(强制生产环境)
- APP_ENV=prod
# 三级存储配置(敏感信息通过 env_file 注入)
- 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:
postgres:
condition: service_healthy
redis:
condition: service_healthy
networks:
- camtalk-net
restart: unless-stopped
postgres:
# 轩辕镜像加速,避免 Docker Hub 拉取超时
image: docker.m.daocloud.io/library/postgres:15-alpine
container_name: camtalk-postgres
env_file:
- /opt/camtalk/.env
environment:
POSTGRES_USER: camtalk
POSTGRES_PASSWORD: camtalk123
POSTGRES_DB: camtalk
volumes:
- pgdata:/var/lib/postgresql/data
@@ -43,7 +49,30 @@ services:
networks:
- camtalk-net
healthcheck:
test: ["CMD-SHELL", "pg_isready -U camtalk -d camtalk"]
test: ["CMD-SHELL", "pg_isready -U $$POSTGRES_USER -d camtalk"]
interval: 5s
timeout: 3s
retries: 10
restart: unless-stopped
redis:
image: docker.m.daocloud.io/library/redis:7-alpine
container_name: camtalk-redis
env_file:
- /opt/camtalk/.env
command: >
sh -c '
if [ -n "$$CAMTALK_REDIS_PASSWORD" ]; then
exec redis-server --appendonly yes --requirepass "$$CAMTALK_REDIS_PASSWORD"
else
exec redis-server --appendonly yes
fi'
volumes:
- redisdata:/data
networks:
- camtalk-net
healthcheck:
test: ["CMD-SHELL", "if [ -n \"$$CAMTALK_REDIS_PASSWORD\" ]; then redis-cli -a \"$$CAMTALK_REDIS_PASSWORD\" ping; else redis-cli ping; fi"]
interval: 5s
timeout: 3s
retries: 10
@@ -51,6 +80,7 @@ services:
volumes:
pgdata:
redisdata:
networks:
camtalk-net:

369
docs/01-架构设计.md Normal file
View 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 API1API Key 安全性2统一的速率限制和成本管控3多模型路由逻辑集中在一处便于维护。
## 核心交互流程
一次完整的"用户提问 → AI 回答"流程:
```mermaid
sequenceDiagram
participant B as 浏览器
participant G as Go 网关Eino Graph
participant S as STT
participant L as LLMChatModel
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无需硬编码端口。

View File

@@ -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
View 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重构为纯接口契约规范移除实现细节

View File

@@ -1,254 +0,0 @@
# 系统架构
## 概述
三层架构:**前端做轻量预处理,后端做智能编排,云端 AI 服务按需调用**。在保证交互体验的同时控制成本。
## 三层架构
| 层级 | 职责 | 关键约束 |
|------|------|---------|
| **客户端(浏览器)** | 媒体采集、边缘预处理、UI 渲染 | 浏览器资源有限,模型需轻量 |
| **Go 网关** | 会话管理、AI 服务编排、流式管道 | 高并发、低延迟、状态管理 |
| **AI 服务** | LLM 推理、语音识别、语音合成 | 按量计费,需控制调用频率 |
> 为什么要单独加一层 Go 网关,而不是让前端直连 AI API1API 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规划中 / MemoryMVP 默认) | 高速 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 | 维护用户会话状态、对话历史 | MemoryMVP 默认)/ 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 GatewayREST API
└── /ws → Go GatewayWebSocket
├── 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`

View File

@@ -4,15 +4,29 @@
本文档记录项目中各项技术的**选型过程、替代方案对比和决策理由**。技术选型没有"绝对正确",只有"更适合"。
**定位**:持久化部分是拓展选型,不阻塞 MVPMVP 用内存存储即可)。前端边缘处理部分是 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 服务栈
│ ├── STT: Deepgram默认 / MiMo ASR
│ ├── LLM: GPT-4o默认 / 通义千问等 OpenAI 兼容模型
│ └── TTS: OpenAI TTS默认 / MiMo TTS
├── 持久化层 → 数据库选型: PostgreSQL规划中MVP 阶段使用内存存储)
│ ├── STT: MiMo ASR默认 / Deepgram
│ ├── LLM: DashScope qwen3-vl-plus默认 / GPT-4o 等 OpenAI 兼容模型
│ └── TTS: MiMo TTS默认 / OpenAI TTS
├── 持久化层
│ ├── 数据库: PostgreSQLpgx/v5手写 SQL
│ ├── 迁移: 嵌入式 SQL 文件,自动执行
│ └── 存储模式: 三级存储 TieredManagerL1 Memory → L2 Redis → L3 PostgreSQL
├── 认证与用户系统
│ ├── 认证方案: JWT (HS256), access 15min + refresh 7day
│ ├── JWT 库: golang-jwt/jwt/v5
@@ -20,48 +34,116 @@
│ ├── 数据库驱动: pgx/v5手写 SQL不用 ORM
│ └── 前端 Token 存储: localStorage
└── 前端边缘处理层
├── 边缘推理: ONNX Runtime Web规划中MVP 使用 Canvas 像素比较
├── 关键帧检测: Canvas 像素比较160x120 降采样
├── 语音检测: @ricky0123/vad-web
└── 媒体采集: 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-4oOpenAI、Claude SonnetAnthropic给照片+问题能"看懂"照片再回答 |
| **STT** | Speech-to-Text语音转文字。流式识别延迟可低于 500ms |
| **TTS** | Text-to-Speech文字转语音。支持流式——边生成边读不必等全部生成完 |
### 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 | 按分钟计费 | 准确率高,支持多语言 |
| FunASR | <500ms | 自部署免费 | 阿里开源,中文优化 |
当前默认使用 Deepgram nova-2,可通过 `ai.stt.provider` 配置切换到 MiMo ASR
当前默认使用 MiMo ASRmimo-v2.5-asr,可通过 `ai.stt.provider` 配置切换到 Deepgram
### LLM多模态大模型
| 方案 | 成本 | 特点 |
|------|------|------|
| **GPT-4o**(默认) | $2.5/1M tokens | 视觉理解能力强API 成熟,流式推理 |
| 通义千问 qwen3-vl-plus | 按量计费 | 阿里云,通过 OpenAI 兼容接口调用 |
| **DashScope qwen3-vl-plus**(默认) | 按量计费 | 阿里云,通过 OpenAI 兼容接口调用,视觉理解能力强 |
| GPT-4o | $2.5/1M tokens | OpenAIAPI 成熟,流式推理 |
| 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语音合成
| 方案 | 成本 | 特点 |
|------|------|------|
| **OpenAI TTS**(默认) | $15/1M 字符 | 音质自然,支持流式,默认模型 tts-1语音 alloy |
| MiMo TTS小米 | 按量计费 | 国产替代,通过配置切换 |
| **MiMo TTS**(默认) | 按量计费 | 国产替代,通过配置切换,模型 mimo-v2.5-tts |
| OpenAI TTS | $15/1M 字符 | 音质自然,支持流式,默认模型 tts-1语音 alloy |
当前默认使用 OpenAI TTStts-1, alloy),可通过 `ai.tts.provider` 配置切换。
当前默认使用 MiMo TTSmimo-v2.5-tts),可通过 `ai.tts.provider` 配置切换到 OpenAI TTS
---
## 、持久化层选型规划中MVP 阶段使用内存存储)
## 、持久化层选型
### 关键术语
| 名词 | 解释 |
|------|------|
| **PostgreSQL** | 关系型数据库,支持 JSONBJSON 二进制格式可建索引、窗口函数、CTE 等高级特性 |
| **Redis** | 内存 KV 数据库,数据放在内存里,读写微秒级。支持 TTL 过期自动清理 |
| **MVCC** | Multi-Version Concurrency Control多版本并发控制PostgreSQL 用此实现高并发读写而不阻塞 |
### 数据特征分析
@@ -146,17 +228,19 @@ ORDER BY created_at DESC
LIMIT 20;
```
### 冷热分离架构
### 冷热分离架构(三级存储)
```
Go Gateway
├── 写入路径 → Redis实时会话状态
│ → PostgreSQL对话历史 + 用量
└── 读取路径 → Redis当前上下文
→ PostgreSQL历史记录
Go Gateway (TieredManager)
├── L1: Memory进程内缓存微秒级
├── L2: Redis分布式缓存毫秒级
└── L3: 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-SHA256JWT 对称签名算法,用同一密钥签名和验证 |
| **bcrypt** | 密码哈希算法,自适应 cost factor抗暴力破解 |
| **pgx** | Go 生态性能最优的 PostgreSQL 驱动,原生协议实现,内置连接池 pgxpool |
### 总览

File diff suppressed because it is too large Load Diff

View File

@@ -14,6 +14,8 @@
麦克风 → VAD → STT → LLM → TTS → 扬声器
```
> 后端 AI 编排基于 Eino Graph 声明式 DAG 实现:`START → STT → History → ChatModel → Msg2Str → Splitter → TTS → Done → END`。详见 [11-Eino框架技术文档](11-Eino框架技术文档.md)。
## 环节一VAD语音活动检测
从持续音频流中检测"人什么时候在说话",避免将环境噪音当作有效输入。**浏览器端完成**,节省 ~70% 带宽。
@@ -41,12 +43,12 @@ vad.start();
| 方案 | 延迟 | 成本 | 特点 |
|------|------|------|------|
| **MiMo ASR**(默认) | ~1s | 按量计费 | 国产替代,兼容 OpenAI 格式HTTP 非流式 |
| **Deepgram** | <500ms | 按分钟计费 | 流式识别,延迟极低 |
| Whisper API | 1-3s | 按分钟计费 | 准确率高,支持多语言 |
| **Deepgram**(默认) | <500ms | 按分钟计费 | 流式识别,延迟极低 |
| **MiMo ASR**(小米) | ~1s | 按量计费 | 国产替代,兼容 OpenAI 格式HTTP 非流式 |
| 浏览器原生 | ~1s | 免费 | 中文效果一般 |
当前实现为**一次性语音识别**(非流式):前端 VAD 检测到用户说完后,将完整音频片段发送到后端,后端调用 `stt.Recognize()` 一次性返回识别结果。流式 STT 为未来优化方向。
当前实现为**一次性语音识别**(非流式):前端 VAD 检测到用户说完后,将完整音频片段发送到后端,后端通过 Eino Graph 的 STT Lambda 节点调用 `stt.Recognize()` 一次性返回识别结果。流式 STT 为未来优化方向。
音频编码格式:前端 `audio.ts` 将 Float32Array 转为 Int16 PCM16kHz, 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**(小米):国产替代,通过配置切换
- **Edge TTS**规划中):微软免费方案,音质不错,延迟略高
- **MiMo TTS**(默认):国产替代,模型 mimo-v2.5-tts通过配置切换
- **OpenAI TTS**:音质好,延迟中等,按字符计费,模型 tts-1
- **Edge TTS**待实现):微软免费方案,音质不错,延迟略高
## 延迟优化要点

View File

@@ -26,39 +26,32 @@
| 用户触发 | 高 | 低 | 只在用户提问时拍照 |
| 本地预筛选 | 中 | 高 | 用轻量模型判断"是否值得问 LLM" |
```typescript
// 混合策略:定时低频 + 事件高频sampling.ts
const IDLE_INTERVAL = 5000; // 空闲 5 秒一帧
const ACTIVE_INTERVAL = 1000; // 用户说话时 1 秒一帧
// SamplingController 根据 VAD 状态切换采样间隔
// detail_level 通过 session config 静态配置,不随说话状态动态变化
```
**实现细节**:参见 `frontend/src/lib/sampling.ts` 中的 SamplingController根据 VAD 状态在空闲模式5s/帧和活跃模式1s/帧)之间切换。
## 策略二:端云协同——把计算推到边缘
不是所有计算都需要上云。可前置到客户端的计算:
- **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 触发
- **敏感内容过滤**规划中NSFW 检测前置,避免无效 API 调用
- **敏感内容过滤**待实现NSFW 检测前置,避免无效 API 调用
## 策略三:模型分级——用对模型做对事(规划中
## 策略三:模型分级——用对模型做对事(待实现
不是每个问题都需要最贵的模型:
```
用户提问 → 问题复杂度判断
├── 简单识别 → GPT-4o-mini ($0.15/1M tokens)
├── 深度分析 → GPT-4o ($2.5/1M tokens)
└── 代码/推理 → o1 ($15/1M tokens)
├── 简单识别 → 轻量模型(如 qwen-turbo
├── 深度分析 → qwen3-vl-plus默认按量计费
└── 代码/推理 → 更强模型(如 o1
```
> 当前 MVP 阶段使用单一模型(默认 GPT-4o模型分级路由为未来优化方向。通过配置 `ai.llm.model` 可手动切换模型
> 当前 MVP 阶段使用单一模型(默认 DashScope qwen3-vl-plus模型分级路由为未来优化方向。LLM 通过 Eino ChatModel 接入,支持任何 OpenAI 兼容接口
## 策略四:缓存与复用(规划中
## 策略四:缓存与复用(待实现
- **语义缓存**规划中):相似问题直接返回缓存结果(如反复问"这是什么"
- **语义缓存**待实现):相似问题直接返回缓存结果(如反复问"这是什么"
- **上下文复用**:连续对话中,未变化的图像不必重复发送(已通过重复画面过滤实现)
- **对话历史裁剪**:前端按 `MAX_HISTORY_ROUNDS = 10` 裁剪,后端按 `defaultHistorySize = 20` 裁剪,限制每轮的固定 token 开销

View 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 handlerLLM token 推送)
├── graph.go # Graph 构建与编译
├── adapter.go # EinoOrchestratorOrchestrator 接口适配器)
├── 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
View 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-USAI 用英语回复 |
| **持久化** | 切换情景后刷新页面 | 情景配置保持,首句仍在历史中 |
| **多情景** | 依次测试所有情景 | 每个情景 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` — 日文翻译

View File

@@ -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-4oOpenAI/ Claude SonnetAnthropic给照片+问题能"看懂"照片再回答。 |
| **STT** | 语音转文字 | Speech-to-Text。Deepgram 流式识别延迟 <500ms。备选 FunASR阿里开源可自部署。 |
| **TTS** | 文字转语音 | Text-to-Speech。OpenAI TTS 音质接近真人。Edge TTS 免费。支持流式——边生成边读,不必等全部生成完。 |
| **GPT-4o-mini** | 轻量分类模型 | 又快又便宜的小模型,用于模型路由——先用小模型判断问题复杂度,简单问题走小模型省 API 费用。 |

View File

@@ -1,5 +0,0 @@
1.视频录制
2.对话翻译
3.对话总结
4.手动对话功能
5.视频框大小可调整,可最小化然后拖动

1111
docs/10-鉴权体系.md Normal file

File diff suppressed because it is too large Load Diff

814
docs/11-令牌桶限流.md Normal file
View 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] = ttlkey 过期时间,秒)
local key = KEYS[1]
local capacity = tonumber(ARGV[1])
local rate = tonumber(ARGV[2])
local now = tonumber(ARGV[3])
local ttl = tonumber(ARGV[4])
local data = redis.call('HMGET', key, 'tokens', 'last_refill')
local tokens = tonumber(data[1]) or capacity
local last_refill = tonumber(data[2]) or now
-- 计算新令牌
local elapsed = math.max(0, now - last_refill)
tokens = math.min(capacity, tokens + elapsed * rate)
local allowed = 0
local retry_after = 0
if tokens >= 1 then
tokens = tokens - 1
allowed = 1
else
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] = ttlkey 过期时间,秒)
local key = KEYS[1]
local capacity = tonumber(ARGV[1])
local rate = tonumber(ARGV[2])
local now = tonumber(ARGV[3])
local ttl = tonumber(ARGV[4])
local data = redis.call('HMGET', key, 'tokens', 'last_refill')
local tokens = tonumber(data[1]) or capacity
local last_refill = tonumber(data[2]) or now
-- 计算新令牌
local elapsed = math.max(0, now - last_refill)
tokens = math.min(capacity, tokens + elapsed * rate)
local allowed = 0
local retry_after = 0
if tokens >= 1 then
tokens = tokens - 1
allowed = 1
else
if rate == 0 then
retry_after = 86400 -- 24小时
else
retry_after = (1 - tokens) / rate
end
end
-- 回写状态
redis.call('HMSET', key, 'tokens', tokens, 'last_refill', now)
redis.call('EXPIRE', key, ttl)
return {allowed, tostring(retry_after)}
```
**设计要点**
- 使用 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/)

View File

@@ -1,750 +0,0 @@
# 持久化与用户系统设计
## 概述
本文档定义用户注册/登录、JWT 认证、对话历史持久化的完整设计方案。核心目标:**用户登录后可在对话列表中选择历史对话继续交谈**。
### 设计决策
| 决策项 | 选择 | 理由 |
|--------|------|------|
| 认证方式 | JWTaccess + 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_tokenDELETE
生成新的 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 签名和过期(不查 DBDB 校验由 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/内存(热) ←→ PostgreSQLwrite-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-throughAppendMessage 同时写 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
View 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 Prompt10-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 字符
- Prompt10-2000 字符
- 首句引导:可选,最多 500 字符
**前端**
- 实时字符计数
- 超长提示
- 必填项高亮
## 未来优化方向
**V1.1**
- Prompt 模板库
- 实时预览效果
- 导入导出功能
- 情景搜索和筛选
**V2.0**
- 情景市场
- 情景分享链接
- AI 辅助优化 Prompt
- 协作编辑(团队情景)
## 参考资料
- [CLAUDE.md](../CLAUDE.md) — 项目开发指南
- [02-接口文档.md](./02-接口文档.md) — WebSocket 和 REST API
- [自建情景功能-权限隔离说明.md](./自建情景功能-权限隔离说明.md) — 安全设计

544
docs/13-日志追踪.md Normal file
View File

@@ -0,0 +1,544 @@
# 日志追踪系统
## 概述
CamTalk 全链路日志追踪系统,通过统一的 trace ID 机制,将 REST API 和 WebSocket 两大入口的所有日志串联起来,实现分布式环境下的请求链路可观测性。
**核心目标**
- 统一 trace ID 贯穿 REST/WebSocket 两大入口
- 所有日志自动附加 trace_id/request_id/session_id
- 保护用户隐私,敏感文本截断或降级
- 支持按 trace_id 快速定位完整请求链路
## Trace ID 作用域
| 标识 | 作用域 | 生成时机 | 用途 |
|-----|--------|---------|------|
| `trace_id` | **连接级**(整个 WebSocket 生命周期)<br/>**请求级**(单次 REST 请求) | REST: 中间件生成<br/>WebSocket: 升级时生成 | 关联同一连接/请求的所有日志 |
| `session_id` | 会话级(对话上下文存储) | ServeWS 时生成 | 标识会话存储 |
| `request_id` | 查询级(单次 WebSocket 查询) | 客户端每次查询传入 | 区分同一连接的不同查询 |
**WebSocket 场景示例**:用户打开页面建立 WebSocket发起 3 次对话查询:
```
连接建立 trace_id=01J5AAA session_id=uuid-123
├─ 查询1 trace_id=01J5AAA request_id=req-001 (问天气)
├─ 查询2 trace_id=01J5AAA request_id=req-002 (问新闻)
└─ 查询3 trace_id=01J5AAA request_id=req-003 (问股票)
```
**REST 场景示例**
```
POST /api/auth/login trace_id=01J5BBB request_id=01J5BBB
GET /api/conversations trace_id=01J5CCC request_id=01J5CCC
```
## 核心组件
```mermaid
graph TB
subgraph trace包["trace 包"]
ID["id.go<br/>ULID 生成器"]
CTX["context.go<br/>context key 管理"]
LOG["logger.go<br/>context-aware logger"]
MW["middleware.go<br/>Gin trace 中间件"]
end
subgraph logger包["logger 包"]
GINLOG["middleware.go<br/>Gin 请求日志"]
GINREC["GinRecovery<br/>panic 恢复"]
end
subgraph 入口层["入口层"]
REST["REST API<br/>trace 中间件注入"]
WS["WebSocket<br/>ServeWS 注入"]
end
subgraph 业务层["业务层"]
HANDLER["Handler"]
ADAPTER["Eino Adapter"]
NODES["Eino Nodes"]
end
subgraph 存储层["存储层"]
PG["PostgreSQL<br/>session/user/message/scenario"]
REDIS["Redis<br/>session/cache/ratelimit"]
end
ID --> MW
CTX --> LOG
LOG --> HANDLER
LOG --> ADAPTER
LOG --> NODES
LOG --> PG
LOG --> REDIS
MW --> REST
GINLOG --> REST
WS --> LOG
```
### trace/id.go — ULID 生成器
使用 ULIDUniversally Unique Lexicographically Sortable Identifier作为 trace ID
- 时间排序:前 48 位是毫秒时间戳,天然按时间排序
- 唯一性:后 80 位随机数,冲突概率极低
- 并发安全:使用 `crypto/rand` + `sync.Pool` 复用 entropy 对象
```go
package trace
import (
cryptorand "crypto/rand"
"sync"
"time"
"github.com/oklog/ulid/v2"
)
var entropyPool = sync.Pool{
New: func() interface{} {
return ulid.Monotonic(cryptorand.Reader, 0)
},
}
// GenerateTraceID 生成并发安全的 ULID trace ID
func GenerateTraceID() string {
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
defer entropyPool.Put(entropy)
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
}
```
### trace/context.go — Context Key 管理
统一管理所有 trace 相关的 context key
```go
package trace
import "context"
type traceIDKey struct{}
type requestIDKey struct{}
type sessionIDKey struct{}
// WithTraceID 将 trace ID 注入 context
func WithTraceID(ctx context.Context, traceID string) context.Context {
return context.WithValue(ctx, traceIDKey{}, traceID)
}
func GetTraceID(ctx context.Context) string {
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
return v
}
return ""
}
// 类似定义 WithRequestID/GetRequestID 和 WithSessionID/GetSessionID
```
### trace/logger.go — Context-Aware Logger
自动从 context 提取 trace 字段并附加到日志:
```go
package trace
import (
"context"
"github.com/hhs/camtalk/internal/logger"
"go.uber.org/zap"
)
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
func FromContext(ctx context.Context) *zap.SugaredLogger {
log := logger.Log
if traceID := GetTraceID(ctx); traceID != "" {
log = log.With("trace_id", traceID)
}
if requestID := GetRequestID(ctx); requestID != "" {
log = log.With("request_id", requestID)
}
if sessionID := GetSessionID(ctx); sessionID != "" {
log = log.With("session_id", sessionID)
}
return log
}
```
**使用模式对比**
```go
// Before: 手动传递字段
logger.Log.Infow("message", "session", sessionID, "request", requestID)
// After: 自动附加
trace.FromContext(ctx).Infow("message")
```
### trace/middleware.go — Gin Trace 中间件
为 REST 请求生成 trace ID 并注入 context
```go
package trace
import "github.com/gin-gonic/gin"
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
func TraceMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
traceID := GenerateTraceID()
ctx := WithTraceID(c.Request.Context(), traceID)
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
c.Request = c.Request.WithContext(ctx)
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
c.Next()
}
}
```
### logger/middleware.go — 请求日志与 Panic 恢复
记录所有 HTTP 请求的 method/path/status/latency
```go
package logger
import (
"time"
"github.com/gin-gonic/gin"
"github.com/hhs/camtalk/internal/trace"
)
// GinLogger 记录每个 HTTP 请求的基础信息
func GinLogger() gin.HandlerFunc {
return func(c *gin.Context) {
start := time.Now()
path := c.Request.URL.Path
c.Next()
latency := time.Since(start).Milliseconds()
status := c.Writer.Status()
log := trace.FromContext(c.Request.Context())
switch {
case status >= 500:
log.Errorw("request completed", "method", c.Request.Method,
"path", path, "status", status, "latency_ms", latency)
case status >= 400:
log.Warnw("request completed", "method", c.Request.Method,
"path", path, "status", status, "latency_ms", latency)
default:
log.Infow("request completed", "method", c.Request.Method,
"path", path, "status", status, "latency_ms", latency)
}
}
}
// GinRecovery 自定义 panic 恢复中间件
func GinRecovery() gin.HandlerFunc {
return func(c *gin.Context) {
defer func() {
if err := recover(); err != nil {
log := trace.FromContext(c.Request.Context())
log.Errorw("panic recovered", "error", err,
"path", c.Request.URL.Path, "method", c.Request.Method)
c.AbortWithStatus(500)
}
}()
c.Next()
}
}
```
## 中间件注册顺序
`cmd/server/main.go` 中,三层中间件按顺序注册:
```go
r := gin.New()
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
r.Use(logger.GinLogger()) // 第二层:记录请求
r.Use(logger.GinRecovery()) // 第三层panic 恢复
```
## 日志输出示例
### REST 请求
```json
{
"level": "info",
"ts": 1718956800.123,
"msg": "login success",
"trace_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
"request_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
"username": "test_user"
}
```
### WebSocket 查询链路(含存储层)
```json
// 1. 查询接收
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query received"}
// 2. 会话加载Redis
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"redis session retrieved", "session_id":"abc-123"}
// 3. STT 完成
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"stt recognition completed", "text_len":45}
// 4. LLM 完成
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"llm generation completed", "tokens":150}
// 5. 消息持久化PostgreSQL
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"user", "tokens_used":45}
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"assistant", "tokens_used":150}
// 6. Pipeline 完成
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query processing completed", "latency_ms":2340}
```
### 限流触发场景
```json
{"level":"warn", "trace_id":"01J5YYY", "msg":"rate limit triggered", "key":"ratelimit:user-456:query", "retry_after_sec":2.5}
```
## 日志查询操作
### 按 trace_id 查询完整链路
**本地开发(文件日志)**
```bash
# 查看完整链路
grep 'trace_id":"01J5XXX"' backend.log | jq .
# 查看链路时间线
grep 'trace_id":"01J5XXX"' backend.log | jq -r '[.ts, .msg] | @tsv'
```
**Grafana Loki**
```logql
{app="camtalk-backend"}
|= "trace_id=01J5XXX"
| json
| line_format "{{.ts}} [{{.level}}] {{.msg}}"
```
### 查询慢请求(延迟 > 5s
```logql
{app="camtalk-backend"}
| json
| msg="query processing completed"
| latency_ms > 5000
```
### 查询数据库错误
```logql
{app="camtalk-backend"}
| json
| level="error"
| msg=~".*failed"
| line_format "{{.trace_id}} {{.msg}} {{.error}}"
```
### 查询 Redis 降级事件
```logql
{app="camtalk-backend"}
| json
| level="warn"
| msg=~"redis.*failed"
```
### 查询错误率
```logql
sum(count_over_time({app="camtalk-backend"} | json | level="error" [5m]))
```
## 敏感内容处理规范
### 完全禁止记录
- 用户明文密码
- JWT token 完整内容(仅记录 "token_present: true"
- API Key 完整值(仅记录前 8 字符 + "..."
### 截断后记录(最多 50 字符)
- 用户输入文本 → `text_preview`
- LLM 生成文本 → `text_preview`
- STT 识别文本 → `text_preview`
**示例**
```go
log.Debugw("stt recognition completed",
"text_len", len(text),
"text_preview", util.Truncate(text, 50))
```
### 仅记录长度/大小
- 图片数据 → `image_size_bytes`
- 音频数据 → `audio_size_bytes`
### 降级为 Debug 级别
所有包含用户文本预览的日志,生产环境默认不输出。
## 日志级别使用准则
| 场景 | 级别 | 示例 |
|-----|------|-----|
| 请求生命周期里程碑 | Info | `"query received"`, `"pipeline completed"` |
| 中间步骤详情 | Debug | `"stt recognition completed"`, `"history assembled"` |
| 敏感内容相关 | Debug | 所有包含用户文本的日志 |
| 预期内的失败 | Warn | `"login failed"`, `"rate limited"` |
| 系统错误 | Error | `"database query failed"`, `"tts synthesis failed"` |
| 严重故障 | Error + stack | `"panic recovered"` |
## 存储层日志实现
### PostgreSQL Repository 层
所有数据库操作统一使用 `trace.FromContext(ctx)` 记录日志:
**已实现文件**
- `backend/internal/store/session_pg.go` — 会话 CRUD
- `backend/internal/store/user_pg.go` — 用户与 refresh token 操作
- `backend/internal/store/message_pg.go` — 对话消息存储
- `backend/internal/store/user_scenario_repository.go` — 用户自定义情景
**日志策略**
```go
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
log := trace.FromContext(ctx)
_, err := r.pool.Exec(ctx, ...)
if err != nil {
log.Errorw("save session failed", "session_id", s.ID, "error", err)
return err
}
log.Debugw("session saved", "session_id", s.ID, "user_id", s.UserID)
return nil
}
```
**NotFound 处理**:预期内的空结果不记录错误:
```go
if errors.Is(err, pgx.ErrNoRows) {
return nil, ErrSessionNotFound // 不记录日志
}
if err != nil {
log.Errorw("find session failed", "session_id", id, "error", err)
return nil, err
}
```
### Redis 服务层
**已实现文件**
- `backend/internal/session/redis.go` — RedisManager会话存储
- `backend/internal/store/cached_user.go` — CachedUserRepository用户缓存装饰器
- `backend/internal/ratelimit/redis_bucket.go` — RedisLimiter令牌桶限流器
**会话存储日志**`redis.go`
```go
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
log := trace.FromContext(ctx)
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
if err != nil {
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
return nil, fmt.Errorf("redis get session: %w", err)
}
if len(vals) == 0 {
return nil, ErrSessionNotFound // 不记录日志
}
log.Debugw("redis session retrieved", "session_id", sessionID)
return session, nil
}
```
**缓存降级日志**`cached_user.go`
```go
if _, err := pipe.Exec(ctx); err != nil {
log := trace.FromContext(ctx)
log.Warnw("redis cache write failed for refresh token", "error", err)
// 降级DB 已写入成功Redis 失败不影响正确性
}
```
**限流触发日志**`redis_bucket.go`
```go
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
log := trace.FromContext(ctx)
result, err := l.script.Run(ctx, ...).Result()
if err != nil {
log.Errorw("rate limit check failed", "key", key, "error", err)
return true, 0 // fail-open 策略
}
if allowed == 0 {
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
return false, retryAfter
}
return true, 0
}
```
**级别选择原则**
- **Error**Redis 连接失败、Lua 脚本执行失败(影响功能)
- **Warn**:缓存写入失败(可降级)、限流触发(预期内异常)
- **Debug**:正常操作完成(避免 Info 级别噪音)
## 编码规范
1. **日志语言**:统一使用英文
2. **结构化**:始终使用 `Infow`/`Errorw`/`Warnw`/`Debugw`
3. **Context 传递**:使用 `trace.FromContext(ctx)` 而非直接引用 `logger.Log`
4. **敏感内容**:禁止在 Info 及以上级别记录用户文本原文
5. **错误日志**:采用"调用方记录"原则,底层函数 return wrapped error
6. **级别约定**
- `Debug`:内部状态跟踪、开发调试信息(数据库/缓存成功操作)
- `Info`:请求/连接生命周期、关键操作里程碑
- `Warn`可降级异常Redis 故障、限流触发)
- `Error`影响用户的操作失败数据库错误、Redis 连接失败)
- `Fatal`:仅启动阶段不可恢复错误
7. **预期内的空结果**`pgx.ErrNoRows``redis.Nil` 等不记录错误日志
## 性能考量
### FromContext 开销
- 有 trace_id~200-300 ns/op
- 无 trace_id~10-20 ns/op仅返回全局 logger
- 1000 QPS 场景额外开销约 0.2ms,可接受
### ULID 生成吞吐量
- 单线程:~500k ops/s
- 并发 8 线程:~2M ops/s
**验收标准**1000 QPS 下trace 系统开销 < 1% CPU< 0.5ms P99 延迟。

View File

@@ -0,0 +1,292 @@
---
tags: [eino, go, ai-agent, rag, tool-calling, project-theme]
create time: 2026-05-05 14:30
---
# EiNO 项目实战:个人知识助手
## 概述
基于字节跳动 EINO 框架Go 语言),从零构建一个**个人知识助手 Agent**。该项目深度融合 **RAG检索增强生成****Tool Calling工具调用** 两大核心能力,采用 Supervisor 多 Agent 编排模式,帮助用户高效检索笔记、整理知识、发现关联。场景聚焦日常知识管理,不依赖外部基础设施,是入门 EINO 框架的理想项目。
## 正文
### 1. 项目背景
> [!question] 思考:当你积累了上千篇笔记,某天想找一篇"半年前记过的 Go 并发模式",只记得大概内容却忘了标题——你会怎么办?
传统做法:逐个翻文件夹 → 搜关键词 → 翻了几分钟还是没找到。而一个智能知识助手可以:
1. **听懂模糊描述**"那个讲 goroutine 泄漏排查的文章"→ 语义检索精准定位
2. **整理碎片知识**"把最近关于 EINO 的笔记汇总成一篇综述"
3. **发现隐藏关联**"这篇 RAG 笔记和那篇向量数据库笔记其实在讲同一件事"
这个场景**天然适合 AI Agent**检索知识库RAG+ 操作笔记Tool Calling+ 多步骤任务ReAct
**用 EINO 的原因**
- Go 原生协程,本地跑也轻量
- 编译时类型检查,工具定义清晰、不易出错
- ADK 内置 Supervisor / Plan-Execute / Interrupt 等模式,开箱即用
---
### 2. 系统架构总览
```mermaid
graph TD
U["用户提问"] --> S["Supervisor Agent<br/>知识总管家"]
S --> R["Retrieval Agent<br/>知识检索专家"]
S --> W["Writer Agent<br/>内容处理专家"]
S --> O["Organizer Agent<br/>知识整理专家"]
R --> VDB["向量数据库<br/>笔记内容索引"]
R --> T1["Tool: 语义搜索<br/>相似笔记召回"]
R --> T2["Tool: 关键词搜索<br/>精确匹配"]
W --> T3["Tool: 创建笔记<br/>写入 Markdown"]
W --> T4["Tool: 摘要提取<br/>生成笔记摘要"]
W --> T5["Tool: 标签推荐<br/>自动打标签"]
O --> T6["Tool: 关联发现<br/>Wiki-link 推荐"]
O --> T7["Tool: 知识图谱<br/>关联关系查询"]
style S fill:#4A90D9,color:#fff
style VDB fill:#27AE60,color:#fff
```
> **Supervisor 模式**Supervisor Agent 接收用户指令,根据意图路由——搜索类交给 Retrieval Agent、写作类交给 Writer Agent、整理类交给 Organizer Agent最终汇总结果返回。
---
### 3. RAG 模块设计
RAG 负责从用户的笔记库中检索相关内容,让 Agent "读懂你的知识库"。
#### 3.1 笔记索引管道
```mermaid
flowchart LR
A["Markdown 笔记库"] --> B["Document Loader<br/>按段落分块"]
B --> C["Embedding<br/>文本向量化"]
C --> D["Indexer<br/>写入向量库"]
D --> E["Hybrid Retriever<br/>混合检索"]
E --> F["Agent<br/>上下文注入"]
```
#### 3.2 核心代码:混合检索器
```go
// RetrieverService 混合检索:语义匹配 + 标签过滤
type RetrieverService struct {
client milvus.Client
embedder embedding.Embedder
}
func (s *RetrieverService) Retrieve(ctx context.Context, query string, tags []string) ([]*schema.Document, error) {
// 1. 将用户查询转为向量
vector, err := s.embedder.Embed(ctx, query)
if err != nil {
return nil, fmt.Errorf("embed query: %w", err)
}
// 2. 构建标量过滤:限定标签范围
// 例如: "tag in ['go', 'concurrency', 'eino']"
expr := buildTagFilterExpr(tags)
// 3. 混合检索Top-K=5
results, err := s.client.Search(ctx, "notes_collection",
nil, expr,
[]string{"content", "title", "tags"},
vector,
milvus.NewTopKMetricType(milvus.L2, 5),
milvus.NewSearchParam(16),
)
// ... 转换为 EINO Document 格式
return s.convertToDocs(results), nil
}
```
> [!tip] 设计要点
> 混合检索 = **语义相似度**"我记得大概意思" + **标签过滤**"应该是 Go 相关的")。相比纯关键词搜索,它能找到表述不同但意思相近的笔记——这正是知识管理中最常见的场景。
---
### 4. Tool Calling 模块设计
Agent 通过工具与笔记系统交互:搜索、创建、整理、发现关联。
#### 4.1 工具清单
| 工具 | 类型 | 描述 |
|------|------|------|
| `semantic_search` | 查询 | 语义搜索笔记,支持模糊自然语言描述 |
| `keyword_search` | 查询 | 精确关键词 + 标签搜索 |
| `create_note` | 写入 | 创建新笔记Markdown + YAML frontmatter |
| `generate_summary` | 处理 | 为指定笔记生成摘要 |
| `suggest_tags` | 处理 | 根据内容自动推荐标签 |
| `find_related` | 查询 | 发现关联笔记,推荐 Wiki-link |
#### 4.2 核心代码:定义工具
```go
// === 写笔记工具 ===
type CreateNoteParams struct {
Title string `json:"title" desc:"笔记标题"`
Content string `json:"content" desc:"Markdown 格式正文"`
Tags []string `json:"tags" desc:"标签列表,如 ['go', 'eino']"`
}
func CreateNoteTool(vaultPath string) componenttool.BaseTool {
return &componenttool.Tool{
Name: "create_note",
Desc: "在知识库中创建一篇新的 Markdown 笔记",
Func: func(ctx context.Context, params *CreateNoteParams) (string, error) {
fullPath := filepath.Join(vaultPath, params.Title+".md")
content := buildMarkdownWithFrontmatter(params)
if err := os.WriteFile(fullPath, []byte(content), 0o644); err != nil {
return "", fmt.Errorf("write note: %w", err)
}
return fmt.Sprintf("笔记已创建: %s", fullPath), nil
},
}
}
// === 关联发现工具 ===
type FindRelatedParams struct {
NoteTitle string `json:"note_title" desc:"目标笔记标题"`
}
func FindRelatedTool(retriever *RetrieverService) componenttool.BaseTool {
return &componenttool.Tool{
Name: "find_related",
Desc: "根据笔记内容,从知识库中发现与之关联的其他笔记,推荐 Wiki-link",
Func: func(ctx context.Context, params *FindRelatedParams) (string, error) {
noteContent := readNote(params.NoteTitle)
related, _ := retriever.Retrieve(ctx, noteContent, nil)
return formatWikiLinkSuggestions(related), nil
},
}
}
```
> [!question] 思考:如果用户说"帮我把最近一周关于 EINO 的笔记整理成一篇综述"Agent 需要依次调用哪些工具?顺序能否调换?
---
### 5. Multi-Agent 编排Supervisor 模式
```go
// === 构建 Supervisor编排三个 Specialist Agent ===
func BuildKnowledgeSupervisor(ctx context.Context) (*adk.Supervisor, error) {
// Retrieval Agent: 负责搜索和检索
retrievalAgent, _ := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
Name: "RetrievalAgent",
Instruction: "你是知识检索专家,擅长从笔记库中找到最相关的内容...",
Model: model,
ToolsConfig: adk.ToolsConfig{
ToolsNodeConfig: compose.ToolsNodeConfig{
Tools: []componenttool.BaseTool{
SemanticSearchTool(retriever),
KeywordSearchTool(),
FindRelatedTool(retriever),
},
},
},
MaxIterations: 10,
})
// Writer Agent: 负责创建和整理内容
writerAgent, _ := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
Name: "WriterAgent",
Model: model,
ToolsConfig: adk.ToolsConfig{
ToolsNodeConfig: compose.ToolsNodeConfig{
Tools: []componenttool.BaseTool{
CreateNoteTool(vaultPath),
SummaryTool(model),
SuggestTagsTool(model),
},
},
},
MaxIterations: 8,
})
// Organizer Agent: 负责关联发现和知识图谱
organizerAgent, _ := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{
// ... 配置 FindRelated 等整理工具
MaxIterations: 8,
})
return adk.NewSupervisor(ctx, &adk.SupervisorConfig{
Name: "KnowledgeSupervisor",
Model: model,
Instruction: "你是知识总管。根据用户意图路由任务...",
SubAgents: []adk.Agent{retrievalAgent, writerAgent, organizerAgent},
})
}
```
```mermaid
sequenceDiagram
participant U as 用户
participant S as Supervisor
participant R as Retrieval Agent
participant W as Writer Agent
U->>S: "帮我整理最近关于 EINO 的笔记,<br/>写一篇学习综述"
S->>S: 意图分类 → 检索 + 写作(复合任务)
S->>R: 委派检索任务
R->>R: Tool: semantic_search("EINO") → 找到 5 篇
R->>R: Tool: find_related → 发现 3 篇关联笔记
R->>S: 返回 8 篇笔记及其内容
S->>W: 委派写作任务(附检索结果)
W->>W: 基于 8 篇笔记撰写综述草稿
W->>W: Tool: suggest_tags → ["eino", "agent", "go"]
W->>S: 返回综述 + 推荐标签
S->>U: 📝 综述全文 + 🏷️ 推荐标签 + 🔗 关联笔记
```
---
### 6. 进阶拓展方向
> [!info] 掌握了基础架构后,可以探索这些方向
1. **会话式检索**:支持多轮追问——"上次那篇关于 goroutine 的"→"不是那篇,是讲泄漏排查的"→ Agent 结合上下文逐步缩小范围,像和人对话一样自然。
2. **定时知识回顾**:用 Cron 触发 Agent每周自动检索本周新增笔记 → 生成"本周知识地图"→ 推送回顾通知。类似间隔重复,但由 AI 驱动。
3. **跨源增强**当本地笔记不足时Agent 可调用 `web_search` 工具获取外部信息作为补充,生成"本地知识 + 外部参考"的混合回答,并标注来源。
4. **知识冲突检测**:当新笔记与已有笔记表述矛盾(如某篇写"Go defer 是栈顺序",另一篇写"是队列顺序"Agent 自动标记冲突,提醒用户核实。
5. **MCP 集成**:将工具标准化为 MCP Server让其他 AI 客户端(如 Claude Desktop、VS Code 插件)也能直接调用你的知识助手。
---
### 7. 关键 EINO 概念速查
| EINO 概念 | 本项目对应 |
|-----------|-----------|
| `ChatModelAgent` | RetrievalAgent / WriterAgent / OrganizerAgent |
| `Supervisor` | KnowledgeSupervisor多 Agent 总调度) |
| `Tool / ToolsNode` | 语义搜索、创建笔记、摘要、标签、关联发现 |
| `Retriever` | 混合检索器(语义 + 标签过滤) |
| `Embedding` | 文本向量化 |
| `Indexer` | 笔记内容写入向量库 |
| `Document Loader` | Markdown 笔记按段落分块加载 |
| `ChatTemplate` | System Prompt + 检索结果注入 |
| `Interrupt / Resume` | 删改操作前的确认拦截(进阶) |
| `Graph / Compile` | 编排所有节点、编译成可执行图 |
---
## 关联笔记
- [[EINO ADK 深入]]

View File

@@ -0,0 +1,245 @@
---
tags: [eino, llm-framework, golang, agent, ai]
create time: 2026-04-29 21:30
---
# Eino 知识索引
## 概述
Eino 是字节跳动CloudWeGo开源的大模型应用开发框架基于 Go 语言。覆盖从组件定义、流程编排到 DevOps 工具链的全流程。本文档作为知识索引,列举 Eino 的关键概念与使用入口,后续可基于此构建知识问答。
---
## 一、是什么 —— Eino 核心定位
| 维度 | 说明 |
|------|------|
| **语言** | Go强类型编译时类型检查 |
| **定位** | 大模型应用开发框架,覆盖全流程 |
| **仓库** | [github.com/cloudwego/eino](https://github.com/cloudwego/eino) + [eino-ext](https://github.com/cloudwego/eino-ext) |
| **特色** | 组件抽象 → 图编排 → ADK Agent → DevOps 工具链,层层递进 |
| **适用场景** | 从简单对话到复杂 Multi-Agent 系统,均可应对 |
> [!question] 思考:为什么选 Go 而不是 Python强类型在大模型应用规模化后能带来什么收益
---
## 二、怎么分层 —— Eino 架构分层
```mermaid
graph TD
A["Eino 框架"] --> B["组件层 Components"]
A --> C["编排层 Orchestration"]
A --> D["ADK 层 Agent 开发套件"]
A --> E["工具层 DevOps"]
B --> B1["ChatModel"]
B --> B2["ChatTemplate"]
B --> B3["Tool / ToolsNode"]
B --> B4["Retriever"]
B --> B5["Document Loader"]
B --> B6["Embedding / Indexer"]
B --> B7["Lambda 自定义"]
C --> C1["Chain 链式"]
C --> C2["Graph 有向图"]
C --> C3["Workflow 字段映射"]
C --> C4["Stream 流处理"]
C --> C5["Callback 横切面"]
D --> D1["Agent 接口"]
D --> D2["ChatModelAgent ReAct"]
D --> D3["WorkflowAgents"]
D --> D4["Multi-Agent 范式"]
D --> D5["Middleware 中间件"]
E --> E1["Tracing 链路追踪"]
E --> E2["IDE 插件 / 可视化"]
E --> E3["Debug 调试"]
```
---
## 三、怎么用 —— 快速上手路径
### 3.1 入门四步走
| 步骤 | 主题 | 入口文档 |
|------|------|----------|
| **Step 1** | ChatModel 与 Message学会调模型 | [[Eino/quick_start/chapter_01_chatmodel_and_message]] |
| **Step 2** | ChatModelAgent + Runner + AgentEvent学会跑 Agent | [[Eino/quick_start/chapter_02_chatmodelagent_runner_agentevent]] |
| **Step 3** | Memory & Session让 Agent 有记忆 | [[Eino/quick_start/chapter_03_memory_and_session]] |
| **Step 4** | Tool & FileSystem让 Agent 能干活 | [[Eino/quick_start/chapter_04_tool_and_filesystem]] |
### 3.2 进阶专题
| 主题 | 关键内容 | 入口文档 |
|------|----------|----------|
| **Middleware** | 横切面注入PlanTask / Summarization / Skill 等内置中间件 | [[Eino/quick_start/chapter_05_middleware]] |
| **Callback & Trace** | 执行过程可观测、可追踪 | [[Eino/quick_start/chapter_06_callback_and_trace]] |
| **Interrupt & Resume** | Human-in-the-loop断点续跑 | [[Eino/quick_start/chapter_07_interrupt_resume]] |
| **Graph & Tool** | 低代码编排,图即是工具 | [[Eino/quick_start/chapter_08_graph_tool]] |
| **A2UI Protocol** | Agent 生成 UI 协议 | [[chapter_10_a2ui_protocol]] |
| **Skill Console** | 技能注册与管理 | [[Eino/quick_start/chapter_09_skill_console]] |
---
## 四、关键概念清单
### 4.1 组件 Component
| 组件 | 职责 | 为什么需要 |
|------|------|-----------|
| **ChatModel** | 与大模型交互的核心接口 | 统一不同模型OpenAI / Ark / Gemini的调用方式 |
| **ChatTemplate** | 构造 Prompt 模板 | 将变量注入系统/用户消息,结构化管理 |
| **Tool / ToolsNode** | 可被模型调用的工具 | 让 LLM 从"只会说"变成"能执行" |
| **Retriever** | 检索外部知识 | RAG 的核心,给模型注入上下文 |
| **Document Loader** | 加载各类文档 | 将 PDF / Markdown / 网页等转为可处理文本 |
| **Embedding / Indexer** | 向量化 + 索引 | 语义检索基础 |
| **Lambda** | 自定义函数作为组件 | 当现有组件不满足需求时,自由扩展 |
> [!question] 思考Lambda 和 Tool 的区别是什么?什么场景用哪个?
### 4.2 编排 Orchestration
| 编排方式 | 特点 | 适用场景 |
|----------|------|----------|
| **Chain** | 简单的顺序执行,有向无环 | 线性 Pipeline |
| **Graph** | 有向图(可含环),灵活分支控制 | ReAct Agent 等复杂路由 |
| **Workflow** | 字段级别的数据映射与传递 | 抖音场景:字段粒度的图映射 |
额外能力:
- **Stream 流处理**:自动处理流式输入输出,支持流的复制、合并、拼接
- **Callback 横切面**:在组件执行前后注入逻辑(日志、追踪、统计)
- **编译时类型检查**Graph Compile 时验证节点类型对齐,而非运行时才发现
### 4.3 ADK —— Agent 开发套件
#### Agent 接口(统一抽象)
```go
type Agent interface {
Name(ctx context.Context) string
Description(ctx context.Context) string
Run(ctx context.Context, input *AgentInput, options ...AgentRunOption) *AsyncIterator[*AgentEvent]
}
```
三个核心要素:**身份Name** + **职责Description** + **标准化执行Run → AsyncIterator**
#### Agent 类型全景
```mermaid
graph LR
A["Agent 接口"] --> B["ChatModelAgent\nReAct 模式"]
A --> C["WorkflowAgents\n流程编排"]
A --> D["Multi-Agent\n协作范式"]
A --> E["自定义 Agent\n实现接口"]
C --> C1["Sequential\n顺序执行"]
C --> C2["Parallel\n并发执行"]
C --> C3["Loop\n循环执行"]
D --> D1["Supervisor\n集中式协调"]
D --> D2["Plan-Execute\n规划-执行-反思"]
D --> D3["DeepAgents\n规划驱动集中协作"]
```
| Agent 类型 | 一句话描述 | 典型场景 |
|------------|-----------|----------|
| **ChatModelAgent** | 基于 ReAct 的思考-行动循环 | 需要模型自主决策和工具调用的场景 |
| **SequentialAgent** | 按顺序依次执行子 Agent | ETL 流水线、CI/CD |
| **ParallelAgent** | 多个子 Agent 并发执行 | 多源数据采集、多渠道推送 |
| **LoopAgent** | 循环执行直到满足退出条件 | 数据同步、迭代优化 |
| **Supervisor** | 中心调度,统一分配与汇总 | 科研项目管理、客服流程 |
| **Plan-Execute** | Planner → Executor → Replanner 闭环 | Excel 处理、多步骤推理 |
| **DeepAgents** | Main Agent + WriteTodos + TaskTool | 长流程阶段性管理、多角色协作 |
#### Agent 协作机制
| 机制 | 使用方式 | 适用场景 |
|------|---------|----------|
| **History** | 框架自动传递前序 Agent 输出 | Agent 间默认信息流 |
| **Shared Session** | KV 存储,`GetSessionValue` / `AddSessionValue` | 跨 Agent 共享状态 |
| **Transfer移交** | `NewTransferToAgentAction` 将任务移交子 Agent | 边界清晰的层级式分工 |
| **ToolCall工具调用** | `NewAgentTool` 将 Agent 封装为 Tool | 仅需参数、无需完整上下文时 |
#### 中断与恢复
- Agent 运行时通过 `Interrupt Action` 主动中断
- `CheckPointStore` 保存断点状态
- `Resume` 方法携带新信息从断点恢复
- 适用场景:需要外部输入、人工审批、长时等待
#### Middleware 中间件体系
| 内置中间件 | 功能 |
|------------|------|
| **PlanTask** | 任务规划与拆解 |
| **Summarization** | 上下文摘要压缩 |
| **ToolSearch** | 动态工具搜索与选择 |
| **ToolReduction** | 工具调用结果精简 |
| **Skill** | 技能注册与匹配 |
| **FileSystem** | 文件读写能力注入 |
| **PatchToolCalls** | 工具调用修正 |
---
## 五、实战案例索引
| 案例 | 核心知识点 | 入口文档 |
|------|-----------|----------|
| **ReAct Agent** | Graph 编排 + Tool 调用 | [[Eino/docs/overview/eino_open_source]] |
| **Excel Agent** | Plan-Execute + Multi-Agent + CodeAgent | [[Eino/docs/overview/eino_adk_excel_agent]] |
| **项目管理 Agent** | Supervisor + 中断恢复 + Transfer | [[Eino/docs/overview/eino_adk0_1]] |
| **字节内部实践** | 豆包、抖音等真实场景 | [[Eino/docs/overview/bytedance_eino_practice]] |
---
## 六、进阶对比与选型
| 对比维度 | Eino | LangChain / LlamaIndex |
|----------|------|------------------------|
| **语言** | Go强类型 | Python动态类型 |
| **类型安全** | 编译时校验 | 运行时才发现 |
| **长期维护** | 类型系统天然可维护 | 动态语言大型项目维护成本高 |
| **并发模型** | goroutine 原生并发 | asyncio / 多线程 |
| **编排能力** | Graph + Workflow 双模式 | LCEL + Graph |
| **Agent 框架** | ADK 统一抽象 + 多种范式 | LangGraph Agent |
> [!question] 思考Eino 的 Graph vs Agent什么时候直接用 Graph什么时候用 ADK Agent
---
## 七、版本演进速览
| 版本 | 关键变化 |
|------|----------|
| **v0.1** | 首个开源版本,核心组件 + Chain/Graph |
| **v0.2** | Callback 体系完善 |
| **v0.3** | 小幅 break change |
| **v0.4** | Compose 优化 |
| **v0.5** | ADK 实现 |
| **v0.6** | JSON Schema 优化 |
| **v0.7** | Interrupt/Resume 重构 |
| **v0.8** | ADK Middleware 体系 |
---
## 八、外部资源
- 项目主页:[https://www.cloudwego.io](https://www.cloudwego.io)
- GitHub[https://github.com/cloudwego/eino](https://github.com/cloudwego/eino)
- 扩展库:[https://github.com/cloudwego/eino-ext](https://github.com/cloudwego/eino-ext)
- 示例:[https://github.com/cloudwego/eino-examples](https://github.com/cloudwego/eino-examples)
- 文档:[https://www.cloudwego.io/zh/docs/eino/](https://www.cloudwego.io/zh/docs/eino/)
---
## 关联笔记
- [[Eino/docs/overview/eino_open_source]] — 开源发布文
- [[Eino/docs/overview/eino_adk0_1]] — ADK 设计模式详解
- [[Eino/docs/overview/eino_adk_excel_agent]] — Excel Agent 实战
- [[Eino/quick_start/README]] — 快速开始总览

View File

@@ -0,0 +1,22 @@
---
Description: ""
date: "2025-07-21"
lastmod: ""
tags: []
title: 核心模块
weight: 4
---
Eino 中的核心模块有如下几个部分:
- **Components 组件**[Eino: Components 组件](/zh/docs/eino/core_modules/components)
Eino 抽象出来的大模型应用中常用的组件,例如 `ChatModel``Embedding``Retriever` 等,这是实现一个大模型应用搭建的积木,是应用能力的基础,也是复杂逻辑编排时的原子对象。
- **Chain/Graph 编排**[Eino: Chain/Graph 编排功能](/zh/docs/eino/core_modules/chain_and_graph_orchestration/chain_graph_introduction)
多个组件混合使用来实现业务逻辑的串联Eino 提供 Chain/Graph 的编排方式,把业务逻辑串联的复杂度封装在了 Eino 内部,提供易于理解的业务逻辑编排接口,提供统一的横切面治理能力。
- **Flow 集成工具 (agents)**: [Eino: Flow 集成组件](/zh/docs/eino/core_modules/flow_integration_components)
Eino 把最常用的大模型应用模式封装成简单、易用的工具,让通用场景的大模型应用开发极致简化,目前提供了 `ReAct Agent``Host Multi Agent`

View File

@@ -0,0 +1,57 @@
---
Description: ""
date: "2025-07-21"
lastmod: ""
tags: []
title: Chain & Graph & Workflow 编排功能
weight: 2
---
在大模型应用中,`Components` 组件是提供 『原子能力』的最小单元,比如:
- `ChatModel` 提供了大模型的对话能力
- `Embedding` 提供了基于语义的文本向量化能力
- `Retriever` 提供了关联内容召回的能力
- `ToolsNode` 提供了执行外部工具的能力
> 详细的组件介绍可以参考: [Eino: Components 组件](/zh/docs/eino/core_modules/components)
一个大模型应用,除了需要这些原子能力之外,还需要根据场景化的业务逻辑,**对这些原子能力进行组合、串联**,这就是 **『编排』**。
大模型应用的开发有其自身典型的特征: 自定义的业务逻辑本身不会很复杂,几乎主要都是对『原子能力』的组合串联。
传统代码开发过程中,业务逻辑用 “代码的执行逻辑” 来表达,迁移到大模型应用开发中时,最直接想到的方法就是 “自行调用组件,自行把结果作为下一组件的输入进行调用”。这样的结果,就是 `代码杂乱``很难复用``没有切面能力`……
当开发者们追求代码『**优雅**』和『**整洁之道**』时,就发现把传统代码组织方式用到大模型应用中时有着巨大的鸿沟。
Eino 的初衷是让大模型应用开发变得非常简单,就一定要让应用的代码逻辑 “简单” “直观” “优雅” “健壮”。
Eino 对「编排」有着这样的洞察:
- 编排要成为在业务逻辑之上的清晰的一层,**不能让业务逻辑融入到编排中**。
- 大模型应用的核心是 “对提供原子能力的组件” 进行组合串联,**组件是编排的 “第一公民”**。
- 抽象视角看编排:编排是在构建一张网络,数据则在这个网络中流动,网络的每个节点都对流动的数据有格式/内容的要求,一个能顺畅流动的数据网络,关键就是 “**上下游节点间的数据格式是否对齐**?”。
- 业务场景的复杂度会反映在编排产物的复杂性上,只有**横向的治理能力**才能让复杂场景不失控。
- 大模型是会持续保持高速发展的,大模型应用也是,只有**具备扩展能力的应用才拥有生命力**。
于是Eino 提供了 “基于 Graph 模型 (node + edge) 的,以**组件**为原子节点的,以**上下游类型对齐**为基础的编排” 的解决方案。
具体来说,实现了如下特性:
- 一切以 “组件” 为核心,规范了业务功能的封装方式,让**职责划分变得清晰**,让**复用**变成自然而然
- 详细信息参考:[Eino: Components 组件](/zh/docs/eino/core_modules/components)
- 业务逻辑复杂度封装到组件内部,编排层拥有更全局的视角,让**逻辑层次变得非常清晰**
- 提供了切面能力callback 机制支持了基于节点的**统一治理能力**
- 详细信息参考:[Eino: Callback 用户手册](/zh/docs/eino/core_modules/chain_and_graph_orchestration/callback_manual)
- 提供了 call option 的机制,**扩展性**是快速迭代中的系统最基本的诉求
- 详细信息参考:[Eino: CallOption 能力与规范](/zh/docs/eino/core_modules/chain_and_graph_orchestration/call_option_capabilities)
- 提供了 “类型对齐” 的开发方式的强化,降低开发者心智负担,把 golang 的**类型安全**特性发挥出来
- 详细信息参考:[Eino: 编排的设计理念](/zh/docs/eino/core_modules/chain_and_graph_orchestration/orchestration_design_principles)
- 提供了 “**流的自动转换**” 能力,让 “流” 在「编排系统的复杂性来源榜」中除名
- 详细信息参考:[Eino 流式编程要点](/zh/docs/eino/core_modules/chain_and_graph_orchestration/stream_programming_essentials)
Graph 本身是强大且语义完备的,可以用这项底层几乎绘制出所有的 “数据流动网络”,比如 “分支”、“并行”、“循环”。
但 Graph 并不是没有缺点的,基于 “点” “边” 模型的 Graph 在使用时,要求开发者要使用 `graph.AddXXXNode()``graph.AddEdge()` 两个接口来创建一个数据通道,强大但是略显复杂。
而在现实的大多数业务场景中,往往仅需要 “按顺序串联” 即可因此Eino 封装了接口更易于使用的 `Chain`。Chain 是对 Graph 的封装,除了 “环” 之外Chain 暴露了几乎所有 Graph 的能力。

View File

@@ -0,0 +1,306 @@
---
Description: ""
date: "2025-11-20"
lastmod: ""
tags: []
title: CallOption 能力与规范
weight: 6
---
**CallOption**: 对 Graph 编译产物进行调用时,直接传递数据给特定的一组节点(Component、Implementation、Node)的渠道
- 和 节点 Config 的区别: 节点 Config 是实例粒度的配置也就是从实例创建到实例消除Config 中的值一旦确定就不需要改变了
- CallOption是请求粒度的配置不同的请求其中的值是不一样的。更像是节点入参但是这个入参是直接由 Graph 的入口直接传入,而不是上游节点传入。
- 举例:给一个 ChatModel 节点传入 Temperature 配置;给一个 Lambda 节点传入自定义 option。
## 组件 CallOption 形态
组件 CallOption 配置,有两个粒度:
- 组件的抽象(Abstract/Interface)统一定义的 CallOption 配置【组件抽象 CallOption】
- 组件的实现(Type/Implementation)定义的该类型组件专用的 CallOption 配置【组件实现 CallOption】
以 ChatModel 这个 Component 为例,介绍 CallOption 的形态
### Model 抽象与实现的目录
```
// 抽象所在代码位置
eino/components/model
├── interface.go
├── option.go // Component 抽象粒度的 CallOption 入参
// 抽象实现所在代码位置
eino-ext/components/model
├── claude
│   ├── option.go // Component 的一种实现的 CallOption 入参
│   └── chatmodel.go
├── ollama
│   ├── call_option.go // Component 的一种实现的 CallOption 入参
│   ├── chatmodel.go
```
### Model 抽象
如上所述,在定义组件的 CallOption 时,需要区分【组件抽象 CallOption】、【组件实现 CallOption】两种场景。 而是否要提供 【组件实现 CallOption】则是由 组件抽象 来决定的。
组件抽象提供的 CallOption 扩展能力如下(以 Model 为例,其他组件类似):
```go
package model
type ChatModel interface {
Generate(ctx context.Context, input []*schema.Message, opts ...Option) (*schema.Message, error)
Stream(ctx context.Context, input []*schema.Message, opts ...Option) (
*schema.StreamReader[*schema.Message], error)
// BindTools bind tools to the model.
// BindTools before requesting ChatModel generally.
// notice the non-atomic problem of BindTools and Generate.
BindTools(tools []*schema.ToolInfo) error
}
// 此结构体是【组件抽象CallOption】的统一定义。 组件的实现可根据自己的需要取用【组件抽象CallOption】的信息
// Options is the common options for the model.
type Options struct {
// Temperature is the temperature for the model, which controls the randomness of the model.
Temperature *float32
// MaxTokens is the max number of tokens, if reached the max tokens, the model will stop generating, and mostly return an finish reason of "length".
MaxTokens *int
// Model is the model name.
Model *string
// TopP is the top p for the model, which controls the diversity of the model.
TopP *float32
// Stop is the stop words for the model, which controls the stopping condition of the model.
Stop []string
}
// Option is the call option for ChatModel component.
type Option struct {
// 此字段是为【组件抽象CallOption】服务的 apply 方法,例如 WithTemperature
// 如果组件抽象不想提供【组件抽象CallOption】可不提供此字段同时不提供 GetCommonOptions() 方法
apply func(opts *Options)
// 此字段是为【组件实现CallOption】服务的 apply 方法。并假设 apply 方法为func(*T)
// 如果组件抽象不想提供【组件实现CallOption】可不提供此字段同时不提供 GetImplSpecificOptions() 方法
implSpecificOptFn any
}
// WithTemperature is the option to set the temperature for the model.
func WithTemperature(temperature float32) Option {
return Option{
apply: func(opts *Options) {
opts.Temperature = &temperature
},
}
}
// WithMaxTokens is the option to set the max tokens for the model.
func WithMaxTokens(maxTokens int) Option {
return Option{
apply: func(opts *Options) {
opts.MaxTokens = &maxTokens
},
}
}
// WithModel is the option to set the model name.
func WithModel(name string) Option {
return Option{
apply: func(opts *Options) {
opts.Model = &name
},
}
}
// WithTopP is the option to set the top p for the model.
func WithTopP(topP float32) Option {
return Option{
apply: func(opts *Options) {
opts.TopP = &topP
},
}
}
// WithStop is the option to set the stop words for the model.
func WithStop(stop []string) Option {
return Option{
apply: func(opts *Options) {
opts.Stop = stop
},
}
}
// GetCommonOptions extract model Options from Option list, optionally providing a base Options with default values.
func GetCommonOptions(base *Options, opts ...Option) *Options {
if base == nil {
base = &Options{}
}
for i := range opts {
opt := opts[i]
if opt.apply != nil {
opt.apply(base)
}
}
return base
}
// 组件实现方基于此方法封装自己的Option函数func WithXXX(xxx string) Option{}
func WrapImplSpecificOptFn[T any](optFn func(*T)) Option {
return Option{
implSpecificOptFn: optFn,
}
}
// GetImplSpecificOptions provides tool author the ability to extract their own custom options from the unified Option type.
// T: the type of the impl specific options struct.
// This function should be used within the tool implementation's InvokableRun or StreamableRun functions.
// It is recommended to provide a base T as the first argument, within which the tool author can provide default values for the impl specific options.
func GetImplSpecificOptions[T any](base *T, opts ...Option) *T {
if base == nil {
base = new(T)
}
for i := range opts {
opt := opts[i]
if opt.implSpecificOptFn != nil {
optFn, ok := opt.implSpecificOptFn.(func(*T))
if ok {
optFn(base)
}
}
}
return base
}
```
### Claude 实现
[https://github.com/cloudwego/eino-ext/blob/main/components/model/claude/option.go](https://github.com/cloudwego/eino-ext/blob/main/components/model/claude/option.go)
```go
package claude
import (
"github.com/cloudwego/eino/components/model"
)
type options struct {
TopK *int32
}
func WithTopK(k int32) model.Option {
return model.WrapImplSpecificOptFn(func(o *options) {
o.TopK = &k
})
}
```
[https://github.com/cloudwego/eino-ext/blob/main/components/model/claude/claude.go](https://github.com/cloudwego/eino-ext/blob/main/components/model/claude/claude.go)
```go
func (c *claude) genMessageNewParams(input []*schema.Message, opts ...model.Option) (anthropic.MessageNewParams, error) {
if len(input) == 0 {
return anthropic.MessageNewParams{}, fmt.Errorf("input is empty")
}
commonOptions := model.GetCommonOptions(&model.Options{
Model: &c.model,
Temperature: c.temperature,
MaxTokens: &c.maxTokens,
TopP: c.topP,
Stop: c.stopSequences,
}, opts...)
claudeOptions := model.GetImplSpecificOptions(&options{TopK: c.topK}, opts...)
// omit mulple lines...
return nil, nil
}
```
## 编排中的 CallOption
[https://github.com/cloudwego/eino/blob/main/compose/runnable.go](https://github.com/cloudwego/eino/blob/main/compose/runnable.go)
Graph 编译产物是 Runnable
```go
type Runnable[I, O any] interface {
Invoke(ctx context.Context, input I, opts ...Option) (output O, err error)
Stream(ctx context.Context, input I, opts ...Option) (output *schema.StreamReader[O], err error)
Collect(ctx context.Context, input *schema.StreamReader[I], opts ...Option) (output O, err error)
Transform(ctx context.Context, input *schema.StreamReader[I], opts ...Option) (output *schema.StreamReader[O], err error)
}
```
Runnable 各方法均接收 compose.Option 列表。
[https://github.com/cloudwego/eino/blob/main/compose/graph_call_options.go](https://github.com/cloudwego/eino/blob/main/compose/graph_call_options.go)
包括 graph run 整体的配置,各类组件的配置,特定 Lambda 的配置等。
```go
// Option is a functional option type for calling a graph.
type Option struct {
options []any
handler []callbacks.Handler
paths []*NodePath
maxRunSteps int
}
// DesignateNode set the key of the node which will the option be applied to.
// notice: only effective at the top graph.
// e.g.
//
// embeddingOption := compose.WithEmbeddingOption(embedding.WithModel("text-embedding-3-small"))
// runnable.Invoke(ctx, "input", embeddingOption.DesignateNode("my_embedding_node"))
func (o Option) DesignateNode(key ...string) Option {
nKeys := make([]*NodePath, len(key))
for i, k := range key {
nKeys[i] = NewNodePath(k)
}
return o.DesignateNodeWithPath(nKeys...)
}
// DesignateNodeWithPath sets the path of the node(s) to which the option will be applied to.
// You can make the option take effect in the subgraph by specifying the key of the subgraph.
// e.g.
// DesignateNodeWithPath({"sub graph node key", "node key within sub graph"})
func (o Option) DesignateNodeWithPath(path ...*NodePath) Option {
o.paths = append(o.paths, path...)
return o
}
// WithEmbeddingOption is a functional option type for embedding component.
// e.g.
//
// embeddingOption := compose.WithEmbeddingOption(embedding.WithModel("text-embedding-3-small"))
// runnable.Invoke(ctx, "input", embeddingOption)
func WithEmbeddingOption(opts ...embedding.Option) Option {
return withComponentOption(opts...)
}
```
compose.Option 可以按需分配给 Graph 中不同的节点。
<a href="/img/eino/graph_runnable_after_compile.png" target="_blank"><img src="/img/eino/graph_runnable_after_compile.png" width="100%" /></a>
```go
// 所有节点都生效的 call option
compiledGraph.Invoke(ctx, input, WithCallbacks(handler))
// 只对特定类型节点生效的 call option
compiledGraph.Invoke(ctx, input, WithChatModelOption(WithTemperature(0.5))
compiledGraph.Invoke(ctx, input, WithToolOption(WithXXX("xxx"))
// 只对特定节点生效的 call option
compiledGraph.Invoke(ctx, input, WithCallbacks(handler).DesignateNode("node_1"))
// 只对特定内部嵌套图或其中节点生效的 Call option
compiledGraph.Invoke(ctx, input, WithCallbacks(handler).DesignateNodeWithPath(NewNodePath("1", "2"))
```

Some files were not shown because too many files have changed in this diff Show More