Compare commits

...

409 Commits

Author SHA1 Message Date
b1c3958e9b docs: 更新 README.md 2026-07-17 10:43:27 +08:00
f9a01269d2 Merge pull request 'refactor: 优化数据库迁移文件约束和注释' (#211) from refatcor/table-structure into main
All checks were successful
Deploy / deploy (push) Successful in 39s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/211
2026-06-23 23:42:40 +08:00
hhs
74f1a8eb34 refactor: 优化数据库迁移文件约束和注释
**迁移文件优化**:
- 001_users.up.sql: 添加完整表和列注释,说明密码哈希算法 (bcrypt) 和 token 哈希算法 (SHA-256)
- 002_messages.up.sql: 收紧 role 字段至 VARCHAR(10),添加 tokens_used 非负约束,补充完整注释
- 003_sessions.up.sql: 收紧 title 字段至 VARCHAR(100) 并添加长度约束 (1-100),详细说明 config JSONB 结构
- 004_user_scenarios.up.sql: 修复 greeting 字段冲突 (VARCHAR(500)),扩大 icon 至 VARCHAR(20) 支持复合 Emoji,prompt 改为 TEXT 无上限,优化约束逻辑

**后端代码修改**:
- auth.go: 移除用户名最小长度限制 (3 字符),仅保留最大长度 64

**前端国际化修改**:
- 更新三种语言的登录/注册表单 placeholder 文本,移除字符长度要求提示
  - zh-CN: "请输入用户名" / "请输入密码"
  - en-US: "Enter username" / "Enter password"
  - ja-JP: "ユーザー名を入力" / "パスワードを入力"

**删除文件**:
- 移除临时修复迁移文件 (005_fix_user_scenarios_description.*)
- 删除临时诊断脚本 (fix_user_scenarios.sql, check_constraints.sql)

所有约束修改都是放宽限制,不影响现有数据。
2026-06-23 23:41:49 +08:00
4dd4138197 Merge pull request 'fix: 减少自定义情景中数据库表对 prompt 的约束' (#210) from fix/database into main
All checks were successful
Deploy / deploy (push) Successful in 33s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/210
2026-06-23 22:58:56 +08:00
hhs
5ef059eb40 fix: 减少自定义情景中数据库表对 prompt 的约束 2026-06-23 22:58:02 +08:00
9523fe6650 Merge pull request 'fix: 修复无法正常自定义情景的问题' (#209) from fix/frontend into main
All checks were successful
Deploy / deploy (push) Successful in 39s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/209
2026-06-23 11:11:08 +08:00
hhs
a219934f4b fix: 修复无法正常自定义情景的问题 2026-06-23 11:10:37 +08:00
72a8d4803e Merge pull request 'fix: 修复前端语法问题' (#208) from fix/frontend into main
All checks were successful
Deploy / deploy (push) Successful in 28s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/208
2026-06-22 21:59:14 +08:00
hhs
dcf53a3783 fix: 修复前端语法问题 2026-06-22 21:58:49 +08:00
b5ec6551bf Merge pull request 'styles:美化样式' (#207) from feat/styles into main
Some checks failed
Deploy / deploy (push) Failing after 17s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/207
2026-06-22 19:32:47 +08:00
79d321c78a fix: 自建情景在聊天面板与系统消息中正确显示图标和名称 2026-06-22 16:02:02 +08:00
57bd8c72b8 Merge pull request 'feat: v2 版本' (#206) from v2 into main
All checks were successful
Deploy / deploy (push) Successful in 18s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/206
2026-06-22 16:01:27 +08:00
d640b5b41b Merge pull request 'docs: 优化首页 README 文档' (#205) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 18s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/205
2026-06-22 15:47:39 +08:00
f1b9e4fcc1 fix: 系统消息最大宽度限制与居中对齐 2026-06-22 15:08:26 +08:00
18aa5f1949 refactor: 场景配置与弹窗样式改用 CSS 变量,适配白色主题 2026-06-22 15:03:58 +08:00
a16628119a Merge pull request 'docs: 优化首页部分内容重复' (#204) from docs/readme into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/204
2026-06-22 14:50:11 +08:00
hhs
3888941f85 docs: 优化首页部分内容重复 2026-06-22 14:49:52 +08:00
0a816049fa Merge pull request 'docs: 优化首页 README 项目部署指南' (#203) from docs/readme into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/203
2026-06-22 14:43:34 +08:00
140a178993 fix: 侧栏收起时使用 visibility 隐藏避免闪烁 2026-06-22 14:43:04 +08:00
hhs
31c1720c00 docs: 优化首页 README 项目部署指南 2026-06-22 14:42:56 +08:00
c6093c01c1 Merge pull request 'docs: 优化首页 README' (#202) from docs/readme into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/202
2026-06-22 14:37:50 +08:00
hhs
3267f35edb docs: 优化首页 README 2026-06-22 14:37:32 +08:00
a309c269e0 Merge pull request 'docs: 完善首页 README 文件' (#201) from docs/readme into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/201
2026-06-22 14:36:20 +08:00
hhs
22930f0080 docs: 完善首页 README 文件 2026-06-22 14:35:53 +08:00
c00b8a83d9 Merge pull request 'fix: 优化侧栏弹出延迟问题' (#200) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 28s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/200
2026-06-22 13:17:18 +08:00
36007babcb Merge pull request 'fix: 优化侧栏弹出延迟问题' (#199) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/199
2026-06-22 13:16:59 +08:00
hhs
f01c84d89e fix: 优化侧栏弹出延迟问题 2026-06-22 13:16:31 +08:00
c3808ff38c Merge pull request 'fix: 修复“设置”侧栏双重展开的问题' (#198) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 29s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/198
2026-06-22 13:03:04 +08:00
2740bea755 Merge pull request 'fix: 修复“设置”侧栏双重展开的问题' (#197) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/197
2026-06-22 13:02:43 +08:00
hhs
69b9694cfe fix: 修复“设置”侧栏双重展开的问题 2026-06-22 13:02:05 +08:00
70ee212ea2 Merge pull request 'fix: 修复前端侧边栏无法平滑弹出的问题' (#196) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 28s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/196
2026-06-22 12:55:43 +08:00
f08f547e76 Merge pull request 'fix: 修复前端侧边栏无法平滑弹出的问题' (#195) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/195
2026-06-22 12:55:24 +08:00
hhs
03a19befe1 fix: 修复前端侧边栏无法平滑弹出的问题 2026-06-22 12:53:09 +08:00
c4b68b77d1 Merge pull request 'fix: 修复空会话返回 null 导致前端崩溃的问题' (#194) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 43s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/194
2026-06-22 12:40:37 +08:00
f2a883bda8 Merge pull request 'fix: 修复空会话返回 null 导致前端崩溃的问题' (#193) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/193
2026-06-22 12:40:18 +08:00
hhs
4cc0713459 fix: 修复空会话返回 null 导致前端崩溃的问题
## 问题描述
当用户创建新会话但未发送任何消息就切换到其他会话时,前端控制台报错:
"TypeError: Cannot read properties of null (reading 'map')"

根本原因:Go 后端未初始化的切片序列化为 JSON 时会变成 `null` 而非 `[]`,
前端尝试对 `null` 调用 `.map()` 导致崩溃。

## 修复方案
采用多层防御策略,同时修复后端和前端:

### 后端修复(确保 API 契约正确)
1. message_pg.go:93 - 将 `var messages []StoredMessage` 改为
   `messages := make([]StoredMessage, 0)`,确保空结果序列化为 `[]`
2. conversation.go - 在两个响应路径(PG 查询 + 内存回退)添加防御性 nil 检查

### 前端防御(多层保护)
1. useSessionList.ts - 在 loadMessages 和 loadSessions 中添加 null 合并操作
   `(res.data.messages || [])` 确保即使后端退化也不会崩溃

## 影响范围
- 所有空会话(新建后未发送消息的对话)现在可以正常切换
- API 响应符合 JSON 最佳实践(数组字段永远是 `[]` 而非 `null`)
2026-06-22 12:39:37 +08:00
hhs
2180751a9b fix: 修复前端侧边栏双重展开动画问题
移除条件渲染和 CSS animation 的冲突,改用纯 CSS transition 控制显示/隐藏。

**问题根源:**
- 组件使用 `if (!open) return null;` 条件渲染,导致挂载时立即出现
- CSS 同时使用 `animation: slideInLeft`,触发从 -100% 的二次滑入
- 两个独立显示机制叠加,造成双重动画闪烁

**修复方案:**
- SessionSidebar: 移除条件渲染,始终保持 DOM 存在,通过动态 className 控制状态
- App.css: 移除 `animation` 和 `@keyframes`,改用 `transition: transform`
- 新增 `.sidebar--collapsed` / `.sidebar--open` 类控制 `translateX`
- 新增 `.sidebar-backdrop--visible` 类控制背景遮罩淡入淡出
- 添加 `pointer-events: none` 确保隐藏状态不响应交互

影响范围:
- frontend/src/components/SessionSidebar/index.tsx
- frontend/src/App.css
2026-06-22 12:05:54 +08:00
964c5c967e Merge pull request 'fix: 修复前端侧边栏弹出异常问题' (#192) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 47s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/192
2026-06-22 11:20:01 +08:00
dcd08d031c Merge pull request 'fix: 修复前端侧边栏弹出异常问题' (#191) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/191
2026-06-22 11:19:34 +08:00
hhs
6bca4e0d47 fix: 修复前端侧边栏弹出异常问题 2026-06-22 11:15:41 +08:00
881c3f9853 Merge pull request 'fix: 修复前端 TypeScript 编译错误' (#190) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 23s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/190
2026-06-22 00:30:00 +08:00
2881941be6 Merge pull request 'fix: 修复前端 TypeScript 编译错误' (#189) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/189
2026-06-22 00:29:42 +08:00
hhs
1b99d24fc1 fix: 修复前端 TypeScript 编译错误
- 移除 App.tsx 中未使用的 activeScenario 变量
- 修复 ChatPanel 中 ReactNode 的 type-only import
- 移除 ChatPanel 中未使用的 onSelectScenario 参数

修复 CI 构建失败问题
2026-06-22 00:28:59 +08:00
f9e68b96b9 Merge pull request 'refactor: 重构视频面板 UI,情景选择改为卡片网格,控制按钮改为 SVG 图标' (#188) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 11s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/188
2026-06-22 00:14:16 +08:00
c58c6b59a5 Merge pull request 'refactor: 重构视频面板 UI,情景选择改为卡片网格,控制按钮改为 SVG 图标' (#187) from feat/stylechange into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/187
2026-06-22 00:13:44 +08:00
e24fb2cc19 refactor: 将情景图标和设备选择器的 emoji 统一替换为 SVG 图标
- 新增 scenarioIcons.tsx 提取情景图标渲染逻辑
- 情景芯片、设备选择器、警告提示等 emoji 全部改为 Feather-style SVG
- 调整 CSS 确保 SVG 图标居中对齐
2026-06-21 23:50:27 +08:00
7a705744d0 Merge pull request 'feat: 实现日志追踪链路' (#186) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 29s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/186
2026-06-21 23:24:17 +08:00
583c33727a Merge pull request 'feat: 实现日志追踪链路' (#185) from feature/log-track into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/185
2026-06-21 23:23:22 +08:00
6cbabb63bb refactor: 情景选择从卡片网格改为横向滚动芯片条
- 将 scenario-picker 网格布局改为 scenario-strip 横向滚动芯片条
- 新增「新建情景」快捷入口芯片(+按钮)
- 新增 scenario.createChip 三语言翻译键
2026-06-21 23:23:22 +08:00
b4fbf8625b Merge pull request 'feat: 实现摄像头和麦克风设备选择功能' (#184) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 39s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/184
2026-06-21 23:21:35 +08:00
hhs
239f8f9877 docs: 同步限流和鉴权文档的日志实现说明
- 更新限流文档:limiter 内部使用 trace.FromContext 自动记录日志
- 更新鉴权文档:Redis 降级策略使用 trace-aware 日志
- 引用 13-日志追踪.md 作为详细说明
- 移除过时的手动 logger.Log 调用示例
2026-06-21 23:19:11 +08:00
hhs
d430e6e5b2 docs: 补充存储层和限流器日志实现说明
添加 PostgreSQL、Redis、限流器三个模块的日志实现文档:
- PostgreSQL 4 个 repository 的日志策略和代码示例
- Redis 会话存储、缓存装饰器、限流器的日志级别选择
- 存储层日志查询示例(数据库错误、Redis 降级)
- 更新架构图,添加存储层节点
2026-06-21 23:11:17 +08:00
hhs
edc66625ba feat: 为 redis_bucket.go 添加限流日志
添加 trace-aware 日志:
- Error: Redis 限流检查失败(fail-open 降级)
- Warn: 限流触发,记录 key 和 retry_after_sec
2026-06-21 23:08:00 +08:00
hhs
9dce107a84 feat: 升级 cached_user.go 日志为 trace-aware
替换 logger.Log 为 trace.FromContext(ctx):
- SaveRefreshToken: Redis 缓存写入失败降级日志
- FindRefreshToken: Redis 缓存读取失败降级日志
- DeleteRefreshToken: Redis 缓存删除失败降级日志
- DeleteUserRefreshTokens: Redis 批量删除失败降级日志
2026-06-21 23:07:49 +08:00
hhs
1c7dd708a0 feat: 升级 session/redis.go 日志为 trace-aware
替换 logger.Log 为 trace.FromContext(ctx):
- CreateWithID: session 创建日志
- Get: session 获取日志(新增错误日志)
- UpdateConfig: 配置更新日志
- UpdateTitle: 标题更新日志
- GetHistory: 无效历史条目警告日志
- Destroy: session 销毁日志
2026-06-21 23:07:02 +08:00
hhs
d0e4bdaeec feat: 为 PostgreSQL store 层添加 trace 日志
为 4 个 PostgreSQL repository 添加 trace-aware 日志:
- session_pg.go: Save/Find/Update/Delete 操作日志
- user_pg.go: 用户 CRUD 和 refresh token 管理日志
- message_pg.go: 消息存储和查询日志
- user_scenario_repository.go: 自定义情景 CRUD 日志

日志策略:
- Error: 数据库操作失败
- Debug: 操作成功(避免 Info 级别噪音)
- NotFound (ErrNoRows) 不记录错误日志
2026-06-21 23:06:50 +08:00
949a707e0f refactor: 重构视频面板 UI,情景选择改为卡片网格,控制按钮改为 SVG 图标
- 情景选择从 ChatPanel header 下拉菜单迁移至视频面板下方卡片网格
- 摄像头/麦克风按钮改为圆形 SVG 图标按钮
- 识别/中断/停止按钮改为药丸形状 SVG 图标按钮
- 设备选择器改为水平布局,优化空间利用
- 移除 ChatPanel 空状态中重复的情景卡片
- .gitignore 忽略 .claudian/ 和笔记目录
2026-06-21 22:58:02 +08:00
hhs
55f7f183a7 docs: 完善日志追踪文档 2026-06-21 22:51:14 +08:00
hhs
6b4b033df3 feat: 日志级别优化(Phase 7)
- nodes_stt.go: STT 识别开始降为 Debug
- nodes_history.go: 历史组装完成降为 Debug
- nodes_tts.go: TTS 流中断降为 Debug
- 保持关键里程碑为 Info:query completed、tts synthesis started/completed
- 中间步骤详情降为 Debug,减少生产环境日志噪音
2026-06-21 22:37:23 +08:00
hhs
76d331c885 feat: Eino nodes 迁移到 trace 包(Phase 6.2)
- nodes_stt.go 使用 trace.FromContext 替换 logger.Log
- nodes_history.go 使用 trace.FromContext
- nodes_tts.go 使用 trace.FromContext
- nodes_done.go 使用 trace.FromContext
- 移除所有 nodes 中的 request_id 手动字段(自动附加)
- 所有日志消息改为英文
2026-06-21 22:36:15 +08:00
hhs
ad700743ef feat: Eino adapter 和 callback 迁移到 trace 包(Phase 6.1)
- 移除 adapter.go 中的 ctxKeySessionID 定义
- 移除 callback.go 中的 ctxKeyRequestID 定义
- 统一使用 trace.WithSessionID/WithRequestID
- adapter.go 使用 trace.FromContext 替换 logger.Log
- callback.go 使用 trace.FromContext
- 移除双重日志,SetActiveRequest 失败直接返回错误
- 更新测试文件导入 trace 包
2026-06-21 22:33:49 +08:00
hhs
6ab4776e08 feat: WebSocket 追踪集成(Phase 5)
- 在 WebSocket 升级后生成连接级 trace_id
- 为每个 query 注入 request_id 到 context
- 使用 trace.FromContext 替换所有 logger.Log
- trace_id 贯穿整个 WebSocket 生命周期
- 自动附加 trace_id/session_id/request_id 到所有日志
2026-06-21 22:31:45 +08:00
hhs
c17798ec67 feat: 添加 HTTP 请求日志中间件(Phase 4)
实现内容:
- 创建 trace/gin_logger.go,实现 GinLogger 和 GinRecovery 中间件
- GinLogger 自动记录所有 HTTP 请求的 method/path/status/latency/client_ip
- GinRecovery 使用 zap 记录 panic 恢复信息,替代 gin.Recovery()
- 修改 main.go 注册三层中间件:TraceMiddleware -> GinLogger -> GinRecovery
- 所有日志自动附加 trace_id 和 request_id 字段

日志示例:
{
  "level": "info",
  "ts": "2026-06-21T22:26:07.610+0800",
  "msg": "request completed",
  "trace_id": "01KVN964KSY6BFDKB9S932NXB8",
  "request_id": "01KVN964KSY6BFDKB9S932NXB8",
  "method": "GET",
  "path": "/api/health",
  "status": 200,
  "latency_ms": 0,
  "client_ip": "::1"
}

测试:已验证健康检查接口日志正常输出
2026-06-21 22:26:43 +08:00
hhs
8a43f4406a feat: 添加限流触发日志 2026-06-21 22:16:38 +08:00
hhs
065673fae2 feat: 添加会话接口日志(创建/销毁) 2026-06-21 22:15:53 +08:00
hhs
03c27e7790 feat: 添加对话接口错误日志(CRUD 操作) 2026-06-21 22:14:41 +08:00
hhs
0a59173476 feat: 添加 JWT 鉴权拒绝日志 2026-06-21 22:13:15 +08:00
hhs
939e43acd0 feat: 添加认证接口日志(登录/注册/刷新/登出) 2026-06-21 22:12:21 +08:00
hhs
87c3e7a8dd refactor: OpenAI TTS 使用 trace.FromContext 支持追踪 2026-06-21 22:08:27 +08:00
hhs
d6e9555a97 refactor: MiMo TTS 使用 trace.FromContext 支持追踪 2026-06-21 22:08:00 +08:00
hhs
9c763ec12a feat: Redis 会话历史解析错误日志截断保护 2026-06-21 22:06:10 +08:00
hhs
1bab02ae84 feat: OpenAI TTS 错误日志截断敏感文本 2026-06-21 22:05:35 +08:00
hhs
89d7b7c17c feat: MiMo TTS 错误日志截断敏感文本 2026-06-21 22:05:13 +08:00
hhs
34d498510e feat: STT 敏感文本降级为 Debug 并截断保护 2026-06-21 22:04:50 +08:00
hhs
104b28efd3 feat: 实现字符串截断工具函数用于敏感内容保护 2026-06-21 22:02:08 +08:00
hhs
a02a8bc374 feat: 实现 Gin trace 中间件生成并注入 trace ID 2026-06-21 22:01:50 +08:00
hhs
0d99d06f06 test: 验证 Eino Graph 正确传递 context.Value 2026-06-21 22:01:27 +08:00
hhs
8b18953010 feat: 实现 context-aware logger 自动附加 trace 字段 2026-06-21 22:00:32 +08:00
hhs
0108ef2064 feat: 实现 trace context 注入与提取功能 2026-06-21 22:00:12 +08:00
hhs
e0dc8272a5 feat: 实现并发安全的 ULID trace ID 生成器 2026-06-21 21:59:43 +08:00
hhs
11c3955bd6 feat: 添加 ULID 依赖用于 trace ID 生成 2026-06-21 21:59:16 +08:00
hhs
a66ab764d9 docs: 统一文档风格 2026-06-21 19:31:17 +08:00
hhs
51117b43f6 fix: 修复 handler_test.go 中 ServeWS 缺少参数 2026-06-21 18:44:03 +08:00
hhs
492fb06c08 merge: 合并 feat/tokentime 到 develop,解决 limiter/scenarioRepo 参数冲突 2026-06-21 17:42:48 +08:00
6967dd7b2e feat: 实现摄像头和麦克风设备选择功能
- 新增 useDeviceList hook,枚举音视频输入设备并监听热插拔
- CameraManager/MicManager 的 startCamera/startMic 支持可选 deviceId 参数
- useVisionSession 集成设备选择:授权后自动枚举、切换设备时热重启
- 连接后显示设备下拉选择器,未连接时隐藏(避免未授权时空列表)
- SessionConfig 新增 cameraDeviceId/micDeviceId 持久化到 localStorage
2026-06-21 16:49:48 +08:00
eea5c07eaa refactor: 删除实时分析/按需识别/纯聊天三种模式切换
三种模式仅为前端 UI 区分,后端无感知,实际使用中价值不大:
- 纯聊天模式名不副实(VAD 仍附带摄像头帧)
- 实时分析增加 API 成本且体验不佳
- 按需识别已是默认且最自然的交互方式

删除后保留按需识别行为:语音带图 + 手动识别按钮,UI 更简洁。
2026-06-21 16:24:09 +08:00
1079e22699 feat: 实现自建情景功能
## 功能概述
- 用户可创建、编辑、删除自定义情景
- 支持自定义情景名称、图标、描述、Prompt、首句引导
- 完整的权限隔离,用户只能管理自己的情景
- 深度集成 Eino 框架,动态加载自建情景 Prompt

## 后端实现
### 数据库
- 新增 user_scenarios 表
- 支持用户配额(最多 20 个)
- 字段验证:description 可选,prompt 最小 10 字符

### API
- GET /api/scenarios - 获取用户情景列表
- POST /api/scenarios - 创建情景
- GET /api/scenarios/:id - 获取详情
- PATCH /api/scenarios/:id - 更新情景
- DELETE /api/scenarios/:id - 删除情景

### Eino 集成
- PipelineState 添加 UserID 字段
- nodes_history 动态加载用户自建情景
- GetScenarioPrompt 支持自建情景优先级

## 前端实现
### 组件
- CreateScenarioModal - 创建情景对话框
- EditScenarioModal - 编辑情景对话框
- ConfigPanel 改造 - 分组显示系统预置和自建情景

### Hook
- useScenarios - 合并系统和自建情景,提供 CRUD 接口

### 国际化
- 中文、英文、日文翻译支持

## 问题修复
- 修复 CORS 问题:使用 Vite 代理
- 统一验证规则:description 可选,prompt 最小 10 字符
- 修复数据库约束:使用 NULLIF 处理空字符串

## 文件变更
新增文件: 13 个
修改文件: 14 个

详见文档: docs/自建情景功能完整文档.md
2026-06-21 15:38:28 +08:00
ad5d90e344 Merge pull request 'docs: 重构文档结构,规范编号并整合冗余内容' (#182) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 26s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/182
2026-06-21 14:52:06 +08:00
c094fe0867 Merge pull request 'docs: 重构文档结构,规范编号并整合冗余内容' (#181) from docs/sync-docs into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/181
2026-06-21 14:50:59 +08:00
hhs
032de796c8 docs: 重构文档结构,规范编号并整合冗余内容
## 主要变更

### 文档重构(减少 1199 行,-23%)
- 01-架构设计.md: 503→369 行 (-27%),删除 DDL/配置示例,精简鉴权/存储描述
- 02-接口文档.md: 1313→570 行 (-57%),删除 Go 接口/Orchestrator 实现/配置管理
- 07-成本控制.md: 65→59 行 (-9%),代码块替换为文件引用

### 文档编号规范化
- 08-功能创意.md → 删除(内容整合到 README.md "功能扩展方向")
- 10-Eino框架与编排设计.md → 08-Eino框架与编排设计.md
- 情景切换.md → 09-情景切换.md
- 12-鉴权体系.md → 10-鉴权体系.md
- 13-令牌桶限流.md → 11-令牌桶限流.md

### 交叉引用更新
- 01-架构设计.md: 更新对鉴权体系/令牌桶限流的引用为新编号
- README.md: 更新文档索引表、推荐阅读顺序、新增功能扩展方向

### 删除过时文档
- 09-技术名词解释.md(内容已整合到 03-技术选型.md)
- 10-Eino重构方案.md(历史记录,已完成)
- 11-Eino框架技术文档.md(已合并到 08)
- 情景切换功能完整文档.md(已规范化为 09)

## 重构原则
- 架构文档聚焦系统结构,移除实现细节
- 接口文档保留纯契约,删除内部实现
- 编号连续(01-11),语义清晰
- 通过交叉引用连接相关文档,避免重复
2026-06-21 14:48:03 +08:00
8b4acb3ce7 Merge pull request 'docs: 统一更新配置文件路径引用' (#180) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 28s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/180
2026-06-21 13:32:20 +08:00
9e5f691056 Merge pull request 'docs: 统一更新配置文件路径引用' (#179) from feature/ratelimit into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/179
2026-06-21 13:32:03 +08:00
hhs
03127aa01a docs: 统一更新配置文件路径引用
将文档中所有 config.yaml 的路径引用更新为 backend/config/config.yaml,
与实际的配置文件组织结构保持一致。

变更文件:
- README.md: 更新配置文件位置说明
- docs/02-接口文档.md: 更新配置文件路径
- docs/13-令牌桶限流设计.md: 更新配置示例路径
2026-06-21 13:31:11 +08:00
hhs
99fcd6bc29 fix: 修复 Dockerfile 配置文件路径 2026-06-21 13:27:29 +08:00
d0f5f5c94d Merge pull request 'feat: 集成限流器到服务' (#178) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 4s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/178
2026-06-21 13:25:06 +08:00
361c5d07d3 Merge pull request 'docs: 规范文档对跨域处理的介绍' (#177) from feature/ratelimit into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/177
2026-06-21 13:11:16 +08:00
hhs
910e71b6f0 docs: 规范文档对跨域处理的介绍 2026-06-21 13:10:48 +08:00
19645be04e Merge pull request 'feat: 集成限流器到服务' (#176) from feature/ratelimit into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/176
2026-06-21 12:48:04 +08:00
hhs
d53265755a refactor: 移除生产环境配置文件中 CORS 配置 2026-06-21 12:39:33 +08:00
hhs
252cdcc8e7 refactor: 开发环境使用与生产环境相同的 AI 模型配置
- config.dev.yaml 补全 AI 服务配置(provider、model、endpoint)
- 保持开发环境超时时间较长(方便调试)
- 确保开发和生产环境使用相同的模型,避免环境差异导致的问题
2026-06-21 12:24:20 +08:00
hhs
515d7ae034 docs: 完善配置环境切换机制
- docker-compose.yml 添加 APP_ENV=prod 强制生产环境
- 更新 .env.example 的 APP_ENV 注释说明(开发/生产差异)
- CLAUDE.md 补充"配置环境切换"章节,说明本地开发和生产部署的配置切换方式
- 明确配置优先级:环境变量 > config.{env}.yaml > config.yaml > 默认值

配置差异:
- 开发环境:debug 日志、关闭限流、允许所有 CORS
- 生产环境:info/json 日志、启用限流、严格 CORS 白名单

验证通过:
- 本地开发:默认加载 config.dev.yaml(env=dev, debug 日志)
- 生产配置:APP_ENV=prod 加载 config.prod.yaml(env=prod, info 日志)
2026-06-21 01:14:29 +08:00
hhs
311330cea1 refactor: 重组配置文件到 config 目录
- 创建 backend/config/ 目录统一管理配置文件
- 移动 config.yaml 到 config/config.yaml
- 新增 config.dev.yaml 开发环境配置(debug 日志、关闭限流)
- 新增 config.prod.yaml 生产环境配置(info 日志、启用限流、严格 CORS)
- 更新配置加载逻辑,优先从 config/ 目录读取,兼容旧路径
- 更新 .gitignore,仅排除 .env,配置文件纳入版本控制
2026-06-21 01:07:47 +08:00
hhs
7adf81c6e5 feat: 集成限流器到服务
- main.go 初始化限流器(根据 Redis 可用性选择内存/Redis 实现)
- WebSocket handler 添加 query 消息限流(按 userID)
- Auth API 添加登录/注册限流(按 IP)
- refresh 和 logout 不限流(避免影响正常用户操作)
- 修复所有测试(传递 nil limiter 参数)
- 所有测试通过(包括 ws 和 api 集成测试)
2026-06-21 00:00:24 +08:00
hhs
ea00939c13 feat: 实现 Gin 限流中间件
- Middleware 函数返回 Gin 中间件
- keyFunc 参数支持灵活提取限流 key(IP/用户 ID 等)
- 限流触发时返回 HTTP 429 + Retry-After header
- 支持 nil limiter(跳过限流)和空 key(跳过限流)
- 完整单元测试(6 个测试用例,全部通过)
- 测试覆盖:允许、拒绝、nil limiter、空 key、keyFunc、Retry-After 舍入
2026-06-20 23:56:44 +08:00
hhs
b74fb3564d feat: 实现 Redis 令牌桶限流器
- RedisLimiter 基于 Lua 脚本保证原子性
- Lua 脚本实现完整令牌桶算法(填充、消耗、TTL)
- fail-open 策略:Redis 故障时允许请求通过
- FormatKey 辅助函数格式化限流 key
- 完整单元测试(10 个测试用例,使用 miniredis)
- 测试覆盖:首次请求、耗尽、不同用户、补充、容量上限、零速率、TTL、故障降级
2026-06-20 23:55:37 +08:00
hhs
3b6226394b feat: 实现内存令牌桶限流器
- Limiter 接口定义(Allow + Stop 方法)
- TokenBucket 实现(容量、填充速率、并发安全)
- MemoryLimiter 管理多用户令牌桶
- 后台 goroutine 定期清理不活跃桶(10 分钟)
- 完整单元测试覆盖(11 个测试用例,全部通过)
- 边界情况处理(rate=0、capacity=0、并发安全)
2026-06-20 23:53:36 +08:00
hhs
a6df8c9131 feat: 添加限流配置层
- Config 结构体新增 RateLimit 字段
- 新增 RateLimitConfig 和 BucketConfig 结构体定义
- config.yaml 新增 ratelimit 配置段(默认关闭)
- 设置默认值:query(10/0.2)、login(5/0.1)、register(3/0.05)
2026-06-20 23:51:45 +08:00
023c834074 Merge pull request 'feat:部署修改' (#175) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 28s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/175
2026-06-20 23:38:29 +08:00
3e00e39e8a Merge pull request 'feat:部署修改' (#174) from feat/tokentime into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/174
2026-06-20 23:36:33 +08:00
9ad486d117 feat:部署修改 2026-06-20 23:35:13 +08:00
6af26ffc91 Merge pull request 'feat: 优化情景切换功能' (#173) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 29s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/173
2026-06-20 23:28:37 +08:00
7ee6918015 Merge pull request 'feat: 优化情景切换功能' (#172) from feat/tokentime into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/172
2026-06-20 23:26:28 +08:00
2fb23852ef feat: 优化情景切换功能 2026-06-20 23:24:32 +08:00
ed95ce56a8 fix: 关闭摄像头后文本输入不再发送图像数据
修复了用户关闭摄像头后,使用文本输入时 AI 仍会分析黑色画面的问题。

变更:
- sendTextMessage: 只在摄像头开启时捕获画面
- 待发消息队列: 根据摄像头状态决定是否携带画面

效果:
- 摄像头关闭时纯文本对话,不提及画面
- 节省 token 消耗和网络带宽

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-20 22:21:03 +08:00
e151c5b665 Merge pull request 'fix: apk 使用阿里云镜像 mirrors.aliyun.com 加速下载' (#171) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 42s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/171
2026-06-20 22:04:05 +08:00
dffd8bd4a5 Merge pull request 'fix: apk 使用阿里云镜像 mirrors.aliyun.com 加速下载' (#170) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/170
2026-06-20 22:03:49 +08:00
hhs
3140a660e4 fix: apk 使用阿里云镜像 mirrors.aliyun.com 加速下载 2026-06-20 22:02:30 +08:00
23e0e22a12 Merge pull request 'fix: CI 环境安装 rsync + docker-cli + docker-cli-compose' (#169) from develop into v2
Some checks failed
Deploy / deploy (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/169
2026-06-20 21:47:57 +08:00
ba7c5ed5ea Merge pull request 'fix: CI 环境安装 rsync + docker-cli + docker-cli-compose' (#168) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/168
2026-06-20 21:47:45 +08:00
hhs
cc7a333b6d fix: CI 环境安装 rsync + docker-cli + docker-cli-compose 2026-06-20 21:46:25 +08:00
f8af2f0ccc Merge pull request 'fix: git clone 前清理残留的 /tmp/camtalk-deploy 目录' (#167) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 3m57s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/167
2026-06-20 21:39:17 +08:00
d6059ae397 Merge pull request 'fix: git clone 前清理残留的 /tmp/camtalk-deploy 目录' (#166) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/166
2026-06-20 21:39:02 +08:00
hhs
d4398e55fe fix: git clone 前清理残留的 /tmp/camtalk-deploy 目录 2026-06-20 21:38:30 +08:00
f5440c9e2d Merge pull request 'fix: 添加 rsync 安装,避免首次部署时 command not found' (#165) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 4m2s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/165
2026-06-20 21:33:39 +08:00
190f40908f Merge pull request 'fix: 添加 rsync 安装,避免首次部署时 command not found' (#164) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/164
2026-06-20 21:33:23 +08:00
hhs
ddb90c58ef fix: 添加 rsync 安装,避免首次部署时 command not found 2026-06-20 21:32:57 +08:00
340c26b7a6 Merge pull request 'fix: 统一使用 /opt/camtalk/.env 路径,act_runner 挂载宿主机目录后直接读取' (#163) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 1s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/163
2026-06-20 21:32:05 +08:00
ca07188eda Merge pull request 'fix: 统一使用 /opt/camtalk/.env 路径,act_runner 挂载宿主机目录后直接读取' (#162) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/162
2026-06-20 21:31:49 +08:00
hhs
15a50f855b fix: 统一使用 /opt/camtalk/.env 路径,act_runner 挂载宿主机目录后直接读取 2026-06-20 21:30:21 +08:00
e37dc7074d chore: 首页代码仓库链接改为 Gitea 地址 2026-06-20 21:12:06 +08:00
f86c2560cd Merge pull request 'fix: env_file 使用绝对路径 /root/camtalk/.env,确保 CI 容器内 docker compose 正确解析' (#161) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 12s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/161
2026-06-20 20:58:51 +08:00
90c5ad7724 Merge pull request 'fix: env_file 使用绝对路径 /root/camtalk/.env,确保 CI 容器内 docker compose 正确解析' (#160) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/160
2026-06-20 20:58:34 +08:00
hhs
45e2c4f37b fix: env_file 使用绝对路径 /root/camtalk/.env,确保 CI 容器内 docker compose 正确解析 2026-06-20 20:57:34 +08:00
b7d0edb6da Merge pull request 'docs: 更新架构设计和令牌桶限流相关文档' (#159) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 13s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/159
2026-06-20 20:54:51 +08:00
159bf278e6 Merge pull request 'docs: 添加令牌桶限流模块设计文档' (#158) from feature/ratelimite into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/158
2026-06-20 20:36:21 +08:00
542126df90 Merge pull request 'docs:更新文档' (#157) from feat/historychat into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/157
2026-06-20 20:20:59 +08:00
e56600e408 docs:更新文档 2026-06-20 20:17:16 +08:00
04215e7a53 Merge pull request 'feat: 优化对话历史功能、增加首页登录页面' (#156) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 32s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/156
2026-06-20 20:11:34 +08:00
9ba4eb4825 Merge pull request 'feat: 优化对话历史功能、增加首页登录页面' (#155) from feat/historychat into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/155
2026-06-20 20:06:39 +08:00
bc7eb7409c Merge pull request 'fix: 通过 docker run -v 桥接同步 /opt/camtalk/.env,每次部署自动获取最新配置' (#154) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 2m52s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/154
2026-06-20 20:05:28 +08:00
0a8a20f3f8 Merge pull request 'fix: 通过 docker run -v 桥接同步 /opt/camtalk/.env,每次部署自动获取最新配置' (#153) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/153
2026-06-20 20:05:11 +08:00
hhs
e20f58b4ae fix: 通过 docker run -v 桥接同步 /opt/camtalk/.env,每次部署自动获取最新配置 2026-06-20 20:04:46 +08:00
ab07e01adf feat: 优化对话历史功能 2026-06-20 19:57:36 +08:00
0447fdacac Merge pull request 'fix: git fetch 改用公网 URL 而非 origin,避免 Docker 内网地址不可达导致卡死' (#152) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 15s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/152
2026-06-20 19:48:04 +08:00
12cf59f1e8 Merge pull request 'fix: git fetch 改用公网 URL 而非 origin,避免 Docker 内网地址不可达导致卡死' (#150) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/150
2026-06-20 19:47:28 +08:00
hhs
4f735e6d29 fix: git fetch 改用公网 URL 而非 origin,避免 Docker 内网地址不可达导致卡死 2026-06-20 19:45:56 +08:00
32aea44f3b Merge pull request 'fix: env_file 改为相对路径,解决 CI 容器无法读取 /opt/camtalk/.env 的问题' (#149) from develop into v2
Some checks failed
Deploy / deploy (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/149
2026-06-20 19:31:24 +08:00
4d51bcefb0 Merge pull request 'fix: env_file 改为相对路径,解决 CI 容器无法读取 /opt/camtalk/.env 的问题' (#148) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/148
2026-06-20 19:31:10 +08:00
hhs
fe5ac20a1f fix: env_file 改为相对路径,解决 CI 容器无法读取 /opt/camtalk/.env 的问题 2026-06-20 19:29:12 +08:00
2a13ee9c89 Merge pull request 'fix: 精简 deploy.yml,checkout + ref + verify 三步确保拉取正确分支' (#147) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 2m12s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/147
2026-06-20 17:50:25 +08:00
df35ff73b5 Merge pull request 'fix: 精简 deploy.yml,checkout + ref + verify 三步确保拉取正确分支' (#146) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/146
2026-06-20 17:50:08 +08:00
hhs
92e6636d05 fix: 精简 deploy.yml,checkout + ref + verify 三步确保拉取正确分支 2026-06-20 17:49:36 +08:00
hhs
12e2d24f37 fix: checkout 后显式 git fetch 切换触发分支,避免始终部署 main 代码;Dockerfile 添加 BuildKit cache mount 加速构建 2026-06-20 17:43:32 +08:00
ae100c4a75 Merge pull request 'fix: 恢复 checkout action 拉取代码,仅添加 ref 参数指定触发分支' (#145) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 2m10s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/145
2026-06-20 17:35:46 +08:00
4a5905307c Merge pull request 'fix: 恢复 checkout action 拉取代码,仅添加 ref 参数指定触发分支' (#144) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/144
2026-06-20 17:35:30 +08:00
2a7d4c74d4 feat:增加首页登录页面 2026-06-20 17:35:14 +08:00
hhs
6bb4773e07 fix: 恢复 checkout action 拉取代码,仅添加 ref 参数指定触发分支 2026-06-20 17:33:59 +08:00
97e125234b Merge pull request 'fix: 完善鉴权模块,修复 CI/CD 中 dubious ownership 错误' (#143) from develop into v2
Some checks failed
Deploy / deploy (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/143
2026-06-20 17:24:06 +08:00
hhs
39073b7673 fix: 添加 git safe.directory 配置,修复 CI/CD 中 dubious ownership 错误 2026-06-20 17:21:40 +08:00
hhs
85d47c1fc4 docs: 添加令牌桶限流模块设计文档
- 新建 docs/13-令牌桶限流设计.md,覆盖算法原理、内存/Redis 双实现、配置设计、接入点、测试用例等
- 更新 docs/01-架构设计.md Rate Limiter 模块行链接至新文档
- 更新 docs/README.md 文档索引和推荐阅读顺序
2026-06-20 17:08:23 +08:00
hhs
4ff8cec312 docs: 添加鉴权体系设计文档,更新认证相关文档
- 新增 12-鉴权体系设计.md,详细描述 JWT 双 token 轮转认证机制
- 更新架构设计文档,补充认证设计章节的安全机制和配置说明
- 更新接口文档,补充 Refresh Token Rotation 安全机制和前端集成示例
- 更新文档索引,添加新文档的推荐阅读顺序
2026-06-20 16:51:03 +08:00
3572b867c0 Merge pull request 'feat: 完善鉴权模块' (#142) from fix/auth into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/142
2026-06-20 16:34:47 +08:00
dfed964f76 Merge pull request 'fix: CI/CD 改用 git clone 直接拉取代码,避免 rsync 同步失败' (#141) from develop into v2
Some checks failed
Deploy / deploy (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/141
2026-06-20 16:25:23 +08:00
hhs
4651ae185b fix: CI/CD 改用 git clone 直接拉取代码,避免 rsync 同步失败 2026-06-20 16:23:59 +08:00
1c65433d40 Merge pull request 'fix: checkout 指定触发分支,避免始终拉取 main 代码' (#140) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 16s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/140
2026-06-20 15:51:41 +08:00
hhs
590c4592ea fix: checkout 指定触发分支,避免始终拉取 main 代码 2026-06-20 15:51:14 +08:00
15cd157f45 Merge pull request 'feat: 对话场景优化' (#139) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 24s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/139
2026-06-20 15:18:38 +08:00
hhs
970f10a274 fix: 修复前端 lint 错误 — 移除 ref 模式,用 useEffect 依赖注入回调
- api.ts: 移除 catch 中未使用的 err 变量
- auth.tsx: 去掉 useRef 模式,直接在 useEffect 中注册回调并声明依赖
2026-06-20 15:07:09 +08:00
hhs
d78cdb509d feat: 前端 401 拦截器 — access token 过期时自动刷新并重试原请求
api.ts:
- 增加 setAuthCallbacks 回调注入机制(避免 api/auth 循环依赖)
- request() 对非公开路径自动附加 Authorization header
- 收到 401 时自动触发 refresh token 刷新,成功后重试原请求
- 并发保护:多个 401 只触发一次 refresh,其余等待同一 Promise
- refreshTokenDirect 内部方法绕过 401 拦截避免递归

auth.tsx:
- 用 ref 保存 persistAuth/scheduleRefresh 最新引用(避免闭包陈旧)
- 初始化时调用 setAuthCallbacks 注入认证回调
2026-06-20 14:57:27 +08:00
hhs
ea70d2efc6 feat: main.go 注入 Redis 到 CachedUserRepository,启用 refresh token 二级缓存
- 将 rdb 变量提升到外层作用域,供 session 和 auth 共用
- Redis 启用时用 CachedUserRepository 包装 userRepo
- backfillTTL 使用 cfg.Auth.RefreshTTL 与 token 实际过期时间一致
2026-06-20 14:54:08 +08:00
hhs
a2a28a9f56 feat: CachedUserRepository — refresh token 的 Redis 缓存装饰器
- 装饰 UserRepository,仅缓存 refresh token 相关操作
- SaveRefreshToken: Write-Through,先写 DB 再写 Redis(SET + SADD)
- FindRefreshToken: Read-Through,Redis miss 时查 DB 并回填
- DeleteRefreshToken: 双删 DB + Redis
- DeleteUserRefreshTokens: 通过 Redis Set 批量清理缓存后删 DB
- Redis 操作失败时降级到纯 DB,不阻断主流程
2026-06-20 14:53:10 +08:00
hhs
898e30b526 feat: Refresh Token 复用检测 — 检测到已 rotation 的 token 被复用时吊销用户全部 refresh token
- Refresh 方法中 FindRefreshToken 返回 not found 时,检查 JWT 是否有效
- JWT 有效但 DB 不存在 → 判定为复用,调用 DeleteUserRefreshTokens 吊销该用户所有 token
- 增加 TestRefresh_ReuseDetectedRevokesAllTokens 测试覆盖复用场景
2026-06-20 14:51:34 +08:00
hhs
87c54b80c0 feat: JWT Claims 增加 TokenType 字段,ValidateAccess/ValidateRefresh 区分校验类型
- Claims 新增 TokenType 字段("access" / "refresh")
- GeneratePair 为 access/refresh token 分别设置 token_type
- ValidateAccess 校验后检查 token_type == "access"
- ValidateRefresh 校验后检查 token_type == "refresh"
- 增加 token 类型交叉校验测试
2026-06-20 14:50:19 +08:00
7288f443f1 Merge pull request 'feat: 对话场景优化' (#138) from fea/models into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/138
2026-06-20 14:49:22 +08:00
488b1e62ba docs: 更新前端三态会话模型,新增结束视频保留对话设计文档 2026-06-20 14:47:25 +08:00
50c84fce88 fix: 修复新对话首次打字出现两个输入框
根因:sendTextMessage 断开状态时立即 setMessages 添加用户消息到 UI,
同时 connect() 触发连接成功后 flush useEffect 再次 setMessages 添加同一条消息,
导致消息重复、ChatPanel 渲染出两个输入区域。

修复:断开状态时只入队不立即显示,由 flush 统一处理。
2026-06-20 14:32:40 +08:00
ae09fba400 fix: 修复聊天输入框 flex 布局导致重复渲染
- chat-input__field 增加 min-width: 0 防止 flex 项溢出
- chat-input__send 增加 flex-shrink: 0 防止按钮被挤压换行
2026-06-20 14:24:52 +08:00
f8a79a3b0b fix: 修复文字对话态误显「正在初始化语音检测」
ChatPanel VAD 初始化提示条件增加 isCameraOn 判断,
视频结束后不再显示语音检测初始化提示。
2026-06-20 14:04:20 +08:00
ae9a27c300 feat: 结束视频后保留对话,支持继续文字聊天
- useVisionSession 新增 stopVideo 回调(停媒体流,保持连接和消息)
- App.tsx 控制区从二态改为三态(initial/video/textOnly)
- 文字对话态显示「重新开始视频」和「结束会话」按钮
- 新增 CSS 样式(video-ended-hint、btn--outline)
- 新增 i18n key(stopVideo/endSession/resumeVideo/video.ended)
- 新增设计方案文档
2026-06-20 13:48:26 +08:00
1f7cd407a8 Merge pull request 'fix: CI 环境安装 rsync 依赖' (#137) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 1m56s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/137
2026-06-20 13:37:27 +08:00
6f81212997 Merge pull request 'fix: CI 环境安装 rsync 依赖' (#136) from fix/config into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/136
2026-06-20 13:37:14 +08:00
hhs
402ad8c949 fix: CI 环境安装 rsync 依赖 2026-06-20 13:36:44 +08:00
52015fa6c6 Merge pull request 'fix: CI/CD 部署固定到 /root/camtalk 目录,确保 compose 能正确管理容器生命周期' (#135) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 4s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/135
2026-06-20 13:34:50 +08:00
97706ea197 Merge pull request 'fix: CI/CD 部署固定到 /root/camtalk 目录,确保 compose 能正确管理容器生命周期' (#134) from fix/config into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/134
2026-06-20 13:33:59 +08:00
hhs
d53de4f33f fix: CI/CD 部署固定到 /root/camtalk 目录,确保 compose 能正确管理容器生命周期 2026-06-20 13:33:05 +08:00
3dc2015a91 docs: 更新技术文档,同步 Eino 重构和默认 provider 变更
- 架构设计:更新为 Eino Graph 声明式编排,增加三级存储架构说明
- 接口文档:AI 编排器章节重写为 Eino Graph,更新 LLM 服务接口
- 技术选型:新增 Eino 框架选型章节,修正 STT/LLM/TTS 默认方案
- 语音交互:Pipeline 描述改为 Eino Graph
- 成本控制:模型引用修正为 qwen3-vl-plus
- 技术名词解释:新增 Eino 框架相关术语
- README:增加 10/11/12 Eino 文档索引
- 10-Eino重构方案:状态更新为已实施
- CLAUDE.md:同步所有变更
2026-06-20 13:24:53 +08:00
de78d60959 Merge pull request 'refactor: 用 Eino 框架重构编排层' (#133) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 2m4s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/133
2026-06-20 12:01:03 +08:00
93d5a90495 Merge pull request 'ci: deploy 工作流增加 main 分支触发' (#132) from fix/redis-use into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/132
2026-06-20 11:59:35 +08:00
hhs
5910e02a66 ci: deploy 工作流增加 main 分支触发 2026-06-20 11:59:07 +08:00
88ee1e2548 Merge pull request 'fix: 修复 TieredManager 创建 session 时 L1/L2 ID 不一致导致 Redis 写入失败的问题' (#131) from fix/redis-use into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/131
2026-06-20 11:55:43 +08:00
hhs
b8fd6a9330 fix: 修复 TieredManager 创建 session 时 L1/L2 ID 不一致导致 Redis 写入失败的问题 2026-06-20 11:54:21 +08:00
c2c9a6d392 Merge pull request 'fix: 统一 docker-compose 环境变量注入方式,修复 postgres/redis 配置读取问题' (#130) from feature/add-redis-docker into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/130
2026-06-20 00:28:18 +08:00
795ace75d6 Merge pull request 'fix: 统一 docker-compose 环境变量注入方式,修复 postgres/redis 配置读取问题' (#129) from feature/add-redis-docker into v2
All checks were successful
Deploy / deploy (push) Successful in 16s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/129
2026-06-20 00:27:20 +08:00
hhs
dd846876ba fix: 统一 docker-compose 环境变量注入方式,修复 postgres/redis 配置读取问题
- postgres 和 redis 服务添加 env_file,统一从 /opt/camtalk/.env 读取配置
- 移除 postgres 中多余的 ${POSTGRES_USER/PASSWORD} 透传(env_file 已直接注入)
- 移除 backend 中多余的 CAMTALK_STORAGE_DSN 和 CAMTALK_REDIS_PASSWORD 透传
- postgres healthcheck 改用 $$POSTGRES_USER(容器内 shell 变量)替代 ${POSTGRES_USER}(compose 变量替换)
- 根本原因:environment 中的 ${VAR} 在 compose 解析时读取 shell 环境,为空时会覆盖 env_file 的值
2026-06-20 00:26:03 +08:00
898d1474ec Merge pull request 'feat: 重构用eino框架' (#128) from fea/newcode into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/128
2026-06-19 23:28:55 +08:00
582f68f68b Merge pull request 'fix: 修复 Dockerfile 运行阶段缺少 config.yaml 导致容器启动 panic' (#127) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 15s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/127
2026-06-19 23:24:45 +08:00
7d328f2552 Merge pull request 'fix: 修复 Dockerfile 运行阶段缺少 config.yaml 导致容器启动 panic' (#126) from feature/add-redis-docker into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/126
2026-06-19 23:24:29 +08:00
eb1b90445f fix: 修复 Eino Graph 类型不匹配和多模态消息问题
- 添加 msg2str 转换节点解决 ChatModel 输出 *schema.Message 与 Splitter 期望 string 的类型不匹配
- 将多模态图片内容从 system 消息移到 user 消息(DashScope API 仅支持 user/tool 角色的多模态内容)
- 修复 Content 和 UserInputMultiContent 不能同时设置的问题
- Splitter 输出改为 StreamReader[string](单句),TTS 改为 TransformableLambda 流式消费
- 修复 .env 中 PostgreSQL DSN 和 Redis ADDR 的 http:// 前缀问题
2026-06-19 23:24:01 +08:00
hhs
e5537aaa1e fix: 修复 Dockerfile 运行阶段缺少 config.yaml 导致容器启动 panic 2026-06-19 23:23:30 +08:00
556666c046 Merge pull request 'fix: 修复 Redis 密码为空时容器启动失败的问题' (#125) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 21s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/125
2026-06-19 23:19:23 +08:00
4235c5cae0 Merge pull request 'fix: 修复 Redis 密码为空时容器启动失败的问题' (#124) from feature/add-redis-docker into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/124
2026-06-19 23:19:01 +08:00
hhs
41f393c09d fix: 修复 Redis 密码为空时容器启动失败的问题
当 CAMTALK_REDIS_PASSWORD 为空时,--requirepass 缺少参数导致
Redis 解析失败退出,healthcheck 无法执行,后端依赖启动失败。

改为 shell 脚本根据密码是否为空动态决定是否启用认证。
2026-06-19 23:18:30 +08:00
04dac0b673 Merge pull request 'feat: 添加 Redis Docker 容器编排及三级存储架构' (#123) from develop into v2
Some checks failed
Deploy / deploy (push) Failing after 5m12s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/123
2026-06-19 23:09:39 +08:00
4c7430d6f4 Merge pull request 'feat: 添加 Redis Docker 容器编排及三级存储架构' (#122) from feature/add-redis-docker into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/122
2026-06-19 23:08:12 +08:00
hhs
765cb34019 feat: 添加 Redis Docker 容器编排及三级存储架构
- docker-compose.yml 新增 Redis 服务及 camtalk-net 桥接网络
- 实现 TieredManager 三级存储(L1 内存 → L2 Redis → L3 PostgreSQL)
- config.go 新增 Redis/Persistence 配置类型及环境变量绑定
- 修复 CAMTALK_REDIS_ADDR 环境变量未被 Viper 绑定的问题
- .env.example 更新为三级存储配置并标注开发/部署地址差异
2026-06-19 22:53:06 +08:00
9576884619 docs: 新增 Eino 框架技术文档和重构实施记录
- docs/11-Eino框架技术文档.md: 框架简介、技术选型对比、核心概念(Lambda/Graph/ChatModel/StreamReader/Callback/State)、CamTalk Graph 设计、目录结构、注意事项
- docs/12-Eino重构实施记录.md: 重构背景、架构变更、四阶段实施详情、代码统计、遗留事项

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-19 22:09:07 +08:00
4ffd84510e refactor: 清理旧编排代码,新增 eino 包单元测试
删除旧代码:
- orchestrator/pipeline.go: 旧 STT→LLM→TTS 手写 goroutine 管道
- orchestrator/splitter.go: 旧句子切分器
- orchestrator/pipeline_test.go: 旧 Pipeline 测试
- ai/llm/openai.go: 旧 LLM OpenAI 实现(被 eino-ext ChatModel 替代)
- ai/llm/openai_test.go: 旧 LLM 测试

保留的接口和工具:
- orchestrator/orchestrator.go: Orchestrator 接口(ws/handler 依赖)
- orchestrator/sender.go: Sender 接口(eino/callback 依赖)
- ai/llm/llm.go: Request/Chunk/TokenUsage 类型定义
- ai/llm/prompt.go: BuildSystemPrompt(eino/nodes_history 依赖)
- ai/llm/scenarios.go: GetScenarioPrompt(eino/nodes_history 依赖)

新增测试:
- eino/graph_test.go: 13 个测试覆盖类型构建、State 并发安全、
  Context 注入、延迟计算、接口实现检查等

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-19 22:04:39 +08:00
4b731b5ac0 feat: 实现 Eino Graph 构建与 Orchestrator 适配器,切换 main.go
- graph.go: 构建 Graph 拓扑 START→STT→History→ChatModel→Splitter→TTS→Done→END
  - 创建 eino-ext ChatModel 对接 DashScope OpenAI 兼容接口
  - 统一使用值类型(PipelineInput/PipelineOutput)
  - Callback 在运行时通过 Stream option 传入
- adapter.go: EinoOrchestrator 实现 orchestrator.Orchestrator 接口
  - 解码 base64 音频/图片,注入 context 值
  - 调用 Graph.Stream() 触发惰性执行并消费输出
  - 追加用户/助手消息到历史
- main.go: 移除旧 llmService + orchestrator.New()
  替换为 eino.NewPipelineGraph() + eino.NewEinoOrchestrator()
- 各节点统一使用值类型,State 传递请求元数据

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-19 21:58:17 +08:00
fd5c7712f8 feat: 引入 Eino 框架并实现 AI 编排层基础设施与节点
- 引入 cloudwego/eino v0.9.9 和 eino-ext/components/model/openai v0.1.13
- 新增 internal/eino/ 包:
  - types.go: PipelineInput/Output、STTOutput、TokenUsage 类型定义
  - state.go: PipelineState 跨节点状态收集(线程安全)
  - callback.go: ChatModel OnEndWithStreamOutput 回调,逐 token 推送 llm_chunk
  - nodes_stt.go: STT Lambda,支持文本/语音输入模式
  - nodes_history.go: 历史组装 Lambda,含多模态图片支持
  - nodes_splitter.go: 句子分割 Transform Lambda
  - nodes_tts.go: TTS Lambda,逐句合成推送音频
  - nodes_done.go: Done Lambda,发送 llm_done 并追加历史

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-19 21:49:28 +08:00
f38fbf0527 Merge pull request 'refactor: 统一配置文件系统并重构项目设计文档' (#121) from develop into v2
All checks were successful
Deploy / deploy (push) Successful in 11s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/121
2026-06-19 19:46:06 +08:00
b10a508356 Merge pull request 'refactor: 统一配置文件系统并重构项目设计文档' (#120) from docs/complete-doc into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/120
2026-06-19 19:43:02 +08:00
hhs
c0b4eeda46 refactor: 统一配置文件系统
- 补全 backend/config.yaml 所有非敏感配置项并添加中文注释
- 重写 config.go:Load(workDir) 显式传参,BindEnv 绑定敏感字段,删除 AutomaticEnv
- setDefaults 默认值与 config.yaml 保持一致(mimo/dashscope)
- backend/.env.example 重写为纯敏感信息模板
- .env 固定在 /opt/camtalk/.env,docker-compose 通过绝对路径加载
- deploy.sh 统一使用 --env-file,移除硬编码 IP
- deploy.yml 删除 CI 写入 .env 的步骤
- Dockerfile 移除 COPY config.yaml
- 修复 deploy.yml 中 POSTGRES_PASSWORD 的 &{{ 拼写错误
2026-06-19 18:41:35 +08:00
hhs
e6481e0faa docs: 添加 Eino 重构代码计划文档 2026-06-19 17:45:41 +08:00
hhs
a04275cc76 docs: 按功能模块重构文档结构
- 新建 01-架构设计.md:合并项目概述+系统架构+持久化设计,含 Mermaid 架构图、模块图、时序图、ER 图、部署图
- 新建 02-接口文档.md:合并接口文档+持久化 API+用户模块 API,统一格式去重
- 重编号 03~09,去掉状态标注,规划中功能标记为待实现
- 删除 PLAN_BACKEND.md、PLAN_USER_MODULE.md 及冗余文档
2026-06-19 15:31:52 +08:00
hhs
dca37f3e48 docs: 同步文档与代码实现状态
- 02-系统架构: Redis/PostgreSQL 标注已实现,模块表新增 Auth/Store/Migrations,更新表设计和前端组件
- 03-接口文档: config 新增 scenario 字段,Manager 接口补全 UpdateTitle/ListByUser,配置结构体同步,扩展接口替换为实际 Repository
- 04-技术选型: 持久化层标注已实现
- 06-语音交互: TTS Voice 更正为 mimo_default
- 11-持久化与用户系统设计: 所有 Phase 标记完成
- PLAN_BACKEND/PLAN_USER_MODULE: 标记完成状态
- README: 新增实现状态总览,补充文档索引
2026-06-19 14:58:43 +08:00
hhs
16302af7d2 docs: 添加 Eino 框架参考文档 2026-06-19 14:35:06 +08:00
hhs
5d8cacf16d docs: 更新 README 文档
Some checks are pending
Deploy / deploy (push) Waiting to run
2026-06-14 23:59:00 +08:00
hhs
dbfdf3c3e5 docs: README
All checks were successful
Deploy / deploy (push) Successful in 19s
2026-06-14 23:58:41 +08:00
54454ae2d7 Merge pull request '部署优化' (#119) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 35s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/119
2026-06-14 21:33:01 +08:00
720c2e2b5b Merge pull request 'feat:部署优化' (#118) from feat/build into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/118
2026-06-14 21:29:04 +08:00
c2bae4e3b7 feat:部署优化 2026-06-14 21:28:18 +08:00
838493145f Merge pull request '部署' (#117) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 17s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/117
2026-06-14 21:25:32 +08:00
943f36e9ec Merge pull request '部署优化' (#116) from feat/build into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/116
2026-06-14 21:25:06 +08:00
083ede8b6d feat: 部署优化 2026-06-14 21:24:11 +08:00
14242eb896 Merge pull request 'fix: 修复 Docker Compose 环境下 PostgreSQL 连接失败' (#115) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 1m30s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/115
2026-06-14 20:50:31 +08:00
6db5bef519 Merge pull request 'fix: 修复 Docker Compose 环境下 PostgreSQL 连接失败' (#114) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/114
2026-06-14 20:48:31 +08:00
hhs
1ddd75d34f fix: 修复 Docker Compose 环境下 PostgreSQL 连接失败 2026-06-14 20:47:00 +08:00
7ea6c0d4f2 Merge pull request '版本升级' (#113) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 1m36s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/113
2026-06-14 20:44:51 +08:00
97357013bb Merge pull request 'feat:添加对话情景功能' (#112) from feat/optimize into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/112
2026-06-14 20:41:50 +08:00
057a226610 Merge remote-tracking branch 'origin/develop' into feat/optimize
# Conflicts:
#	frontend/src/lib/errors.ts
2026-06-14 20:40:05 +08:00
b60e513317 feat:添加对话情景功能 2026-06-14 20:32:30 +08:00
638a387550 Merge pull request 'fix: 修复 ChatPanel vadError 属性类型不匹配导致构建失败' (#110) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 1m30s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/110
2026-06-14 20:11:19 +08:00
9dc9025d10 Merge pull request 'fix: 修复 ChatPanel vadError 属性类型不匹配导致构建失败' (#109) from fix/frontend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/109
2026-06-14 20:10:59 +08:00
hhs
92e47718bc fix: 修复 ChatPanel vadError 属性类型不匹配导致构建失败 2026-06-14 20:10:31 +08:00
6ce36e3ed8 Merge pull request 'fix: 更换 PostgreSQL 镜像源' (#108) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 1m22s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/108
2026-06-14 20:00:02 +08:00
6032608e49 Merge pull request 'fix: 更换 PostgreSQL 镜像源' (#107) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/107
2026-06-14 19:59:41 +08:00
hhs
668da26e23 fix: 更换 PostgreSQL 镜像源 2026-06-14 19:58:22 +08:00
44f4f3ce1a Merge pull request 'feat:语音对话优化' (#106) from feat/onload into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/106
2026-06-14 19:55:02 +08:00
c3118dd72a feat:语音对话优化 2026-06-14 19:54:22 +08:00
f22243a2d2 Merge pull request 'fix: 更换 PostgreSQL 镜像源为轩辕镜像,修复阿里云镜像拉取失败' (#105) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 14s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/105
2026-06-14 19:52:04 +08:00
8ff3c519ac Merge pull request 'fix: 更换 PostgreSQL 镜像源为轩辕镜像,修复阿里云镜像拉取失败' (#104) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/104
2026-06-14 19:51:50 +08:00
hhs
ffc9b4a320 fix: 更换 PostgreSQL 镜像源为轩辕镜像,修复阿里云镜像拉取失败 2026-06-14 19:51:15 +08:00
eba0607f94 Merge pull request 'Merge pull request 'fix: 移除未使用的变量和导入,修复 tsc 编译错误'' (#103) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 17s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/103
2026-06-14 19:44:35 +08:00
e89585e453 Merge pull request 'fix: 移除未使用的变量和导入,修复 tsc 编译错误' (#102) from fix/actions into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/102
2026-06-14 19:43:59 +08:00
hhs
e11934cf70 fix: 移除未使用的变量和导入,修复 tsc 编译错误 2026-06-14 19:37:16 +08:00
1f19435951 Merge pull request 'Merge pull request 'feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本'' (#101) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 1m32s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/101
2026-06-14 19:25:37 +08:00
16105f5ac6 Merge pull request 'feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本' (#99) from feature/reload into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/99
2026-06-14 19:01:45 +08:00
hhs
386c918591 feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本 2026-06-14 19:00:06 +08:00
hhs
966b30218a feat: 实现 sessions 持久化,支持会话恢复 2026-06-14 18:54:18 +08:00
d724a8e9f9 Merge pull request 'feat: 对接 PostgreSQL 存储层' (#98) from build/backend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/98
2026-06-14 18:46:05 +08:00
8360e31328 Merge pull request 'feat/talk' (#97) from feat/talk into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/97
2026-06-14 18:44:47 +08:00
1109ab41f7 refactor: 将退出登录功能从顶部导航栏移至设置面板 2026-06-14 18:42:47 +08:00
hhs
542720695d feat: 对接 PostgreSQL 存储层 2026-06-14 18:40:37 +08:00
8de956719c feat: 实现前端认证系统,WebSocket 连接携带 JWT token
- 新增 AuthContext/AuthProvider:登录/注册状态管理,JWT 自动刷新
- 新增 AuthPage 组件:登录/注册表单,支持模式切换和前端校验
- 新增 api.ts:封装 auth REST API 客户端(register/login/refresh/logout)
- WebSocket 连接时拼接 ?token=<jwt>,重连自动携带
- useVisionSession 接受 accessToken 参数并透传
- App.tsx 包裹 AuthProvider,未登录时显示登录页
- storage.ts 新增 token/user 的 localStorage 存储
- i18n 新增中/英/日三语 auth 翻译
- App.css 新增 auth 页面和用户 badge 样式
2026-06-14 18:38:45 +08:00
a1ae7740f1 Merge pull request 'feat: 构建用户模块,实现用户对话历史持久化,完善接口文档' (#96) from build/backend into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/96
2026-06-14 18:08:14 +08:00
hhs
b57adf153b docs: 补充用户模块 REST API 接口契约
- PLAN_USER_MODULE.md 新增「前端 API 接口参考」章节
- 03-接口文档.md 同步认证接口、对话接口、WebSocket 认证变更
- 新增错误码 USERNAME_TAKEN / INVALID_CREDENTIALS / INVALID_TOKEN / INVALID_INPUT
- 数据模型补充 User / ConversationSummary / StoredMessage 及对应 TypeScript 类型
- 配置结构体补充 AuthConfig(JWTSecret / AccessTTL / RefreshTTL)
2026-06-14 18:05:50 +08:00
hhs
c3a32ce276 feat: Phase 8.5 — ConversationSummary 查询优化,支持 SQL 聚合获取消息统计 2026-06-14 17:58:42 +08:00
hhs
7b745018c5 feat: Phase 8.4 — 实现 LoadSession 和 LoadSessionFromDB,支持从 PostgreSQL 恢复会话到内存 2026-06-14 17:56:33 +08:00
hhs
f4515ce5e4 feat: Phase 8.3 — Session Manager 注入 MessageRepository,AppendMessage 启用 Write-Through 2026-06-14 17:54:33 +08:00
hhs
96f4bc7abb feat: Phase 8.2 — 实现 PostgreSQL MessageRepository 2026-06-14 17:53:14 +08:00
hhs
dae5722945 feat: Phase 8.1 — 定义 MessageRepository 接口和消息表迁移脚本 2026-06-14 17:52:21 +08:00
hhs
3c5c4943e8 feat: Phase 7.7 — 编写 WS 认证测试
- TestWS_AuthMissingToken: 无 token 返回 401
- TestWS_AuthInvalidToken: 无效 token 返回 401
- TestWS_AuthExpiredToken: 过期 token 返回 401
- TestWS_AuthValidToken: 有效 token 成功连接
- TestWS_AuthConversationIDResume: conversation_id 恢复已有对话
- TestWS_AuthConversationIDNotFound: 不存在的 conversation_id 返回 401
- TestWS_AuthConversationIDOwnership: 非 owner 访问返回 401
2026-06-14 17:48:47 +08:00
hhs
80c6b1b56e feat: Phase 7.3 — WS conversation_id 处理
- ?conversation_id=xxx 存在时校验 session 归属(UserID 匹配)
- 校验失败返回 401 SESSION_NOT_FOUND
- 校验通过则复用已有 session;否则创建新 session
2026-06-14 17:47:34 +08:00
hhs
905b56640e feat: Phase 7.2 — WS 连接 JWT 认证
- 从 ?token=xxx 查询参数提取 access_token
- 校验失败返回 401(missing token / invalid token)
- 校验成功后将 userID 用于创建会话
- 更新现有测试:setupTestServer 自动生成有效 token
2026-06-14 17:47:05 +08:00
hhs
2aa3c98ab6 feat: Phase 7.1 — 修改 ServeWS 签名,新增 tokenMgr 参数
- ServeWS 和 serveWS 函数新增 *auth.TokenManager 参数
- main.go 传入 tokenMgr 到 ServeWS
- handler_test.go 适配新签名
2026-06-14 17:46:11 +08:00
hhs
d01aeaba68 feat: 编写 ConversationHandler API 测试(httptest + mock SessionManager)
- List: 成功、分页参数、未认证 401
- Create: 成功、自定义配置
- Get: 成功、404 未找到、404 权限不足(隐藏信息)
- UpdateTitle: 成功、空标题校验、超长标题校验
- Delete: 成功、权限不足
- GetMessages: 成功、limit 分页、before 偏移分页、权限不足
- 共 17 个 Conversation 测试用例,全部通过
2026-06-14 17:41:28 +08:00
hhs
f902e05e31 feat: main.go 接入 ConversationHandler 路由注册 2026-06-14 17:39:39 +08:00
hhs
62baa656ee feat: 实现 ConversationHandler(对话 CRUD + 消息查询 + 权限校验 + 路由注册)
- List: GET /api/conversations 获取当前用户对话列表(分页)
- Create: POST /api/conversations 创建新对话
- Get: GET /api/conversations/:id 获取对话详情
- UpdateTitle: PATCH /api/conversations/:id 更新对话标题
- Delete: DELETE /api/conversations/:id 删除对话
- GetMessages: GET /api/conversations/:id/messages 获取消息列表
- 所有端点通过 AuthMiddleware 认证
- getSessionForUser 校验 session.UserID == claims.UserID
- 返回 404 而非 403 避免信息泄露
2026-06-14 17:39:20 +08:00
hhs
ec8555d44b feat: 编写 Session Manager 新方法测试(ListByUser 分页、UpdateTitle、自动标题生成) 2026-06-14 17:36:38 +08:00
hhs
6487a8ecab feat: 扩展 Session Manager 接口,新增 ListByUser、UpdateTitle 方法
- Manager.Create 签名新增 userID 参数
- 新增 ConversationSummary 类型和 ListByUser 分页查询
- 新增 UpdateTitle 方法
- MemoryManager 实现:ListByUser 遍历+过滤+排序,UpdateTitle,自动标题生成
- RedisManager 实现:user:{id}:sessions 索引,ListByUser 通过 SMEMBERS 查询
- AppendMessage 自动更新标题(首条 user 消息时,取前 20 字符)
- 更新 ws handler、api/session.go、orchestrator mock 的 Create 调用
2026-06-14 17:35:20 +08:00
hhs
9ff971fd89 feat: 扩展 Session 模型,新增 UserID、Title、UpdatedAt 字段 2026-06-14 17:30:59 +08:00
hhs
c01ee1d14c feat: 编写 AuthHandler API 测试(httptest + mock AuthService,覆盖全部端点) 2026-06-14 17:28:13 +08:00
hhs
9aaed88cc1 feat: main.go 接入认证服务(TokenManager + AuthService + AuthHandler 路由注册) 2026-06-14 17:26:50 +08:00
hhs
15b437043f feat: 实现 AuthHandler(注册/登录/刷新/登出端点 + 输入校验 + 路由注册) 2026-06-14 17:25:50 +08:00
hhs
fdea3aa9d1 feat: 新增认证相关错误码(USERNAME_TAKEN, INVALID_CREDENTIALS, INVALID_TOKEN, INVALID_INPUT) 2026-06-14 17:24:57 +08:00
430811ffc4 Merge pull request 'feat: 实现 Phase 3 JWT + 认证服务' (#95) from feature/user-model-phase3 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/95
2026-06-14 17:20:35 +08:00
hhs
f878988048 feat: 编写 AuthService 单元测试(mock UserRepository) 2026-06-14 17:18:22 +08:00
hhs
e455af6b2d feat: 编写 TokenManager 单元测试 2026-06-14 17:17:28 +08:00
hhs
5211d14d50 feat: 定义并实现 AuthService(注册/登录/刷新/登出) 2026-06-14 17:14:52 +08:00
hhs
5b8cbb025f feat: 实现 JWT 认证中间件(AuthMiddleware) 2026-06-14 17:13:54 +08:00
hhs
fca9108882 feat: 实现密码哈希与校验工具(bcrypt) 2026-06-14 17:07:51 +08:00
hhs
20597101f5 feat: 实现 TokenManager(JWT 签发/校验/哈希) 2026-06-14 17:07:12 +08:00
5e83531d54 Merge pull request 'feat: 定义用户模块 User 模型与 Repository 层' (#94) from feature/user-mode-phase2 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/94
2026-06-14 17:03:46 +08:00
hhs
9842b4a457 feat: 编写 UserRepository 单元测试 2026-06-14 16:56:29 +08:00
hhs
ce13e5a048 feat: 实现内存 UserRepository(测试用) 2026-06-14 16:55:44 +08:00
hhs
36827edfea feat: 实现 PostgreSQL UserRepository 2026-06-14 16:55:12 +08:00
hhs
af73e78aa3 feat: 定义 UserRepository 接口及 User 数据模型 2026-06-14 16:54:21 +08:00
hhs
c9c174f400 feat: 扩展 models 新增 User 结构体 2026-06-14 16:53:54 +08:00
c0972856c6 Merge pull request 'feat: 配置扩展与 PostgreSQL 连接池' (#93) from feature/user-model-phase1 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/93
2026-06-14 16:50:25 +08:00
hhs
0ff810a297 feat: main.go 条件初始化 PostgreSQL 连接池
- 添加 store 包导入
- 当 storage.driver 为 postgres 时创建 pgxpool
- 使用 defer 确保连接池正确关闭
- 添加日志记录连接状态
2026-06-14 16:43:07 +08:00
hhs
bf1e24453c feat: 编写用户表 schema 迁移脚本
- 新增 migrations/001_users.up.sql:创建 users 表和 refresh_tokens 表
- 新增 migrations/001_users.down.sql:回滚脚本
- users 表包含 id, username, password_hash, created_at, updated_at
- refresh_tokens 表包含 id, user_id, token_hash, expires_at, created_at
- 添加必要的索引优化查询性能
2026-06-14 16:42:12 +08:00
hhs
0edafbbf8a feat: 实现 PostgreSQL 连接池
- 新增 internal/store/db.go
- 实现 NewPostgresPool(ctx, dsn) 创建连接池
- 添加 pgx/v5 依赖
- 最大连接数设置为 10
2026-06-14 16:41:50 +08:00
hhs
4b30e67c2e feat: 扩展 Config 结构体,新增 AuthConfig 配置
- 新增 AuthConfig 结构体(JWTSecret, AccessTTL, RefreshTTL)
- 在 Config 中添加 Auth 字段
- 设置默认值:access_ttl=15分钟,refresh_ttl=10080分钟(7天)
- JWTSecret 必须通过环境变量 CAMTALK_AUTH_JWT_SECRET 设置
2026-06-14 16:40:06 +08:00
07211675e9 Merge pull request 'feat/tips' (#92) from feat/tips into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/92
2026-06-14 16:32:18 +08:00
7a0443ebcc style: 样式优化 2026-06-14 16:31:23 +08:00
fe053f94dc fix: 隐藏观察模式按钮,清理未使用的 toggleMode 引用 2026-06-14 15:56:59 +08:00
9433612c4b Merge pull request '聊天框与会话历史栏优化' (#91) from fix/language into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/91
2026-06-14 15:46:40 +08:00
92f3f45c41 feat: 新增会话历史侧边栏,支持新建/切换/删除/重命名会话
类似 ChatGPT 的侧边栏布局,消息历史持久化到 localStorage。

- 新增 SessionSummary 类型和 localStorage 读写函数
- 新增 useSessionList Hook 管理会话列表 CRUD
- 新增 SessionSidebar 组件(展开/折叠、行内重命名)
- App.tsx 协调 useSessionList 和 useVisionSession
- 消息变化时自动持久化,切换/新建会话时保存并加载
- 新增 i18n 翻译 key(中/英/日)
2026-06-14 15:44:25 +08:00
74cd6c7d2d Merge pull request 'chore: 硬编码公网 IP 地址' (#90) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 10s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/90
2026-06-14 15:29:29 +08:00
8f8e9dea5f Merge pull request 'chore: 硬编码公网 IP 地址' (#89) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/89
2026-06-14 15:29:19 +08:00
hhs
2bb72096de chore: 硬编码公网 IP 地址 2026-06-14 15:29:02 +08:00
6b99336941 Merge pull request 'chore: 部署完成回显显示公网 IP' (#88) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 8s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/88
2026-06-14 15:28:09 +08:00
ae91f877fe Merge pull request 'chore: 部署完成回显显示公网 IP' (#87) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/87
2026-06-14 15:27:57 +08:00
hhs
f23a425f1e chore: 部署完成回显显示公网 IP 2026-06-14 15:26:33 +08:00
12eec6fc0e Merge pull request 'chore: 前端端口改为 9000' (#86) from develop into main
All checks were successful
Deploy / deploy (push) Successful in 9s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/86
2026-06-14 15:24:50 +08:00
a0769956f2 Merge pull request 'chore: 前端端口改为 9000' (#85) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/85
2026-06-14 15:23:37 +08:00
hhs
0f3bf432fd chore: 前端端口改为 9000 2026-06-14 15:21:47 +08:00
0ba2f9807c Merge pull request 'ci: 添加国内镜像代理和 BuildKit 缓存优化构建速度' (#84) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 5m26s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/84
2026-06-14 15:15:13 +08:00
9379d3ef82 Merge pull request 'ci: 添加国内镜像代理和 BuildKit 缓存优化构建速度' (#83) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/83
2026-06-14 15:15:01 +08:00
hhs
cb0c747133 ci: 添加国内镜像代理和 BuildKit 缓存优化构建速度 2026-06-14 15:14:28 +08:00
91fdc007f7 Merge pull request 'ci: 简化 Docker CLI 安装,仅使用 apk' (#82) from develop into main
Some checks failed
Deploy / deploy (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/82
2026-06-14 15:11:38 +08:00
18c889b9b8 Merge pull request 'ci: 简化 Docker CLI 安装,仅使用 apk' (#81) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/81
2026-06-14 15:11:26 +08:00
hhs
ffcb11625a ci: 简化 Docker CLI 安装,仅使用 apk 2026-06-14 15:11:08 +08:00
da8818c3be fix: 隐藏顶部状态栏的 token 数显示 2026-06-14 15:10:43 +08:00
443f613fbd Merge pull request 'ci: 修复部署脚本找不到 docker 命令的问题' (#80) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 18s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/80
2026-06-14 15:09:52 +08:00
f46f8ffa5a Merge pull request 'ci: 修复部署脚本找不到 docker 命令的问题' (#79) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/79
2026-06-14 15:09:37 +08:00
hhs
159db37a60 ci: 修复部署脚本找不到 docker 命令的问题 2026-06-14 15:08:52 +08:00
15b1dd56fa feat: 聊天框与视频功能解耦,输入文字即可开始对话
之前必须点击「开始对话」连接 WebSocket 后才能使用聊天框。
现在聊天输入框始终可用,输入文字自动连接 WebSocket 并发送消息;
左侧按钮改为「开始视频通话」,单独控制摄像头/麦克风。

- ChatPanel 移除 !isConnected early return,始终显示输入框
- useVisionSession 新增待发消息队列,sendTextMessage 支持自动连接
- startSession 改为仅负责视频(摄像头 + 麦克风 + VAD)
- 修复 statusRef 在 useWebSocketManager 之前声明导致的 TDZ 错误
- 更新 i18n 翻译(新增 controls.startVideo、video.placeholder)
2026-06-14 15:07:39 +08:00
35372438a5 Merge pull request 'ci: 合并 workflow,仅保留 push main 时自动部署' (#78) from develop into main
Some checks failed
Deploy / deploy (push) Failing after 1s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/78
2026-06-14 15:02:54 +08:00
daa4d23895 Merge pull request 'ci: 合并 workflow,仅保留 push main 时自动部署' (#77) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/77
2026-06-14 15:02:41 +08:00
hhs
9e813f0a08 ci: 合并 workflow,仅保留 push main 时自动部署 2026-06-14 15:01:46 +08:00
97484735cf Merge pull request 'fix: 修复并优化部署流程' (#75) from develop into main
Some checks failed
Backend CI / ci (push) Successful in 1m50s
Deploy / verify (push) Has been skipped
Deploy / deploy (push) Failing after 1s
Frontend CI / ci (push) Successful in 2m11s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/75
2026-06-14 14:57:20 +08:00
a128ed1471 Merge pull request 'fix: 修复并优化部署流程' (#74) from fix/workflow-fix into develop
Some checks failed
Frontend CI / ci (pull_request) Has been cancelled
Backend CI / ci (pull_request) Has been cancelled
Deploy / verify (pull_request) Failing after 1s
Deploy / deploy (pull_request) Has been skipped
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/74
2026-06-14 14:56:52 +08:00
hhs
e33b5650a5 chore: 移除 deploy.sh 中未使用的 warn/error 函数 2026-06-14 14:55:17 +08:00
hhs
b358bf05b5 ci: 移除 deploy 工作流中多余的依赖安装和健康检查步骤 2026-06-14 14:54:59 +08:00
hhs
4ef1e9f08b refactor: 彻底移除 .env 依赖,配置统一由 config.yaml + 环境变量驱动 2026-06-14 14:54:55 +08:00
hhs
d32112e074 chore: 添加 .dockerignore 优化 Docker 构建上下文 2026-06-14 14:54:52 +08:00
81e50de7f6 Merge pull request 'fix: 修复 deploy 工作流缺少 Docker 导致构建失败的问题' (#73) from fix/workflow-fix into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/73
2026-06-14 14:38:04 +08:00
2f05b5fa2b Merge pull request 'Merge pull request 'fix: 修复 deploy 工作流缺少 Node.js 导致 checkout 失败的问题'' (#72) from develop into main
Some checks failed
Backend CI / ci (push) Successful in 1m49s
Deploy / verify (push) Has been skipped
Deploy / deploy (push) Failing after 3s
Frontend CI / ci (push) Successful in 1m22s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/72
2026-06-14 14:37:26 +08:00
hhs
304a45b0ad fix: 修复 deploy 工作流缺少 Docker 导致构建失败的问题 2026-06-14 14:37:00 +08:00
e95b1603c1 feat: 前端 UI 国际化,支持中/英/日三语切换
之前语言设置只影响后端 AI 行为,前端 UI 文字始终为中文。
新增轻量 i18n 模块(React Context + 翻译字典),切换语言后所有 UI 文字即时刷新。

- 新增 i18n 核心模块和 zh-CN/en-US/ja-JP 翻译文件(约 72 个 key)
- App.tsx 拆分为 App(Provider)+ AppContent(业务),确保 Hook 可访问 Context
- 所有组件通过 useI18n().t() 获取翻译文本
- errors.ts 改为接受翻译函数参数
- storage.ts saveConfig 时派发事件通知 locale 变化
2026-06-14 14:34:30 +08:00
7153a7adf6 Merge pull request 'fix: 修复 deploy 工作流缺少 Node.js 导致 checkout 失败的问题' (#71) from fix/actions into develop
Some checks failed
Backend CI / ci (pull_request) Successful in 1m52s
Deploy / verify (pull_request) Failing after 17s
Deploy / deploy (pull_request) Has been skipped
Frontend CI / ci (pull_request) Successful in 34s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/71
2026-06-14 14:27:59 +08:00
hhs
73eae8a118 fix: 修复 deploy 工作流缺少 Node.js 导致 checkout 失败的问题 2026-06-14 14:27:04 +08:00
1a6325698a Merge pull request 'chore: 删除 deepgram_test.go 测试文件' (#69) from refactor/ai-model-refact into develop
Some checks failed
Backend CI / ci (pull_request) Successful in 1m51s
Deploy / verify (pull_request) Failing after 1s
Deploy / deploy (pull_request) Has been skipped
Frontend CI / ci (pull_request) Successful in 1m10s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/69
2026-06-14 14:05:05 +08:00
hhs
4ebac835de chore: 删除 deepgram_test.go 测试文件 2026-06-14 14:04:30 +08:00
d8bbdb1744 Merge pull request 'fix: 修复语音与文字不同步的问题,改为句子级流式播放' (#67) from frontend-12 into develop
Some checks failed
Backend CI / ci (pull_request) Failing after 2m14s
Deploy / verify (pull_request) Failing after 1s
Deploy / deploy (pull_request) Has been skipped
Frontend CI / ci (pull_request) Successful in 56s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/67
2026-06-14 13:55:47 +08:00
da5c727d8c fix: 修复语音与文字不同步的问题,改为句子级流式播放
之前前端 TTSPlayer 攒齐所有音频片段后才播放,导致文字全部显示后才开始语音。
改为后端每句 TTS 发送 is_last: true,前端收到每句即加入播放队列,第一句到达即开始播放。

- 后端 Chunk 结构体新增 Final 字段,区分句子结束和整轮结束
- 前端 TTSPlayer 重写为队列式播放,onended 回调自动衔接下一句
- 同步更新接口文档和测试用例
2026-06-14 13:54:49 +08:00
7b3f4706e1 Merge pull request 'chore: 添加前后端 Dockerfile 和 DockerCompose 编排文件,添加部署脚本和工作流' (#66) from chore/add-docker into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/66
2026-06-14 13:41:23 +08:00
hhs
65718a8460 ci: 添加 Gitea Actions 部署工作流
- PR 到 main 时验证 Docker 镜像构建
- push 到 main 时自动构建并部署,含健康检查
2026-06-14 13:38:35 +08:00
hhs
84f5303632 feat: 添加部署脚本 deploy.sh
- 支持 build / up / down / restart / logs / status 命令
- 首次运行自动从 .env.example 生成 .env 模板
2026-06-14 13:38:32 +08:00
hhs
1c74fdb894 feat: 添加 Docker Compose 编排文件
- frontend 和 backend 两个服务,共享 camtalk-net 网络
- backend 通过 .env 注入 API Key,不对外暴露端口
- frontend 通过 Nginx 反代对外暴露 80 端口
2026-06-14 13:38:29 +08:00
hhs
d75cdf95c8 feat: 添加后端 Dockerfile
- 多阶段构建:golang:1.24-alpine 编译,alpine:3.20 运行
- CGO_ENABLED=0 静态链接,包含 ca-certificates 和时区数据
2026-06-14 13:38:25 +08:00
hhs
7108d2a58b feat: 添加前端 Dockerfile 和 Nginx 配置
- 多阶段构建:node:22-alpine 构建,nginx:stable-alpine 运行
- Nginx 配置:静态文件服务 + /api /ws 反代到后端
2026-06-14 13:38:22 +08:00
80e45bb8ff Merge pull request '优化无麦克风和摄像头场景的使用' (#65) from frontend-11 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/65
2026-06-14 13:19:15 +08:00
81c2b64e5f fix: 修复无摄像头/麦克风时聊天框不显示的问题
- 重构 ChatPanel 逻辑,连接后始终显示输入框
- 将欢迎消息移入消息容器内
- 确保输入框在任何状态下都可见
2026-06-14 13:17:14 +08:00
dffc8aa538 feat: 支持无摄像头/麦克风模式打字聊天
- 修改 startSession 逻辑,摄像头和麦克风为可选
- 连接成功后始终显示聊天输入框
- 更新 UI 提示文字,引导用户打字对话
2026-06-14 13:13:29 +08:00
dfccc824be Merge pull request '修复对话麦克风问题' (#64) from frontend-10 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/64
2026-06-14 13:08:46 +08:00
312e762aff fix: 修复文本输入时 LLM API 报错的问题
- 移除 contentPart.Text 的 omitempty 标签
- 确保 text 字段始终存在于 JSON 请求中
2026-06-14 13:02:07 +08:00
01026d1846 Merge pull request 'refactor: 前端 WebSocket 地址改为动态推导,添加 Vite 开发代理' (#63) from fix/frontend-proxy into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/63
2026-06-14 13:01:32 +08:00
hhs
d8630a8f26 refactor: 前端 WebSocket 地址改为动态推导,添加 Vite 开发代理
- websocket.ts: 移除硬编码 localhost:8080,基于 window.location 动态构建 WS URL
- vite.config.ts: 添加 /ws 和 /api 的 server.proxy 配置
- 同步更新 02-系统架构.md 和 03-接口文档.md 中的开发环境说明
2026-06-14 13:00:54 +08:00
41cacaa740 fix: 修复麦克风关闭后重新打开无法使用的问题
- 关闭麦克风时同时停止 VAD
- 打开麦克风时重新初始化 VAD
2026-06-14 12:54:58 +08:00
51ed6aa563 feat: 添加文字输入对话功能
- 后端 WsQuery 添加 text 字段,支持文本输入模式
- Pipeline 支持文本查询时跳过 STT 直接使用输入文本
- 前端 ChatPanel 添加文字输入框,连接状态下可用
- useVisionSession 添加 sendTextMessage 方法
- 更新接口文档,添加文本输入模式说明
2026-06-14 12:52:39 +08:00
563ca12790 Merge pull request 'fix/hard-code' (#62) from fix/hard-code into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/62
2026-06-14 12:31:55 +08:00
bfc896ed35 Merge pull request 'fix: 修复新旧 TTS 语音重叠播放的问题' (#61) from fix-redio into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/61
2026-06-14 12:31:17 +08:00
f6a2348c04 fix:修复语音对话配置 2026-06-14 12:29:46 +08:00
24a6ce5295 fix: 修复新旧 TTS 语音重叠播放的问题
- TTSPlayer 新增 stopCurrentAudio() 方法,play() 开始前先停止旧音频
- onSpeechEnd 发送新 query 前停止上一轮 TTS 播放
2026-06-14 12:01:34 +08:00
hhs
96cf009228 test: 更新测试文件适配构造函数签名变更 2026-06-14 11:55:29 +08:00
hhs
544716fe27 fix: pipeline.go TTS 输出格式/采样率从配置读取
- New() 改为接收 *config.Config 参数
- OutputFmt/SampleRate 从 config.AI.TTS 读取
2026-06-14 11:55:24 +08:00
hhs
8a966c29b7 fix: TTS 服务去除重复默认值,新增 httpClientTimeout 参数
- 去除 model/voice/endpoint 的 fallback 默认值
- HTTP Client 超时从 config 传入
2026-06-14 11:55:20 +08:00
hhs
aae0629b2d fix: LLM 服务去除重复默认值,新增 httpClientTimeout 参数
- 去除 model/endpoint 的 fallback 默认值
- HTTP Client 超时从 config 传入
2026-06-14 11:55:19 +08:00
hhs
03b3566822 fix: STT 服务去除重复默认值,新增 timeout 参数
- Deepgram/MiMo STT 超时从 config 传入
- 去除 model/endpoint 的 fallback 默认值,由 config 层保证
2026-06-14 11:55:18 +08:00
hhs
f1ce28966c fix: handler.go 心跳/版本号/CheckOrigin/历史上限改为配置驱动
- 心跳间隔和超时从 config.Server 读取
- ServerVersion 从 config.App.Version 读取
- CheckOrigin 通过 config.Server.AllowedOrigins 控制
- GetHistory limit 从 config.Session.MaxHistory 读取
2026-06-14 11:55:15 +08:00
hhs
63f8cc279d fix: main.go 使用配置值替代硬编码
- 版本号支持 -ldflags 构建时注入
- Session TTL/maxHistory 从配置读取
- 优雅关闭超时从配置读取
- AI 服务构造函数传入 HTTP Client 超时参数
2026-06-14 11:55:10 +08:00
hhs
4a7cf89f29 fix: 前端 WebSocket URL 改为环境变量或自动推导,消除 localhost 硬编码 2026-06-14 11:55:07 +08:00
hhs
22aefbc216 feat: 扩展配置结构,新增 Session/Heartbeat/Shutdown/STT Timeout/HTTP Client/TTS Output 配置项 2026-06-14 11:55:06 +08:00
c34be3a996 Merge pull request 'fix: 修复 pipeline_test.go 中 New 函数调用参数不足的问题' (#60) from test/pipeline-test into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/60
2026-06-14 11:19:42 +08:00
hhs
00abc6c1e1 fix: 修复 pipeline_test.go 中 New 函数调用参数不足的问题 2026-06-14 11:18:58 +08:00
32f3efbb7a Merge pull request 'fix:修复了模型语音问题' (#59) from fix-redio into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/59
2026-06-14 11:13:39 +08:00
eb7ebecfc0 fix:修复了模型语音问题 2026-06-14 11:12:49 +08:00
b1e71be3d9 Merge pull request 'feat:增加前端样式优化skill和mcp工具' (#58) from frontend10 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/58
2026-06-14 10:21:26 +08:00
1d938055e8 Merge pull request 'feat: 添加 AI 服务初始化时的模型配置日志' (#57) from feature/log-track into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/57
2026-06-14 10:19:43 +08:00
bf0bb83f7d feat: 增加前端样式优化skill和mcp工具 2026-06-14 10:19:31 +08:00
hhs
37af0dbe06 feat: 添加 AI 服务初始化时的模型配置日志 2026-06-14 10:18:48 +08:00
70dfa1c8f8 feat:样式美化 2026-06-14 10:18:36 +08:00
5aa383cab9 Merge pull request 'feat: 添加 MiMo TTS 语音合成服务' (#56) from feature/mimo-adapter into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/56
2026-06-14 10:06:41 +08:00
hhs
e967d89e7e feat: 添加 MiMo TTS 语音合成服务
实现基于 Xiaomi MiMo TTS API 的语音合成 provider,使用冰糖音色。
- 新增 MiMoService 实现 tts.Service 接口
- 使用 chat/completions 格式,与 MiMo STT 保持一致
- 在 main.go 添加 TTS provider 切换逻辑(mimo/xiaomi)
- 配置文件默认音色设为冰糖
- 包含完整测试覆盖(9 个测试用例)
2026-06-14 10:05:56 +08:00
b8ed49be63 Merge pull request 'docs: 跟进项目完成状况' (#55) from docs/update-docs into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/55
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-14 09:42:55 +08:00
hhs
2db0e3b0b6 docs: 编写项目 README.md
- 项目简介与三层架构图
- 技术栈、项目结构树
- 快速开始指南(前端/后端启动、配置说明)
- WebSocket 协议概览与文档索引
2026-06-14 09:15:05 +08:00
hhs
ca09a1fb72 chore: 清理配置文件并移动功能创意文档
- .gitignore 添加 .env,防止密钥泄露
- 删除误提交的 backend/.env
- TTS 模型名修正为 mimo-v2.5-tts
- 功能创意.md 移至 docs/10-功能创意.md
2026-06-14 08:53:12 +08:00
hhs
6e0c67e1cb docs: 文档与代码一致性检查与修复
- 修复心跳 Bug:应用层 ping 不更新 lastPong,60 秒后连接会被错误断开
- 03-接口文档:audio/mpeg→audio/mp3、STTConfig/TTSConfig 补充 Model 字段、
  APP_ENV 环境变量名修正、配置搜索路径补充、.env 加载说明、Vite proxy 说明修正
- 02-系统架构:补充 ConfigPanel/Toast 组件、Model Router/Rate Limiter 标注规划中、
  补充 Gin 框架、MVP 存储改为 Memory、AI 服务 provider 更新、Orchestrator 伪代码对齐
- 04-技术选型:新增 AI 服务栈选型章节(STT/LLM/TTS)、PostgreSQL 标注规划中
- 06-语音交互:VAD 参数名修正、STT 改为一次性识别描述、音频编码格式补充
- 07-视觉理解:关键帧检测代码改为 TypeScript、分辨率修正、阈值逻辑统一
- 08-成本控制:变量名修正、未实现功能标注规划中、对话历史裁剪策略补充
- CLAUDE.md:同步更新技术栈、模块结构、存储策略描述
2026-06-14 08:52:36 +08:00
a6ce6d9c4f Merge pull request '功能优化' (#54) from fix-model into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/54
2026-06-13 22:21:42 +08:00
a7d679d04f feat: 功能优化 2026-06-13 22:20:31 +08:00
c31ecc46fc Merge pull request 'fix: MiMo ASR 认证头改为 Authorization: Bearer' (#53) from fix/stt400 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/53
2026-06-13 21:58:50 +08:00
hhs
322061c387 fix: MiMo ASR 认证头改为 Authorization: Bearer
MiMo API 使用标准 Bearer Token 认证,而非自定义 api-key 头。
2026-06-13 21:57:56 +08:00
4458ee82b2 Merge pull request 'fix: 修复 MiMo ASR endpoint 路径重复拼接导致 404 的问题' (#52) from fix/stt400 into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/52
2026-06-13 21:51:52 +08:00
hhs
420871cd7e fix: 修复 MiMo ASR endpoint 路径重复拼接导致 404 的问题
config.yaml 中 endpoint 已包含 /chat/completions,而 mimo.go 会自动拼接该路径,
导致实际请求地址变为 .../chat/completions/chat/completions。

- config.yaml: endpoint 改为 base URL https://api.xiaomimimo.com/v1
- .env.example: STT endpoint 示例更新为 MiMo 地址
2026-06-13 21:51:26 +08:00
03a05b89a0 Merge pull request 'chore: 更新 STT API Key 和 endpoint' (#51) from fix/config into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/51
2026-06-13 21:44:47 +08:00
hhs
e3eb4f264b chore: 更新 STT API Key 和 endpoint 2026-06-13 21:43:43 +08:00
a404fb97f9 Merge pull request 'feat: 添加 Xiaomi MiMo ASR 语音识别提供者' (#50) from fix/redis-config into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/50
2026-06-13 21:32:39 +08:00
hhs
3d3de828fc feat: 添加 Xiaomi MiMo ASR 语音识别提供者
- 新增 MiMoService 实现 stt.Service 接口,通过 HTTP POST 调用 OpenAI 兼容的 /chat/completions 接口
- 自动将原始 PCM 数据封装为 WAV 格式(MiMo 仅支持 mp3/wav)
- 语言代码映射:zh-CN→zh、en-US→en、其他→auto
- main.go 添加 provider 选择逻辑(mimo/xiaomi → MiMo,其他 → Deepgram)
- 更新 config.yaml 使用正确的 model 名称 mimo-v2.5-asr
- 添加完整单元测试
2026-06-13 21:31:58 +08:00
4dd89be79f Merge pull request 'fix: 修复返回给前端的 totalTokens 为0的问题并提供示例 .env' (#49) from fix/config-file into develop
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/49
2026-06-13 21:12:16 +08:00
hhs
b0a7ce885e fix: 修复返回给前端的 totalTokens 为0的问题并提供示例 .env 2026-06-13 21:09:56 +08:00
213aa67e9f Merge pull request 'develop' (#48) from develop into main
Some checks failed
Backend CI / ci (push) Failing after 2m29s
Frontend CI / ci (push) Successful in 1m9s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/48
2026-06-13 20:58:46 +08:00
89386ac00d Merge pull request 'ci: 更新 Go 容器版本为 1.26.2' (#47) from fix/cicd into develop
Some checks failed
Backend CI / ci (pull_request) Failing after 3m55s
Frontend CI / ci (pull_request) Successful in 3m21s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/47
2026-06-13 20:58:33 +08:00
hhs
a666b82e64 ci: 更新 Go 容器版本为 1.26.2 2026-06-13 20:57:31 +08:00
5eeb592bf0 Merge pull request 'Merge pull request '添加计时器,麦克风,摄像头开关' (#45) from develop-frontend8 into develop 33分钟前' (#46) from develop into main
Some checks failed
Backend CI / ci (push) Failing after 1m24s
Frontend CI / ci (push) Successful in 38s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/46
2026-06-13 20:43:07 +08:00
836a61dfd2 Merge pull request '添加计时器,麦克风,摄像头开关' (#45) from develop-frontend8 into develop
Some checks failed
Frontend CI / ci (pull_request) Successful in 55s
Backend CI / ci (pull_request) Failing after 1m21s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/45
2026-06-13 20:09:03 +08:00
00f10088d0 feat: 底部控制栏添加摄像头/麦克风开关,说话时麦克风脉冲动画
- useVisionSession 新增 isCameraOn/isMicOn 状态 + toggleCamera/toggleMic
- 底部控制栏添加 📷🎤 图标按钮,开启绿色/关闭红色
- 说话时麦克风按钮绿色脉冲缩放动画(micPulse)
- 按钮样式优化:nowrap 防换行,flex-wrap 自适应

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-13 20:01:37 +08:00
188782fa66 feat: 添加 LIVE 计时器和语言按钮组(P0 优化)
- 视频区左上角显示 LIVE 绿色圆点 + 连接时长计时器
- Header 添加中文/EN/日 语言按钮组,一键切换语言
- live-badge、lang-group 样式

Co-Authored-By: Claude <noreply@anthropic.com>
2026-06-13 19:01:33 +08:00
46cf35967e Merge pull request 'docs' (#25) from docs into main
All checks were successful
Backend CI / ci (push) Successful in 1m6s
Frontend CI / ci (push) Successful in 3m33s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/25
2026-06-13 13:35:48 +08:00
d5ee6deeb9 Merge pull request 'develop' (#23) from develop into main
Some checks failed
Backend CI / ci (push) Has been cancelled
Frontend CI / ci (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/23
2026-06-13 11:46:56 +08:00
31dfabd882 Merge pull request 'develop' (#21) from develop into main
All checks were successful
Backend CI / ci (push) Successful in 1m4s
Frontend CI / ci (push) Successful in 27s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/21
2026-06-13 11:42:08 +08:00
bab67a8749 Merge pull request 'develop' (#19) from develop into main
Some checks failed
Backend CI / ci (push) Failing after 2s
Frontend CI / ci (push) Successful in 25s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/19
2026-06-13 11:39:28 +08:00
5d70bc8862 Merge pull request 'develop' (#17) from develop into main
Some checks failed
Frontend CI / ci (push) Failing after 0s
Backend CI / ci (push) Failing after 1s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/17
2026-06-13 11:31:34 +08:00
c0279e94c7 Merge pull request 'develop' (#15) from develop into main
Some checks failed
Backend CI / ci (push) Failing after 2s
Frontend CI / ci (push) Failing after 1s
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/15
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-13 11:21:08 +08:00
c68657bf30 Merge pull request 'develop' (#13) from develop into main
Some checks failed
Frontend CI / ci (push) Failing after 2s
Backend CI / ci (push) Has been cancelled
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/13
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-13 10:44:50 +08:00
5e89ab01bd Merge pull request 'develop' (#11) from develop into main
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/11
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-13 10:36:28 +08:00
60fc3ac5d1 Merge pull request 'develop' (#9) from develop into main
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/9
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-13 10:27:59 +08:00
0a0eaf7eb9 Merge pull request 'develop' (#7) from develop into main
Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/7
Reviewed-by: cfy777 <3087823110@qq.com>
2026-06-13 10:08:12 +08:00
175 changed files with 26453 additions and 4658 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=wss://api.deepgram.com/v1/listen
# 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

@@ -1,38 +0,0 @@
name: Backend CI
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
ci:
runs-on: aliyun
container: golang:1.23-alpine
defaults:
run:
working-directory: backend
steps:
- name: Setup Node.js
run: |
sed -i 's/dl-cdn.alpinelinux.org/mirrors.aliyun.com/g' /etc/apk/repositories
apk add --no-cache nodejs
working-directory: /
- name: Checkout
uses: "http://8.161.227.145:3000/huanghaosheng/checkout@releases/v4"
- name: Download Dependencies
run: go mod download
env:
GOPROXY: https://goproxy.cn,direct
- name: Vet
run: go vet ./...
- name: Build
run: go build ./cmd/server
- name: Test
run: go test ./...

View File

@@ -0,0 +1,35 @@
name: Deploy
on:
push:
branches: [main, v2]
jobs:
deploy:
runs-on: aliyun
steps:
- name: Deploy
run: |
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

View File

@@ -1,27 +0,0 @@
name: Frontend CI
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
ci:
runs-on: aliyun
container: node:22-alpine
defaults:
run:
working-directory: frontend
steps:
- name: Checkout
uses: "http://8.161.227.145:3000/huanghaosheng/checkout@releases/v4"
- name: Install Dependencies
run: npm ci
- name: Lint
run: npm run lint
- name: Type Check & Build
run: npm run build

7
.gitignore vendored
View File

@@ -22,8 +22,9 @@ Thumbs.db
# ---- Playwright MCP ----
.playwright-mcp/
# ---- 截图 ----
*.png
# ---- Obsidian ----
.obsidian/
.claudian/
修改过程笔记/
学习复盘/
docs/follow-up/

176
CLAUDE.md
View File

@@ -1,104 +1,128 @@
# CLAUDE.md
本文件为 Claude Code (claude.ai/code) 在本仓库中工作时提供指引。
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
## 项目概述
CamTalk — 多模态实时 AI 视觉对话助手(摄像头 + 麦克风 + 视觉 + 语音 AI
CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头和麦克风与 AI 交互AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应。项目目前处于设计文档阶段,源代码正在逐步构建
> **文档优先原则:** 开发前先读 `docs/` 设计文档,以文档为准;若代码与文档不一致,优先更新文档(尤其接口文档)。详细设计见 `docs/01-13` 系列文档。**注意**`docs/Eino/` 框架文档内容庞大(~75 个文件),仅在需要了解 Eino Graph/节点/Callback 等框架细节时才读取
> **文档优先原则:** 执行任何开发任务前,先读取 `docs/` 下的相关设计文档(架构、接口、技术选型等),以文档为最高依据。代码实现应与文档一致;若有偏差,优先更新文档(尤其是接口文档)。
## 常用命令
```bash
# === 前端frontend/ 目录)===
npm run dev # Vite 开发服务器http://localhost:5173代理 /ws 和 /api 到 :8080
npm run build # 生产构建tsc -b && vite build输出到 dist/
npm run lint # ESLint 代码检查
npm run preview # 预览生产构建
# === 后端backend/ 目录)===
go run ./cmd/server # 启动服务(监听 :8080启动时自动执行数据库迁移
golangci-lint run # Go 代码检查
# 后端测试
go test ./... # 单元测试
go test -tags=integration ./... # 集成测试(需要 PostgreSQL
go test -v -run TestXxx ./path/ # 运行单个测试
# === Docker 部署 ===
./deploy.sh build # 构建 Docker 镜像
./deploy.sh up # 启动服务4 容器frontend/backend/postgres/redis
./deploy.sh down # 停止服务
./deploy.sh logs # 查看日志(可加服务名:./deploy.sh logs backend
./deploy.sh status # 查看服务状态
```
## 架构
三层系统:
三层系统:前端React + Vite→ Go 网关Gin + WebSocket + Eino Graph AI 编排)→ AI 服务DashScope LLM, MiMo STT/TTS
1. **浏览器客户端**React 18 + TypeScript, Vite—— 媒体采集、边缘预处理VAD 通过 `@ricky0123/vad-web`、关键帧检测通过 ONNX Runtime Web、UI 渲染。核心 Hook`useVisionSession()`
2. **Go 网关**gorilla/websocket, Redis, Viper, Zap—— WebSocket 服务器、会话管理、模型路由、AI 编排、速率限制。每个 WebSocket 连接一个 goroutine。
3. **云端 AI 服务** —— 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 故障自动降级
**存储**冷热分离 —— Redis 存实时会话状态PostgreSQL 存对话历史和用量统计MVP 后引入。Repository 接口模式(`HistoryRepository``UsageRepository`MVP 用内存实现
**鉴权**JWT 双 token 轮转Access 120min + Refresh 7d重放攻击检测DB hash 校验Redis 缓存装饰器
## 技术栈
| 层级 | 技术 |
前端React 18 + TypeScript + ViteVAD@ricky0123/vad-webONNX Runtime国际化zh-CN / en-US / ja-JP
后端Go 1.25+, Gin, WebSocket, Viper, Zap, CloudWeGo Eino Graph
AIDashScope qwen3-vl-plus, MiMo ASR/TTS可切换 Deepgram/OpenAI TTS
存储PostgreSQL 15 + Redis 7
CI/CDGitea Actions`.gitea/workflows/deploy.yml`push main/v2 自动构建部署到自建 aliyun runner
前端测试:**暂无**package.json 无 test 脚本,无测试框架配置)
## 配置体系
配置优先级:**环境变量 > `config.{APP_ENV}.yaml` > `config.yaml` > 代码默认值**
配置文件位于 `backend/config/`
- `config.yaml` — 基础配置dev 默认值)
- `config.dev.yaml` — 开发环境覆盖(可选)
- `config.prod.yaml` — 生产环境覆盖(可选)
环境切换:`APP_ENV=dev|prod`dev 默认prod 启用限流 + 严格 CORS + Release 模式)
敏感信息API Key、JWT Secret、数据库密码**只能通过环境变量或 `.env` 文件注入**,不写入 YAML 配置文件。核心环境变量(参考 `backend/.env.example`
| 变量 | 说明 |
|------|------|
| 前端 | React 18, TypeScript, Vite, ONNX Runtime Web, @ricky0123/vad-web |
| 后端 | Go, gorilla/websocket, Redis, Viper, Zap |
| LLM | GPT-4o, Claude Sonnet |
| STT | Deepgram, FunASR自部署备选 |
| TTS | OpenAI TTS, Edge TTS免费替代 |
| 模型路由 | GPT-4o-mini 用于轻量分类 |
| `CAMTALK_AI_LLM_API_KEY` | LLM API KeyDashScope |
| `CAMTALK_AI_STT_API_KEY` | STT API KeyMiMo/Deepgram |
| `CAMTALK_AI_TTS_API_KEY` | TTS API KeyMiMo/OpenAI |
| `CAMTALK_AUTH_JWT_SECRET` | JWT 签名密钥 |
| `CAMTALK_STORAGE_DSN` | PostgreSQL 连接字符串 |
| `CAMTALK_REDIS_ADDR` | Redis 地址 |
| `CAMTALK_REDIS_PASSWORD` | Redis 密码 |
## 构建与运行命令
## 数据库迁移
```bash
# 前端
cd frontend && npm install
npm run dev # Vite 开发服务器
npm run build # 生产构建
npm run lint # ESLint 检查
npm run test # Vitest 测试
迁移 SQL 文件位于 `backend/migrations/``001_*.up.sql` 等),通过 Go `//go:embed` 嵌入二进制(见 `backend/migrations/embed.go`)。应用启动时**自动执行**未应用的迁移,无需手动运行迁移命令。迁移通过 `schema_migrations` 表追踪执行状态。
# 后端
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 ./... # 静态分析
```
回滚脚本为同目录下的 `*.down.sql` 文件,需手动执行。
基础设施Redis 为会话状态必需。PostgreSQL 为 MVP 可选(内存回退)。
## 协议与 API
## WebSocket 协议
**WebSocket**`ws://localhost:8080/ws?token=<jwt>&conversation_id=<uuid>`
- 客户端消息:`query`(图像/音频 Base64, `config`, `interrupt`, `ping`
- 服务端消息:`connected`, `stt_result`, `llm_chunk`, `llm_done`, `tts_audio`, `error`, `pong`
- 心跳:客户端 30s ping服务端 60s 超时断连;重连:指数退避 1s→30s
- 实现:`CamTalkWebSocket` 单例(`frontend/src/lib/websocket.ts`),订阅模式,自动重连
- WebSocket 地址自动从当前页面协议/主机推导,也可通过 `VITE_WS_URL` 环境变量显式指定(如 `wss://api.example.com/ws`
端点:`ws://localhost:8080/ws`
**REST API**`/api/auth/*`(注册/登录/刷新/登出),`/api/conversations/*`CRUD + 消息分页),`/api/scenarios/*`(用户自定义情景 CRUD`/api/health`
所有消息为 JSON 文本帧,统一信封格式 `{type, request_id?, timestamp?}`。完整契约见 `docs/03-接口文档.md`
**错误码**`INVALID_MESSAGE`, `SESSION_NOT_FOUND`, `RATE_LIMITED`, `IMAGE_TOO_LARGE`, `LLM_TIMEOUT`, `STT/TTS/LLM_ERROR`, `INVALID_TOKEN`
**客户端 → 服务端**`query`(图像 Base64 + 音频 Base64`config``interrupt``ping`
**服务端 → 客户端**`connected``stt_result``llm_chunk``llm_done``tts_audio``error``pong`
## 关键文件路径
**心**客户端每 30 秒 ping服务端 60 秒无 ping 断开连接。
**重连**:指数退避 + 抖动 —— 1s, 2s, 4s, 8s… 最大 30s。
**后端核**
- `backend/cmd/server/main.go` — 入口依赖注入与启动流程存储→AI 服务→Graph→路由→Server
- `backend/internal/eino/` — Eino Graph 编排层graph.go 构建、adapter.go 适配、callback.go 推送、state.go 状态、nodes_*.go 各节点实现)
- `backend/internal/session/tiered.go` — 三级会话存储TieredManager
- `backend/internal/store/` — 持久化层Repository 接口 + PG 实现 + Redis 缓存装饰器)
- Repository 模式:接口定义在 `user.go`/`session.go`/`message.go`PG 实现在 `*_pg.go`Redis 缓存装饰器在 `cached_user.go`
- `backend/internal/ws/handler.go` — WebSocket 连接管理(升级→认证→收发循环→清理)
- `backend/internal/ai/` — AI 服务抽象层llm/stt/tts 各子目录,统一 `Service` 接口)
- `backend/internal/auth/` — JWT/bcrypt/中间件
- `backend/internal/ratelimit/` — 令牌桶限流(内存/Redis 两种后端)
- `backend/migrations/` — 嵌入式 SQL 迁移文件embed.go + *.sql
## REST API辅助
- `GET /api/health` — 健康检查(版本、运行时间、活跃会话数
- `POST /api/sessions` — 创建会话可选MVP 在 WS 连接时自动创建
- `DELETE /api/sessions/{id}` — 销毁会话
## 错误码
`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 + 关键帧检测ONNX Runtime |
| `WebSocketManager` | WebSocket 连接生命周期管理 |
| `ChatPanel` | 消息展示 |
| `VideoPreview` | 摄像头画面预览 |
## 后端模块结构
| 模块 | 职责 |
|------|------|
| WebSocket Hub | 连接管理、广播/定向推送 |
| Session Manager | 会话状态、对话历史Redis + TTL |
| Model Router | 按请求选择 AI 模型(规则引擎 + 成本阈值) |
| AI Orchestrator | 并行/串行 AI 调用编排context 超时控制 |
| Rate Limiter | 按用户的令牌桶速率限制 |
**前端核心**
- `frontend/src/hooks/useVisionSession.ts` — 核心会话 Hook~500 行,编排整个采集→发送→接收→播放流程)
- `frontend/src/lib/websocket.ts` — WebSocket 客户端单例(心跳/重连/订阅模式
- `frontend/src/lib/auth.tsx` — AuthProviderJWT 自动刷新 + React Context
- `frontend/src/lib/api.ts` — REST 客户端401 拦截 + token 刷新)
- `frontend/src/lib/ttsPlayer.ts` — 流式 TTS 音频播放队列
- `frontend/src/lib/i18n/` — 国际化zh-CN / en-US / ja-JP
- `frontend/src/components/` — UI 组件LandingPage/CameraManager/MicManager/WebSocketManager/ChatPanel/SessionSidebar/ConfigPanel/VideoPreview 等)
- `frontend/vite.config.ts` — VAD 模型文件自动复制 + ONNX WASM MIME 处理 + 代理配置
## 编码规范
- **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` 字段区分
- **存储层模式**Repository 接口 + PostgreSQL 实现 + Redis 缓存装饰器(`CachedUserRepository` 包装模式)
- **CORS**:禁止后端代码/配置文件配置 CORS统一由代理层处理开发环境 Vite proxy生产环境 Nginx
- **提交信息**Conventional Commits中文描述`feat: 添加 WebSocket 心跳`
- **禁止自动 push**:除非用户明确要求
- **文档优先**:开发前先读 `docs/` 设计文档,代码与文档不一致时优先更新文档

559
README.md
View File

@@ -1,2 +1,561 @@
# CamTalk
<div align="center">
**多模态实时 AI 视觉对话助手**
用户通过摄像头和麦克风与 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/) • [在线体验](https://camtalk.goanchor.top) • [文档](docs/README.md)
</div>
---
<!-- > ⚠️ **在线体验提示**:由于演示环境使用 HTTP 协议,需配置 Chrome 允许非 HTTPS 下访问摄像头/麦克风:
>
> 1. 访问 `chrome://flags/#unsafely-treat-insecure-origin-as-secure`
> 2. 启用该选项,并在输入框填入 `http://8.161.227.145:9000`
> 3. 点击 **Relaunch** 重启浏览器
![CamTalk 界面截图](docs/pictures/1.png) -->
## ✨ 核心特性
- 🎥 **多模态理解**:摄像头视觉 + 麦克风语音双输入AI 理解完整场景
- 🗣️ **自然对话**:基于 VAD 的端到端语音交互,低延迟流式响应
- 🚀 **实时推送**LLM 文本流 + TTS 音频流并行推送,感知延迟 < 0.5 秒
- 🎭 **情景模式**:自由对话、面试官、英语老师等多场景支持
- 💾 **对话历史**:自动保存会话,支持搜索、重命名、删除、时间分组
- 🔐 **安全认证**JWT 双 token 轮转 + Refresh Token Rotation 防重放
- 📊 **三级存储**Memory → Redis → PostgreSQL 自动降级,保障可靠性
- 🌐 **国际化**:支持中文、英文、日文界面
## 🏗️ 系统架构
CamTalk 采用**三层架构**:前端轻量预处理 → Go 网关智能编排 → 云端 AI 按需调用
```mermaid
graph TB
subgraph Browser["🌐 浏览器客户端"]
UI["React UI 渲染"]
VAD["VAD 语音检测"]
Media["媒体采集"]
end
subgraph Gateway["⚙️ Go 网关 (Eino Graph)"]
WS["WebSocket Handler"]
Auth["JWT 认证"]
Session["会话管理 (三级存储)"]
Orch["AI 编排器 (7节点DAG)"]
end
subgraph AI["☁️ 云端 AI 服务"]
STT["STT (MiMo/Deepgram)"]
LLM["LLM (qwen3-vl-plus)"]
TTS["TTS (MiMo/OpenAI)"]
end
Browser <-->|"WebSocket<br/>(JWT + query/config)"| Gateway
Orch --> STT
Orch --> LLM
Orch --> TTS
```
### AI 编排流水线Eino Graph
基于 [CloudWeGo Eino](https://github.com/cloudwego/eino) 框架的声明式 7 节点 DAG
```
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/ # 🌐 浏览器客户端
│ └── src/
│ ├── components/ # UI 组件
│ │ ├── LandingPage/ # 登录着陆页 + LoginModal
│ │ ├── CameraManager/ # 摄像头流采集
│ │ ├── MicManager/ # 麦克风音频采集 + VAD
│ │ ├── WebSocketManager/ # WS 连接生命周期
│ │ ├── ChatPanel/ # 消息展示 + 流式回复
│ │ ├── SessionSidebar/ # 对话历史侧边栏
│ │ └── ConfigPanel/ # 配置面板(主题/TTS/语言/场景)
│ ├── hooks/ # 自定义 Hooks
│ │ ├── useVisionSession.ts # 核心会话 Hook (~500 行)
│ │ ├── useSessionList.ts # 对话列表管理
│ │ └── useObservationMode.ts # 观察模式
│ ├── lib/ # 工具库
│ │ ├── 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/ # 服务入口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 提示词与场景
│ │ ├── stt/ # STT 服务MiMo/Deepgram
│ │ └── tts/ # TTS 服务MiMo/OpenAI
│ ├── ws/ # WebSocket Handler
│ ├── 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.25
- **PostgreSQL** >= 15可选 Docker
- **Redis** >= 7可选用于缓存加速
### 本地开发
#### 1. 克隆项目
```bash
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 migrate
# 启动服务(监听 :8080
go run ./cmd/server
```
#### 4. 启动前端
```bash
cd frontend
# 安装依赖
npm install
# 启动开发服务器http://localhost:5173
npm run dev
```
#### 5. 访问应用
打开浏览器访问 [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) | 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 全链路追踪、日志级别 |
### 功能文档
| 文档 | 内容 |
|------|------|
| [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>

9
backend/.dockerignore Normal file
View File

@@ -0,0 +1,9 @@
bin
tmp
.git
.gitignore
*.md
.env*
.vscode
.idea
vendor

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

2
backend/.gitignore vendored
View File

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

38
backend/Dockerfile Normal file
View File

@@ -0,0 +1,38 @@
# ---- 构建阶段 ----
FROM golang:1.26-alpine AS builder
WORKDIR /app
# 使用国内 Go 代理
ENV GOPROXY=https://goproxy.cn,https://goproxy.io,direct
# 先复制依赖清单,利用 Docker 缓存层
COPY go.mod go.sum ./
# --mount=type=cache 复用 Go module 缓存,依赖不变时跳过下载
RUN --mount=type=cache,target=/go/pkg/mod \
go mod download
# 复制源码并构建
COPY . .
# 复用 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
RUN apk add --no-cache ca-certificates tzdata
WORKDIR /app
# 复制二进制和配置文件(敏感配置通过 docker-compose env_file 注入覆盖)
COPY --from=builder /camtalk .
COPY config/ ./config/
EXPOSE 8080
ENTRYPOINT ["./camtalk"]

View File

@@ -5,27 +5,38 @@ import (
"errors"
"net/http"
"os/signal"
"strings"
"syscall"
"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/ai/llm"
"github.com/hhs/camtalk/internal/auth"
"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"
)
// Version 通过构建时 -ldflags 注入,如:
// go build -ldflags "-X main.Version=v1.0.0" ./cmd/server
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())
}
@@ -39,19 +50,176 @@ func main() {
"addr", cfg.Server.Addr(),
)
// 初始化 Session ManagerMVP 默认内存实现
// 初始化存储层三级存储架构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 // 数据库连接池
// 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")
}
var err error
pool, err = store.NewPostgresPool(ctx, dsn)
if err != nil {
logger.Log.Fatalw("failed to connect to postgres", "error", err)
}
defer pool.Close()
// 执行数据库迁移
if err := store.RunMigrations(ctx, pool, migrations.FS); err != nil {
logger.Log.Fatalw("failed to run migrations", "error", err)
}
userRepo = store.NewPgUserRepository(pool)
msgRepo = store.NewPgMessageRepository(pool)
sessRepo = store.NewPgSessionRepository(pool)
logger.Log.Infow("L3 PostgreSQL storage initialized", "driver", cfg.Storage.Persistence.Driver)
} else {
userRepo = store.NewMemUserRepository()
logger.Log.Info("using in-memory user storage")
}
// 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
// TODO: 当 Redis 配置非空时切换为 RedisManager
sessionMgr = session.NewMemoryManager(30*time.Minute, 20)
defer sessionMgr.(*session.MemoryManager).Stop()
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)")
}
// 初始化 AI 服务
sttService := stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, logger.Log)
llmService := llm.NewOpenAIService(cfg.AI.LLM.APIKey, cfg.AI.LLM.Model, cfg.AI.LLM.Endpoint, cfg.AI.LLM.Timeout, logger.Log)
ttsService := tts.NewOpenAIService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Model, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Speed, cfg.AI.TTS.Timeout, logger.Log)
logger.Log.Infow("initializing AI services",
"stt.provider", cfg.AI.STT.Provider,
"stt.model", cfg.AI.STT.Model,
"llm.provider", cfg.AI.LLM.Provider,
"llm.model", cfg.AI.LLM.Model,
"tts.provider", cfg.AI.TTS.Provider,
"tts.model", cfg.AI.TTS.Model,
"tts.voice", cfg.AI.TTS.Voice,
)
// 初始化 Orchestrator
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr, cfg.AI.LLM.Model)
var sttService stt.Service
switch strings.ToLower(cfg.AI.STT.Provider) {
case "mimo", "xiaomi":
sttService = stt.NewMiMoService(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", "mimo", "model", cfg.AI.STT.Model, "endpoint", cfg.AI.STT.Endpoint)
default:
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)
}
var ttsService tts.Service
switch strings.ToLower(cfg.AI.TTS.Provider) {
case "mimo", "xiaomi":
ttsService = tts.NewMiMoService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Model, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Timeout, cfg.AI.TTS.HTTPClientTimeout, logger.Log)
logger.Log.Infow("TTS service initialized", "provider", "mimo", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "endpoint", cfg.AI.TTS.Endpoint)
default:
ttsService = tts.NewOpenAIService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Model, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Speed, cfg.AI.TTS.Timeout, cfg.AI.TTS.HTTPClientTimeout, logger.Log)
logger.Log.Infow("TTS service initialized", "provider", "openai", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "speed", cfg.AI.TTS.Speed)
}
// 初始化 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(
cfg.Auth.JWTSecret,
time.Duration(cfg.Auth.AccessTTL)*time.Minute,
time.Duration(cfg.Auth.RefreshTTL)*time.Minute,
)
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" {
@@ -59,20 +227,45 @@ func main() {
}
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")
{
apiGroup.GET("/health", healthHandler(sessionMgr))
apiGroup.GET("/health", healthHandler(sessionMgr, cfg))
}
// Session REST 端点
sessionHandler := api.NewSessionHandler(sessionMgr)
sessionHandler.RegisterRoutes(apiGroup)
// Auth REST 端点
authHandler := api.NewAuthHandler(authService, tokenMgr)
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))
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter, userScenarioRepo))
// HTTP Server
srv := &http.Server{
@@ -96,7 +289,7 @@ func main() {
<-ctx.Done()
logger.Log.Info("shutting down...")
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(cfg.Server.ShutdownTimeout)*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
@@ -106,11 +299,15 @@ func main() {
}
// healthHandler 健康检查。
func healthHandler(sessionMgr session.Manager) gin.HandlerFunc {
func healthHandler(sessionMgr session.Manager, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
version := Version
if version == "" {
version = cfg.App.Version
}
c.JSON(200, gin.H{
"status": "ok",
"version": "0.1.0",
"version": version,
"uptime_seconds": int(time.Since(startTime).Seconds()),
"active_sessions": sessionMgr.ActiveCount(),
})

View File

@@ -1,39 +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: Xiaomi MiMo
model: mimo-v2.5
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
llm:
provider: dashscope
model: qwen3-vl-plus
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
timeout: 30
tts:
provider: Xiaomi MiMo
model: mimo-v2.5
voice: alloy
speed: 1.0
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
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

@@ -1,24 +1,38 @@
module github.com/hhs/camtalk
go 1.24
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
@@ -27,16 +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/joho/godotenv v1.5.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
@@ -45,14 +68,18 @@ 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
golang.org/x/text v0.28.0 // indirect
golang.org/x/text v0.29.0 // indirect
google.golang.org/protobuf v1.34.1 // indirect
gopkg.in/yaml.v3 v3.0.1 // 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,44 +67,97 @@ 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=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
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=
@@ -86,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=
@@ -103,37 +189,60 @@ 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/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng=
golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU=
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=
google.golang.org/protobuf v1.34.1/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15 h1:YR8cESwS4TdDjEe65xsg0ogRM/Nc3DYOhEAlW+xobZo=
gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/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

@@ -15,10 +15,11 @@ type Service interface {
// Request 推理请求。
type Request struct {
Image []byte // JPEG 图片(已从 Base64 解码)
Text string // 用户语音识别后的文本
History []models.Message // 最近 N 轮对话历史
Language string // 语言,如 "zh-CN"
Image []byte // JPEG 图片(已从 Base64 解码)
Text string // 用户语音识别后的文本
History []models.Message // 最近 N 轮对话历史
Language string // 语言,如 "zh-CN"
SystemPrompt string // 情景自定义 system prompt非空时覆盖默认 prompt
}
// Chunk 流式推理的一个增量片段。

View File

@@ -1,237 +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 服务。
func NewOpenAIService(apiKey, model, endpoint string, timeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
if model == "" {
model = "gpt-4o"
}
if endpoint == "" {
endpoint = "https://api.openai.com/v1"
}
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 10 * time.Second
}
return &OpenAIService{
apiKey: apiKey,
model: model,
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: 60 * time.Second}, // HTTP client timeout > LLM timeout
}
}
// --- 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,omitempty"`
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)
}
// 创建带超时的 context
ctx, cancel := context.WithTimeout(ctx, o.timeout)
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/chat/completions", bytes.NewReader(payload))
if err != nil {
cancel()
return nil, fmt.Errorf("llm: create request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+o.apiKey)
resp, err := o.client.Do(httpReq)
if err != nil {
cancel()
return nil, fmt.Errorf("llm: send request: %w", err)
}
if resp.StatusCode != http.StatusOK {
cancel()
bodyBytes, _ := io.ReadAll(resp.Body)
resp.Body.Close()
return nil, fmt.Errorf("llm: api error (status %d): %s", resp.StatusCode, string(bodyBytes))
}
// 启动 goroutine 解析 SSE 流
ch := make(chan Chunk, 64)
go func() {
defer close(ch)
defer cancel()
defer resp.Body.Close()
o.parseSSEStream(resp.Body, ch)
}()
return ch, nil
}
// parseSSEStream 解析 SSE 流,将 delta 发送到 channel。
func (o *OpenAIService) parseSSEStream(body io.Reader, ch chan<- Chunk) {
scanner := bufio.NewScanner(body)
scanner.Buffer(make([]byte, 0, 64*1024), 256*1024)
var fullText strings.Builder
var lastModel string
for scanner.Scan() {
line := scanner.Text()
// SSE 格式data: {...}
if !strings.HasPrefix(line, "data: ") {
continue
}
data := strings.TrimPrefix(line, "data: ")
if data == "[DONE]" {
// 流结束,发送最终 chunk
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
return
}
var delta streamDelta
if err := json.Unmarshal([]byte(data), &delta); err != nil {
o.logger.Warnw("llm: unmarshal delta failed", "error", err, "data", data)
continue
}
if delta.Model != "" {
lastModel = delta.Model
}
// 提取增量文本
if len(delta.Choices) > 0 {
content := delta.Choices[0].Delta.Content
if content != "" {
fullText.WriteString(content)
ch <- Chunk{Delta: content, Done: false, Model: lastModel}
}
// 某些模型在最后一个 choice 中携带 usage
if delta.Choices[0].FinishReason != nil && delta.Usage != nil {
ch <- Chunk{
Delta: "",
Done: true,
Model: lastModel,
TokensUsed: &TokenUsage{
Prompt: delta.Usage.PromptTokens,
Completion: delta.Usage.CompletionTokens,
Total: delta.Usage.TotalTokens,
},
}
return
}
}
}
// scanner 结束但没收到 [DONE]
if err := scanner.Err(); err != nil {
o.logger.Warnw("llm: scan error", "error", err)
}
ch <- Chunk{Delta: "", Done: true, Model: lastModel}
}
// buildMessages 构建 OpenAI Chat API 的 messages 数组。
func (o *OpenAIService) buildMessages(req Request) []chatMessage {
var messages []chatMessage
// System prompt
messages = append(messages, chatMessage{
Role: "system",
Content: []contentPart{{Type: "text", Text: BuildSystemPrompt(req.Language, "")}},
})
// 历史消息
for _, msg := range req.History {
messages = append(messages, chatMessage{
Role: msg.Role,
Content: []contentPart{{Type: "text", Text: msg.Content}},
})
}
// 当前用户消息(图像 + 文本)
var parts []contentPart
if len(req.Image) > 0 {
b64 := base64.StdEncoding.EncodeToString(req.Image)
parts = append(parts, contentPart{
Type: "image_url",
ImageURL: &imageURL{URL: "data:image/jpeg;base64," + b64},
})
}
parts = append(parts, contentPart{Type: "text", Text: req.Text})
messages = append(messages, chatMessage{Role: "user", Content: parts})
return messages
}

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, 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, 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, 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, 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, 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, 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,10 +2,26 @@ package llm
import "strings"
// BuildSystemPrompt 根据语言细节级别构建系统提示词。
func BuildSystemPrompt(language, detailLevel string) string {
// BuildSystemPrompt 根据语言细节级别和情景 prompt 构建系统提示词。
// scenarioPrompt 非空时,覆盖默认视觉助手 prompt。
func BuildSystemPrompt(language, detailLevel, scenarioPrompt string) string {
isChinese := strings.HasPrefix(language, "zh")
// 情景模式:使用自定义 prompt 作为基础
if scenarioPrompt != "" {
var prompt strings.Builder
prompt.WriteString(scenarioPrompt)
if detailLevel == "high" {
if isChinese {
prompt.WriteString(" 请在涉及视觉内容时提供更详细的描述,包括颜色、位置、数量等细节。")
} else {
prompt.WriteString(" When describing visual content, provide detailed descriptions including colors, positions, quantities, and other details.")
}
}
return prompt.String()
}
// 默认模式:视觉助手
var prompt strings.Builder
if isChinese {
prompt.WriteString("你是一个视觉助手。用户通过摄像头看到一个场景,并用语音向你提问。请用简洁自然的中文回答。如果涉及视觉描述,先说\"我看到……\"。回答控制在3-5句话以内除非用户要求详细说明。")

View File

@@ -0,0 +1,148 @@
package llm
import "strings"
// scenarioPrompt 定义单个情景的多语言 system prompt 和首句引导。
type scenarioPrompt struct {
ZH string
EN string
JA string
GreetingZH string // 首句引导(中文)
GreetingEN string // 首句引导(英文)
GreetingJA string // 首句引导(日文)
}
// scenarioPrompts 预置情景 → prompt 映射表。
// key 为情景 ID与前端 Scenario.id 对齐)。
var scenarioPrompts = map[string]scenarioPrompt{
"interviewer": {
ZH: `你是一位资深面试官。你通过摄像头观察面试者,并根据他们的背景和表现提出面试问题。
【角色定位】
- 你是面试官,不是助手或顾问
- 你的目标是评估候选人的能力
- 保持专业、客观、礼貌
【交互规则】
1. 每次只问一个问题,等用户回答后再追问
2. 问题要有层次:自我介绍 → 专业问题 → 情景题 → 反问环节
3. 对用户的回答给出简短点评(优点+不足),然后追问
4. 如果摄像头能看到用户的环境,可以结合环境提出相关话题
5. 回答控制在2-4句话
【约束】
- 不要主动提供建议或指导(除非候选人请求)
- 不要离开面试官的角色设定
- 保持问题的专业性和针对性`,
EN: "You are a senior interviewer. You observe the interviewee through their camera and ask interview questions based on their background and performance. Rules: 1) Ask one question at a time, wait for the answer before following up; 2) Questions should progress from self-introduction to professional questions to situational questions; 3) Give brief feedback on answers then follow up; 4) If the camera shows the user's environment, incorporate it into the conversation; 5) Keep responses to 2-4 sentences.",
JA: "あなたはベテラン面接官です。カメラで面接者を見て、バックグラウンドと実績に基づいて面接質問をします。ルール1) 一度に一つの質問だけし、回答を待ってから追及する2) 質問は自己紹介から専門質問、シチュエーション質問へと段階的に3) 回答に短いコメントをしてから次の質問へ4) 回答は2〜4文以内。",
GreetingZH: "你好!我是今天的面试官。让我们先从自我介绍开始,请简单介绍一下你自己和你应聘的岗位。",
GreetingEN: "Hello! I'm your interviewer today. Let's start with a self-introduction. Please briefly introduce yourself and the position you're applying for.",
GreetingJA: "こんにちは!本日の面接官です。まず自己紹介から始めましょう。あなた自身と応募職種について簡単に教えてください。",
},
"english_teacher": {
ZH: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material (e.g., 'I can see a bookshelf behind you. What's your favorite book?'); 6) Keep responses to 3-5 sentences.",
EN: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
JA: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
GreetingZH: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
GreetingEN: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
GreetingJA: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
},
"debate": {
ZH: `你是一位辩论赛对手。用户提出一个观点,你需要站在反方进行反驳。
【角色定位】
- 你是辩论对手,不是评委或顾问
- 你的目标是通过逻辑论证反驳对方观点
- 保持理性、严谨、尊重对手
【交互规则】
1. 逻辑严密,用事实和论据反驳,不要人身攻击
2. 每次提出1-2个核心反驳点并给出简要论据
3. 如果用户论证有力,承认其合理性但仍要寻找突破口
4. 适时提出反问,引导用户深入思考
5. 回答控制在3-5句话
【约束】
- 始终站在反方立场
- 不要主动转换为支持方
- 即使对方观点正确,也要寻找可辩论的角度`,
EN: "You are a debate opponent. The user presents a viewpoint, and you argue against it. Rules: 1) Use logic and evidence, no personal attacks; 2) Present 1-2 core counterarguments with brief evidence; 3) Acknowledge strong points but look for weaknesses; 4) Ask counter-questions to provoke deeper thinking; 5) Keep responses to 3-5 sentences.",
JA: "あなたはディベートの相手です。ユーザーが提示した观点に対して反論します。ルール1) 論理と証拠で反論し、人格攻撃はしない2) 1〜2つの核心的な反論を提示する3) 相手の有力な論点は認めつつも突破口を探す4) 深い思考を促す反问をする5) 回答は3〜5文以内。",
GreetingZH: "你好!我是你的辩论对手。请提出一个你坚信的观点,我会站在反方立场与你辩论,帮你锻炼逻辑思维。",
GreetingEN: "Hello! I'm your debate opponent. Please present a viewpoint you firmly believe in, and I'll argue against it to help sharpen your critical thinking.",
GreetingJA: "こんにちは!あなたのディベート相手です。あなたが信じる观点を提示してください。反対の立場から論じて、論理的思考を鍛えます。",
},
"interpreter": {
ZH: "你是一名同声翻译员。将用户说的话实时翻译为目标语言。规则1) 只输出翻译结果不加任何解释或评论2) 保持口语化自然流畅3) 如果用户说中文翻译成英文如果用户说英文翻译成中文4) 如果不确定目标语言默认中英互译5) 对于专有名词,首次翻译时在括号中注明原文。",
EN: "You are a simultaneous interpreter. Translate what the user says in real-time. Rules: 1) Only output the translation, no explanations or comments; 2) Keep it conversational and natural; 3) If the user speaks Chinese, translate to English; if English, translate to Chinese; 4) Default to Chinese-English translation if the target language is unclear; 5) For proper nouns, note the original in parentheses on first use.",
JA: "あなたは同時通訳者です。ユーザーの発言をリアルタイムで翻訳します。ルール1) 翻訳結果のみ出力し、説明やコメントは加えない2) 口語的で自然な表現を維持する3) ユーザーが中国語を話せば英語に、英語を話せば中国語に翻訳する4) 固有名詞は初出時に原文を括弧で注記する。",
GreetingZH: "我是你的同声翻译。请开始说话,我会实时将中文翻译成英文,或将英文翻译成中文。",
GreetingEN: "I'm your simultaneous interpreter. Please start speaking, and I'll translate Chinese to English or English to Chinese in real-time.",
GreetingJA: "私はあなたの同時通訳者です。お話しください。中国語を英語に、または英語を中国語にリアルタイムで翻訳します。",
},
}
// GetScenarioPrompt 根据情景 ID 和语言获取对应的 system prompt。
// 支持系统预置情景和用户自建情景。
// customScenarios: 用户自建情景映射表scenarioID → prompt可为 nil
// 返回空字符串表示无此情景(使用默认 prompt
func GetScenarioPrompt(scenarioID, language string, customScenarios map[string]string) string {
if scenarioID == "" || scenarioID == "free_chat" {
return ""
}
// 1. 优先查找系统预置情景
if p, ok := scenarioPrompts[scenarioID]; ok {
switch {
case strings.HasPrefix(language, "zh"):
return p.ZH
case strings.HasPrefix(language, "ja"):
return p.JA
default:
return p.EN
}
}
// 2. 查找用户自建情景
if customScenarios != nil {
if customPrompt, ok := customScenarios[scenarioID]; ok {
return customPrompt
}
}
// 3. 默认空字符串
return ""
}
// GetScenarioGreeting 根据情景 ID 和语言获取对应的首句引导。
// 支持系统预置情景和用户自建情景。
// customGreetings: 用户自建情景的首句引导映射表scenarioID → greeting可为 nil
// 返回空字符串表示无此情景或不需要引导(自由对话)。
func GetScenarioGreeting(scenarioID, language string, customGreetings map[string]string) string {
if scenarioID == "" || scenarioID == "free_chat" {
return ""
}
// 1. 优先查找系统预置情景
if p, ok := scenarioPrompts[scenarioID]; ok {
switch {
case strings.HasPrefix(language, "zh"):
return p.GreetingZH
case strings.HasPrefix(language, "ja"):
return p.GreetingJA
default:
return p.GreetingEN
}
}
// 2. 查找用户自建情景
if customGreetings != nil {
if customGreeting, ok := customGreetings[scenarioID]; ok {
return customGreeting
}
}
// 3. 默认空字符串
return ""
}

View File

@@ -18,21 +18,22 @@ type DeepgramService struct {
apiKey string
model string
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
}
// NewDeepgramService 创建 Deepgram STT 服务。
func NewDeepgramService(apiKey, model, endpoint string, logger *zap.SugaredLogger) *DeepgramService {
if model == "" {
model = "nova-2"
}
if endpoint == "" {
endpoint = "wss://api.deepgram.com/v1/listen"
// model、endpoint 由 config 层保证非空timeoutSec 为 0 时默认 5 秒。
func NewDeepgramService(apiKey, model, endpoint string, timeoutSec int, logger *zap.SugaredLogger) *DeepgramService {
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 5 * time.Second
}
return &DeepgramService{
apiKey: apiKey,
model: model,
endpoint: endpoint,
timeout: timeout,
logger: logger,
}
}
@@ -57,8 +58,8 @@ func (d *DeepgramService) Recognize(ctx context.Context, audio []byte, opts Opti
// 构建 WebSocket URL附带查询参数
wsURL := d.buildURL(opts)
// 5 秒总超时
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
// 总超时
ctx, cancel := context.WithTimeout(ctx, d.timeout)
defer cancel()
// 建立 WebSocket 连接

View File

@@ -1,191 +0,0 @@
package stt
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gorilla/websocket"
"go.uber.org/zap"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
}
// newMockDeepgram 创建模拟 Deepgram WebSocket 服务。
// 返回 httptest.Server 和对应的 ws:// URL。
func newMockDeepgram(t *testing.T, handler func(conn *websocket.Conn)) *httptest.Server {
t.Helper()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
t.Logf("upgrade error: %v", err)
return
}
defer conn.Close()
handler(conn)
}))
return srv
}
// wsToWss 将 http:// 转换为 ws://。
func wsToWss(httpURL string) string {
return "ws" + strings.TrimPrefix(httpURL, "http")
}
func TestDeepgramService_Recognize_Success(t *testing.T) {
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
// 读取音频数据
_, _, err := conn.ReadMessage()
if err != nil {
t.Errorf("read audio: %v", err)
return
}
// 发送中间结果(非 final
intermediate := deepgramResponse{
IsFinal: false,
}
intermediate.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: "你好", Confidence: 0.9}}
data, _ := json.Marshal(intermediate)
_ = conn.WriteMessage(websocket.TextMessage, data)
// 发送最终结果
final := deepgramResponse{
IsFinal: true,
}
final.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: "你好世界", Confidence: 0.95}}
data, _ = json.Marshal(final)
_ = conn.WriteMessage(websocket.TextMessage, data)
// 等待客户端关闭
_, _, _ = conn.ReadMessage()
})
defer srv.Close()
svc := NewDeepgramService("test-key", "", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
text, err := svc.Recognize(context.Background(), []byte("fake-pcm-audio"), Options{
Encoding: "pcm_s16le",
SampleRate: 16000,
Language: "zh-CN",
})
if err != nil {
t.Fatalf("Recognize() error: %v", err)
}
if text != "你好世界" {
t.Errorf("Recognize() = %q, want %q", text, "你好世界")
}
}
func TestDeepgramService_Recognize_EmptyAudio(t *testing.T) {
svc := NewDeepgramService("test-key", "", "ws://localhost", zap.NewNop().Sugar())
_, err := svc.Recognize(context.Background(), nil, Options{})
if err == nil {
t.Fatal("Recognize() with empty audio should return error")
}
}
func TestDeepgramService_Recognize_ConnectError(t *testing.T) {
svc := NewDeepgramService("test-key", "", "ws://localhost:1", zap.NewNop().Sugar())
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
_, err := svc.Recognize(ctx, []byte("audio"), Options{})
if err == nil {
t.Fatal("Recognize() with bad endpoint should return error")
}
}
func TestDeepgramService_Recognize_Timeout(t *testing.T) {
// 模拟一个永不响应的服务端
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
// 读取音频但不发送任何结果,让客户端超时
_, _, _ = conn.ReadMessage()
time.Sleep(10 * time.Second)
})
defer srv.Close()
svc := NewDeepgramService("test-key", "", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
defer cancel()
_, err := svc.Recognize(ctx, []byte("audio"), Options{})
if err == nil {
t.Fatal("Recognize() should timeout")
}
}
func TestDeepgramService_Recognize_MultipleFinals(t *testing.T) {
srv := newMockDeepgram(t, func(conn *websocket.Conn) {
_, _, _ = conn.ReadMessage()
// 发送多个 final 结果(多句话场景)
for _, text := range []string{"你好", "世界"} {
resp := deepgramResponse{IsFinal: true}
resp.Channel.Alternatives = []struct {
Transcript string `json:"transcript"`
Confidence float64 `json:"confidence"`
}{{Transcript: text, Confidence: 0.9}}
data, _ := json.Marshal(resp)
_ = conn.WriteMessage(websocket.TextMessage, data)
}
_, _, _ = conn.ReadMessage()
})
defer srv.Close()
svc := NewDeepgramService("test-key", "", wsToWss(srv.URL)+"/v1/listen", zap.NewNop().Sugar())
text, err := svc.Recognize(context.Background(), []byte("audio"), Options{})
if err != nil {
t.Fatalf("Recognize() error: %v", err)
}
if text != "你好世界" {
t.Errorf("Recognize() = %q, want %q", text, "你好世界")
}
}
func TestDeepgramService_buildURL(t *testing.T) {
svc := NewDeepgramService("key", "", "wss://api.deepgram.com/v1/listen", zap.NewNop().Sugar())
tests := []struct {
name string
opts Options
want []string // URL 中应包含的参数
}{
{
name: "defaults",
opts: Options{},
want: []string{"encoding=pcm_s16le", "sample_rate=16000", "language=zh-CN"},
},
{
name: "custom",
opts: Options{Encoding: "wav", SampleRate: 44100, Language: "en"},
want: []string{"encoding=wav", "sample_rate=44100", "language=en"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
u := svc.buildURL(tt.opts)
for _, param := range tt.want {
if !strings.Contains(u, param) {
t.Errorf("buildURL() = %q, should contain %q", u, param)
}
}
})
}
}

View File

@@ -0,0 +1,234 @@
package stt
import (
"bytes"
"context"
"encoding/base64"
"encoding/binary"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"go.uber.org/zap"
)
// MiMoService 基于 Xiaomi MiMo ASR HTTP API 的语音识别实现。
// 接口兼容 OpenAI chat/completions 格式,音频仅支持 mp3/wav。
type MiMoService struct {
apiKey string
model string
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
}
// NewMiMoService 创建 MiMo STT 服务。
// model、endpoint 由 config 层保证非空timeoutSec 为 0 时默认 10 秒。
func NewMiMoService(apiKey, model, endpoint string, timeoutSec int, logger *zap.SugaredLogger) *MiMoService {
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 10 * time.Second
}
return &MiMoService{
apiKey: apiKey,
model: model,
endpoint: endpoint,
timeout: timeout,
logger: logger,
}
}
// mimoRequest MiMo ASR 请求体。
type mimoRequest struct {
Model string `json:"model"`
Messages []mimoMessage `json:"messages"`
ASROptions *mimoASROptions `json:"asr_options,omitempty"`
}
type mimoMessage struct {
Role string `json:"role"`
Content []mimoContent `json:"content"`
}
type mimoContent struct {
Type string `json:"type"`
InputAudio *mimoAudioIn `json:"input_audio,omitempty"`
}
type mimoAudioIn struct {
Data string `json:"data"` // data URL: data:{mime};base64,{data}
}
type mimoASROptions struct {
Language string `json:"language"`
}
// mimoResponse MiMo ASR 非流式响应。
type mimoResponse struct {
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
}
// Recognize 实现 stt.Service。将音频发送到 MiMo ASR API返回识别文本。
func (m *MiMoService) Recognize(ctx context.Context, audio []byte, opts Options) (string, error) {
if len(audio) == 0 {
return "", fmt.Errorf("stt: empty audio")
}
// MiMo 仅支持 mp3/wav若输入为原始 PCM 则封装为 WAV
audioData := audio
mimeType := "audio/wav"
if !isWAV(audio) && !isMP3(audio) {
wav, err := pcmToWAV(audio, opts.SampleRate, 1)
if err != nil {
return "", fmt.Errorf("stt: pcm to wav: %w", err)
}
audioData = wav
} else if isMP3(audio) {
mimeType = "audio/mpeg"
}
b64 := base64.StdEncoding.EncodeToString(audioData)
dataURL := fmt.Sprintf("data:%s;base64,%s", mimeType, b64)
// 映射语言代码
language := mapLanguage(opts.Language)
reqBody := mimoRequest{
Model: m.model,
Messages: []mimoMessage{
{
Role: "user",
Content: []mimoContent{
{
Type: "input_audio",
InputAudio: &mimoAudioIn{
Data: dataURL,
},
},
},
},
},
}
if language != "" {
reqBody.ASROptions = &mimoASROptions{Language: language}
}
body, err := json.Marshal(reqBody)
if err != nil {
return "", fmt.Errorf("stt: marshal request: %w", err)
}
url := strings.TrimRight(m.endpoint, "/") + "/chat/completions"
ctx, cancel := context.WithTimeout(ctx, m.timeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
if err != nil {
return "", fmt.Errorf("stt: create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+m.apiKey)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return "", fmt.Errorf("stt: request mimo: %w", err)
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return "", fmt.Errorf("stt: read response: %w", err)
}
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("stt: mimo returned %d: %s", resp.StatusCode, string(respBody))
}
var mResp mimoResponse
if err := json.Unmarshal(respBody, &mResp); err != nil {
return "", fmt.Errorf("stt: unmarshal response: %w", err)
}
if len(mResp.Choices) == 0 {
// MiMo 返回空结果,视为无法识别(非错误),返回空文本
return "", nil
}
text := strings.TrimSpace(mResp.Choices[0].Message.Content)
return text, nil
}
// mapLanguage 将标准语言代码映射为 MiMo 支持的值auto/zh/en
func mapLanguage(lang string) string {
switch {
case lang == "":
return "auto"
case strings.HasPrefix(lang, "zh"):
return "zh"
case strings.HasPrefix(lang, "en"):
return "en"
default:
return "auto"
}
}
// isWAV 检查数据是否为 WAV 格式RIFF 头)。
func isWAV(data []byte) bool {
return len(data) > 4 && string(data[:4]) == "RIFF"
}
// isMP3 检查数据是否为 MP3 格式ID3 标签或帧同步字)。
func isMP3(data []byte) bool {
if len(data) > 3 && string(data[:3]) == "ID3" {
return true
}
// 帧同步字0xFF 0xFB/0xF3/0xF2
return len(data) > 2 && data[0] == 0xFF && (data[1]&0xE0) == 0xE0
}
// pcmToWAV 将原始 PCM 数据封装为 WAV 文件。
func pcmToWAV(pcm []byte, sampleRate, channels int) ([]byte, error) {
if sampleRate == 0 {
sampleRate = 16000
}
if channels == 0 {
channels = 1
}
bitsPerSample := 16
byteRate := sampleRate * channels * bitsPerSample / 8
blockAlign := channels * bitsPerSample / 8
dataSize := len(pcm)
var buf bytes.Buffer
// RIFF header
buf.WriteString("RIFF")
binary.Write(&buf, binary.LittleEndian, uint32(36+dataSize))
buf.WriteString("WAVE")
// fmt 子块
buf.WriteString("fmt ")
binary.Write(&buf, binary.LittleEndian, uint32(16)) // 子块大小
binary.Write(&buf, binary.LittleEndian, uint16(1)) // PCM 格式
binary.Write(&buf, binary.LittleEndian, uint16(channels)) // 通道数
binary.Write(&buf, binary.LittleEndian, uint32(sampleRate)) // 采样率
binary.Write(&buf, binary.LittleEndian, uint32(byteRate)) // 字节率
binary.Write(&buf, binary.LittleEndian, uint16(blockAlign)) // 块对齐
binary.Write(&buf, binary.LittleEndian, uint16(bitsPerSample)) // 每样本位数
// data 子块
buf.WriteString("data")
binary.Write(&buf, binary.LittleEndian, uint32(dataSize))
buf.Write(pcm)
return buf.Bytes(), nil
}

View File

@@ -0,0 +1,258 @@
package stt
import (
"context"
"encoding/base64"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"testing"
"go.uber.org/zap"
)
func newTestMiMoService(handler http.HandlerFunc) (*MiMoService, *httptest.Server) {
srv := httptest.NewServer(handler)
s := NewMiMoService("test-key", "mimo-v2.5-asr", srv.URL, 0, zap.NewNop().Sugar())
return s, srv
}
func TestMiMoService_Recognize_Success(t *testing.T) {
s, srv := newTestMiMoService(func(w http.ResponseWriter, r *http.Request) {
// 验证请求
if r.Header.Get("Authorization") != "Bearer test-key" {
t.Errorf("expected Authorization Bearer test-key, got %s", r.Header.Get("Authorization"))
}
if r.URL.Path != "/chat/completions" {
t.Errorf("expected path /chat/completions, got %s", r.URL.Path)
}
var req mimoRequest
body, _ := io.ReadAll(r.Body)
if err := json.Unmarshal(body, &req); err != nil {
t.Fatalf("unmarshal request: %v", err)
}
if req.Model != "mimo-v2.5-asr" {
t.Errorf("expected model mimo-v2.5-asr, got %s", req.Model)
}
if len(req.Messages) == 0 || req.Messages[0].Role != "user" {
t.Error("expected user message")
}
if req.ASROptions == nil || req.ASROptions.Language != "zh" {
t.Errorf("expected language zh, got %v", req.ASROptions)
}
resp := mimoResponse{
Choices: []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
}{
{Message: struct {
Content string `json:"content"`
}{Content: "你好世界"}},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
})
defer srv.Close()
// 发送一个简单的有效 WAV44 字节头 + 少量 PCM
wav := makeValidWAV([]byte{0x00, 0x00, 0x00, 0x00})
text, err := s.Recognize(context.Background(), wav, Options{Language: "zh-CN"})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if text != "你好世界" {
t.Errorf("expected '你好世界', got '%s'", text)
}
}
func TestMiMoService_Recognize_EmptyAudio(t *testing.T) {
s, srv := newTestMiMoService(func(w http.ResponseWriter, r *http.Request) {})
defer srv.Close()
_, err := s.Recognize(context.Background(), nil, Options{})
if err == nil {
t.Fatal("expected error for empty audio")
}
}
func TestMiMoService_Recognize_ServerError(t *testing.T) {
s, srv := newTestMiMoService(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
w.Write([]byte("internal error"))
})
defer srv.Close()
wav := makeValidWAV([]byte{0x00, 0x00})
_, err := s.Recognize(context.Background(), wav, Options{})
if err == nil {
t.Fatal("expected error for 500 response")
}
}
func TestMiMoService_Recognize_EmptyChoices(t *testing.T) {
s, srv := newTestMiMoService(func(w http.ResponseWriter, r *http.Request) {
resp := mimoResponse{Choices: nil}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
})
defer srv.Close()
wav := makeValidWAV([]byte{0x00, 0x00})
text, err := s.Recognize(context.Background(), wav, Options{})
if err != nil {
t.Fatalf("unexpected error for empty choices: %v", err)
}
if text != "" {
t.Errorf("expected empty string for empty choices, got %q", text)
}
}
func TestMiMoService_Recognize_PCMAutoWrap(t *testing.T) {
// 测试原始 PCM 数据自动封装为 WAV
s, srv := newTestMiMoService(func(w http.ResponseWriter, r *http.Request) {
var req mimoRequest
body, _ := io.ReadAll(r.Body)
if err := json.Unmarshal(body, &req); err != nil {
t.Fatalf("unmarshal request: %v", err)
}
// 验证 data URL 格式
if len(req.Messages) == 0 || len(req.Messages[0].Content) == 0 {
t.Fatal("empty message content")
}
dataURL := req.Messages[0].Content[0].InputAudio.Data
if len(dataURL) < 22 || dataURL[:14] != "data:audio/wav" {
t.Errorf("expected wav data URL, got prefix: %s", dataURL[:min(len(dataURL), 30)])
}
// 验证 base64 可解码
b64Part := dataURL[22:] // skip "data:audio/wav;base64,"
decoded, err := base64.StdEncoding.DecodeString(b64Part)
if err != nil {
t.Fatalf("base64 decode failed: %v", err)
}
// 应该是有效 WAVRIFF 头)
if len(decoded) < 44 || string(decoded[:4]) != "RIFF" {
t.Error("decoded data is not a valid WAV")
}
resp := mimoResponse{
Choices: []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
}{
{Message: struct {
Content string `json:"content"`
}{Content: "test"}},
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
})
defer srv.Close()
// 发送原始 PCM非 WAV/MP3
pcm := []byte{0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07}
text, err := s.Recognize(context.Background(), pcm, Options{SampleRate: 16000})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if text != "test" {
t.Errorf("expected 'test', got '%s'", text)
}
}
func TestMapLanguage(t *testing.T) {
tests := []struct {
input string
want string
}{
{"", "auto"},
{"zh-CN", "zh"},
{"zh", "zh"},
{"en-US", "en"},
{"en", "en"},
{"ja", "auto"},
}
for _, tt := range tests {
got := mapLanguage(tt.input)
if got != tt.want {
t.Errorf("mapLanguage(%q) = %q, want %q", tt.input, got, tt.want)
}
}
}
func TestIsWAV(t *testing.T) {
if !isWAV([]byte("RIFF....")) {
t.Error("expected true for RIFF header")
}
if isWAV([]byte("ID3...")) {
t.Error("expected false for ID3 header")
}
if isWAV([]byte{0x00}) {
t.Error("expected false for short data")
}
}
func TestIsMP3(t *testing.T) {
if !isMP3([]byte("ID3\x03")) {
t.Error("expected true for ID3 header")
}
if !isMP3([]byte{0xFF, 0xFB, 0x00}) {
t.Error("expected true for MP3 sync word")
}
if isMP3([]byte("RIFF")) {
t.Error("expected false for RIFF header")
}
}
func makeValidWAV(pcm []byte) []byte {
// 构造一个最小有效 WAV
wav := make([]byte, 44+len(pcm))
copy(wav[:4], "RIFF")
// little-endian size = 36 + len(pcm)
size := uint32(36 + len(pcm))
wav[4] = byte(size)
wav[5] = byte(size >> 8)
wav[6] = byte(size >> 16)
wav[7] = byte(size >> 24)
copy(wav[8:12], "WAVE")
copy(wav[12:16], "fmt ")
// fmt chunk size = 16
wav[16] = 16
// PCM format = 1
wav[20] = 1
// channels = 1
wav[22] = 1
// sample rate = 16000
wav[24] = 0x80
wav[25] = 0x3E
// byte rate = 32000
wav[28] = 0x00
wav[29] = 0x7D
// block align = 2
wav[32] = 2
// bits per sample = 16
wav[34] = 16
copy(wav[36:40], "data")
dSize := uint32(len(pcm))
wav[40] = byte(dSize)
wav[41] = byte(dSize >> 8)
wav[42] = byte(dSize >> 16)
wav[43] = byte(dSize >> 24)
copy(wav[44:], pcm)
return wav
}
func min(a, b int) int {
if a < b {
return a
}
return b
}

View File

@@ -0,0 +1,214 @@
package tts
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/hhs/camtalk/internal/trace"
"github.com/hhs/camtalk/internal/util"
"go.uber.org/zap"
)
// MiMoService 基于 Xiaomi MiMo TTS API 的语音合成实现。
// 接口兼容 OpenAI chat/completions 格式,通过 messages 传递待合成文本与风格指令。
type MiMoService struct {
apiKey string
model string
voice string
endpoint string
timeout time.Duration
logger *zap.SugaredLogger
client *http.Client
}
// NewMiMoService 创建 MiMo TTS 服务。
// model、voice、endpoint 由 config 层保证非空。
func NewMiMoService(apiKey, model, voice, endpoint string, timeoutSec, httpClientTimeoutSec int, logger *zap.SugaredLogger) *MiMoService {
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 5 * time.Second
}
httpClientTimeout := time.Duration(httpClientTimeoutSec) * time.Second
if httpClientTimeout <= 0 {
httpClientTimeout = 30 * time.Second
}
return &MiMoService{
apiKey: apiKey,
model: model,
voice: voice,
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: httpClientTimeout},
}
}
// mimoTTSRequest MiMo TTS API 请求体。
type mimoTTSRequest struct {
Model string `json:"model"`
Messages []mimoTTSMessage `json:"messages"`
Audio mimoTTSAudio `json:"audio"`
Stream bool `json:"stream"`
}
// mimoTTSMessage MiMo TTS 消息。
type mimoTTSMessage struct {
Role string `json:"role"` // "user"(风格指令)| "assistant"(待合成文本)
Content string `json:"content"`
}
// mimoTTSAudio MiMo TTS 音频配置。
type mimoTTSAudio struct {
Format string `json:"format"` // "mp3" | "wav" | "pcm16"
Voice string `json:"voice"` // 预置音色 ID
}
// mimoTTSResponse MiMo TTS 非流式响应。
type mimoTTSResponse struct {
Choices []struct {
Message struct {
Audio struct {
Data string `json:"data"` // base64 编码的音频数据
} `json:"audio"`
} `json:"message"`
} `json:"choices"`
}
// mimoTTSStreamResponse MiMo TTS 流式响应。
type mimoTTSStreamResponse struct {
Choices []struct {
Delta struct {
Audio struct {
Data string `json:"data"` // base64 编码的音频数据片段
} `json:"audio"`
} `json:"delta"`
} `json:"choices"`
}
// SynthesizeStream 实现 tts.Service。从 textStream 读取句子,逐句调用 MiMo TTS API。
func (m *MiMoService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error) {
voice := opts.Voice
if voice == "" {
voice = m.voice
}
ch := make(chan Chunk, 4)
go func() {
defer close(ch)
for text := range textStream {
if text == "" {
continue
}
audio, err := m.synthesize(ctx, text, voice)
if err != nil {
log := trace.FromContext(ctx)
log.Warnw("mimo tts: synthesize failed",
"error", err,
"text_len", len(text),
"text_preview", util.Truncate(text, 100))
// 静默跳过,不中断整个流
continue
}
select {
case ch <- Chunk{Audio: audio, IsLast: true, Final: false}:
case <-ctx.Done():
return
}
}
// textStream 关闭,发送 Final 标记
select {
case ch <- Chunk{Audio: nil, IsLast: false, Final: true}:
case <-ctx.Done():
}
}()
return ch, nil
}
// synthesize 调用 MiMo TTS API 合成单个句子。
// 使用非流式调用返回完整音频数据base64 解码后)。
func (m *MiMoService) synthesize(ctx context.Context, text, voice string) ([]byte, error) {
// 单句超时
ctx, cancel := context.WithTimeout(ctx, m.timeout)
defer cancel()
// 构建 MiMo TTS 请求:文本放在 assistant 消息中
reqBody := mimoTTSRequest{
Model: m.model,
Messages: []mimoTTSMessage{
{
Role: "assistant",
Content: text,
},
},
Audio: mimoTTSAudio{
Format: "mp3",
Voice: voice,
},
Stream: false,
}
payload, err := json.Marshal(reqBody)
if err != nil {
return nil, fmt.Errorf("mimo tts: marshal request: %w", err)
}
url := strings.TrimRight(m.endpoint, "/") + "/chat/completions"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
if err != nil {
return nil, fmt.Errorf("mimo tts: create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("api-key", m.apiKey)
resp, err := m.client.Do(req)
if err != nil {
return nil, fmt.Errorf("mimo tts: send request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
errBody, _ := io.ReadAll(resp.Body)
return nil, fmt.Errorf("mimo tts: api error (status %d): %s", resp.StatusCode, string(errBody))
}
// 非流式响应:解析 JSON提取 base64 音频数据
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("mimo tts: read response: %w", err)
}
var ttsResp mimoTTSResponse
if err := json.Unmarshal(respBody, &ttsResp); err != nil {
return nil, fmt.Errorf("mimo tts: unmarshal response: %w", err)
}
if len(ttsResp.Choices) == 0 {
return nil, fmt.Errorf("mimo tts: empty choices in response")
}
audioData := ttsResp.Choices[0].Message.Audio.Data
if audioData == "" {
return nil, fmt.Errorf("mimo tts: empty audio data in response")
}
// base64 解码音频数据
audio, err := base64.StdEncoding.DecodeString(audioData)
if err != nil {
return nil, fmt.Errorf("mimo tts: decode audio base64: %w", err)
}
return audio, nil
}

View File

@@ -0,0 +1,414 @@
package tts
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"go.uber.org/zap"
)
// mockMiMoTTSServer 创建模拟 MiMo TTS API 的 HTTP 服务器。
func mockMiMoTTSServer(t *testing.T, handler http.HandlerFunc) *httptest.Server {
t.Helper()
return httptest.NewServer(handler)
}
// buildMiMoTTSResponse 构造 MiMo TTS 非流式响应 JSON。
func buildMiMoTTSResponse(audioData string) []byte {
resp := mimoTTSResponse{
Choices: []struct {
Message struct {
Audio struct {
Data string `json:"data"`
} `json:"audio"`
} `json:"message"`
}{
{
Message: struct {
Audio struct {
Data string `json:"data"`
} `json:"audio"`
}{
Audio: struct {
Data string `json:"data"`
}{Data: audioData},
},
},
},
}
data, _ := json.Marshal(resp)
return data
}
func TestMiMoService_SynthesizeStream_Success(t *testing.T) {
var callCount int32
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
atomic.AddInt32(&callCount, 1)
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)
}
// 验证 api-key 认证头
apiKey := r.Header.Get("api-key")
if apiKey != "test-key" {
t.Errorf("api-key = %q, want %q", apiKey, "test-key")
}
// 验证请求体
body, _ := io.ReadAll(r.Body)
var req mimoTTSRequest
if err := json.Unmarshal(body, &req); err != nil {
t.Errorf("unmarshal request: %v", err)
}
if req.Model != "mimo-v2.5-tts" {
t.Errorf("model = %q, want %q", req.Model, "mimo-v2.5-tts")
}
if len(req.Messages) != 1 || req.Messages[0].Role != "assistant" {
t.Errorf("expected 1 assistant message, got %d messages", len(req.Messages))
}
if req.Audio.Voice != "冰糖" {
t.Errorf("voice = %q, want %q", req.Audio.Voice, "冰糖")
}
if req.Audio.Format != "mp3" {
t.Errorf("format = %q, want %q", req.Audio.Format, "mp3")
}
// 返回假音频数据base64 编码)
audioB64 := base64.StdEncoding.EncodeToString([]byte("fake-mp3-data"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好", "世界", "")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{
Voice: "冰糖", OutputFmt: "mp3", SampleRate: 24000,
})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 应该有 3 个音频 chunk + 1 个 Final 标记
if len(chunks) != 4 {
t.Fatalf("got %d chunks, want 4", len(chunks))
}
// 验证前 3 个有音频数据IsLast 为 true每句结束
for i := 0; i < 3; i++ {
if string(chunks[i].Audio) != "fake-mp3-data" {
t.Errorf("chunk[%d].Audio = %q, want %q", i, string(chunks[i].Audio), "fake-mp3-data")
}
if !chunks[i].IsLast {
t.Errorf("chunk[%d].IsLast should be true (sentence end)", i)
}
if chunks[i].Final {
t.Errorf("chunk[%d].Final should be false", i)
}
}
// 验证最后一个是 Final整轮结束
if !chunks[3].Final {
t.Error("last chunk should be Final")
}
if chunks[3].Audio != nil {
t.Error("last chunk Audio should be nil")
}
// 验证调用了 3 次 API3 个句子)
if atomic.LoadInt32(&callCount) != 3 {
t.Errorf("API called %d times, want 3", callCount)
}
}
func TestMiMoService_SynthesizeStream_APIError(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, "internal error")
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
// 应该只有一个 Final chunk音频被跳过
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1 (Final only)", len(chunks))
}
if !chunks[0].Final {
t.Error("chunk should be Final")
}
}
func TestMiMoService_SynthesizeStream_Timeout(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
time.Sleep(3 * time.Second)
audioB64 := base64.StdEncoding.EncodeToString([]byte("late-mp3"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
// 1 秒超时
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 1, 30, zap.NewNop().Sugar())
textStream := sendSentences("很长的句子")
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
ch, err := svc.SynthesizeStream(ctx, textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 超时后音频被跳过,只有 Final
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1", len(chunks))
}
if !chunks[0].Final {
t.Error("chunk should be Final")
}
}
func TestMiMoService_SynthesizeStream_EmptyText(t *testing.T) {
var callCount int32
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
atomic.AddInt32(&callCount, 1)
audioB64 := base64.StdEncoding.EncodeToString([]byte("mp3"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
// 空句子应该被跳过
textStream := sendSentences("", "你好", "")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 只有 "你好" 应该被合成
if atomic.LoadInt32(&callCount) != 1 {
t.Errorf("API called %d times, want 1", callCount)
}
// 1 个音频IsLast: true+ 1 个 Final
if len(chunks) != 2 {
t.Fatalf("got %d chunks, want 2", len(chunks))
}
}
func TestMiMoService_SynthesizeStream_ContextCancelled(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
audioB64 := base64.StdEncoding.EncodeToString([]byte("mp3"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := make(chan string, 3)
textStream <- "第一句"
textStream <- "第二句"
textStream <- "第三句"
close(textStream)
ctx, cancel := context.WithCancel(context.Background())
// 立即取消
cancel()
ch, err := svc.SynthesizeStream(ctx, textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
// 消费 channel应该很快结束
var count int
for range ch {
count++
}
// 可能收到 0 个或 1 个 chunk取决于时序
t.Logf("received %d chunks after context cancel", count)
}
func TestMiMoService_SynthesizeStream_PartialFailure(t *testing.T) {
var callCount int32
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
n := atomic.AddInt32(&callCount, 1)
if n == 2 {
// 第二个句子失败
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, "error")
return
}
audioB64 := base64.StdEncoding.EncodeToString([]byte(fmt.Sprintf("mp3-%d", n)))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("第一句", "第二句", "第三句")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 2 个成功音频IsLast: true+ 1 个 Final第二句被跳过
if len(chunks) != 3 {
t.Fatalf("got %d chunks, want 3", len(chunks))
}
if !chunks[0].IsLast {
t.Error("first audio chunk should be IsLast")
}
if !chunks[1].IsLast {
t.Error("second audio chunk should be IsLast")
}
if !chunks[len(chunks)-1].Final {
t.Error("last chunk should be Final")
}
}
func TestMiMoService_SynthesizeStream_CustomVoice(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
var req mimoTTSRequest
if err := json.Unmarshal(body, &req); err != nil {
t.Errorf("unmarshal request: %v", err)
}
if req.Audio.Voice != "茉莉" {
t.Errorf("voice = %q, want %q", req.Audio.Voice, "茉莉")
}
audioB64 := base64.StdEncoding.EncodeToString([]byte("mp3"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{Voice: "茉莉"})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
for range ch {
}
}
func TestMiMoService_SynthesizeStream_DefaultVoice(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
var req mimoTTSRequest
if err := json.Unmarshal(body, &req); err != nil {
t.Errorf("unmarshal request: %v", err)
}
// 未指定 voice 时应使用默认 "冰糖"
if req.Audio.Voice != "冰糖" {
t.Errorf("voice = %q, want %q (default)", req.Audio.Voice, "冰糖")
}
audioB64 := base64.StdEncoding.EncodeToString([]byte("mp3"))
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(audioB64))
})
defer srv.Close()
// 不指定 voice
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
for range ch {
}
}
func TestMiMoService_SynthesizeStream_EmptyAudioData(t *testing.T) {
srv := mockMiMoTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
// 返回空音频数据
w.Header().Set("Content-Type", "application/json")
w.Write(buildMiMoTTSResponse(""))
})
defer srv.Close()
svc := NewMiMoService("test-key", "mimo-v2.5-tts", "冰糖", srv.URL, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{})
if err != nil {
t.Fatalf("SynthesizeStream() error: %v", err)
}
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
// 空音频数据导致错误,句子被跳过,只有 Final
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1", len(chunks))
}
if !chunks[0].Final {
t.Error("chunk should be Final")
}
}

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"
)
@@ -25,23 +27,19 @@ type OpenAIService struct {
}
// NewOpenAIService 创建 OpenAI TTS 服务。
func NewOpenAIService(apiKey, model, voice, endpoint string, speed float64, timeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
if model == "" {
model = "tts-1"
}
if voice == "" {
voice = "alloy"
}
// model、voice、endpoint 由 config 层保证非空。
func NewOpenAIService(apiKey, model, voice, endpoint string, speed float64, timeoutSec, httpClientTimeoutSec int, logger *zap.SugaredLogger) *OpenAIService {
if speed <= 0 {
speed = 1.0
}
if endpoint == "" {
endpoint = "https://api.openai.com/v1"
}
timeout := time.Duration(timeoutSec) * time.Second
if timeout <= 0 {
timeout = 5 * time.Second
}
httpClientTimeout := time.Duration(httpClientTimeoutSec) * time.Second
if httpClientTimeout <= 0 {
httpClientTimeout = 30 * time.Second
}
return &OpenAIService{
apiKey: apiKey,
model: model,
@@ -50,7 +48,7 @@ func NewOpenAIService(apiKey, model, voice, endpoint string, speed float64, time
endpoint: endpoint,
timeout: timeout,
logger: logger,
client: &http.Client{Timeout: 30 * time.Second},
client: &http.Client{Timeout: httpClientTimeout},
}
}
@@ -85,21 +83,25 @@ 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
}
select {
case ch <- Chunk{Audio: audio, IsLast: false}:
case ch <- Chunk{Audio: audio, IsLast: true, Final: false}:
case <-ctx.Done():
return
}
}
// textStream 关闭,发送 IsLast 标记
// textStream 关闭,发送 Final 标记
select {
case ch <- Chunk{Audio: nil, IsLast: true}:
case ch <- Chunk{Audio: nil, IsLast: false, Final: true}:
case <-ctx.Done():
}
}()

View File

@@ -58,7 +58,7 @@ func TestOpenAIService_SynthesizeStream_Success(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好", "世界", "")
@@ -74,24 +74,27 @@ func TestOpenAIService_SynthesizeStream_Success(t *testing.T) {
chunks = append(chunks, c)
}
// 应该有 3 个音频 chunk + 1 个 IsLast 标记
// 应该有 3 个音频 chunk + 1 个 Final 标记
if len(chunks) != 4 {
t.Fatalf("got %d chunks, want 4", len(chunks))
}
// 验证前 3 个有音频数据
// 验证前 3 个有音频数据IsLast 为 true每句结束
for i := 0; i < 3; i++ {
if string(chunks[i].Audio) != "fake-mp3-data" {
t.Errorf("chunk[%d].Audio = %q, want %q", i, string(chunks[i].Audio), "fake-mp3-data")
}
if chunks[i].IsLast {
t.Errorf("chunk[%d].IsLast should be false", i)
if !chunks[i].IsLast {
t.Errorf("chunk[%d].IsLast should be true (sentence end)", i)
}
if chunks[i].Final {
t.Errorf("chunk[%d].Final should be false", i)
}
}
// 验证最后一个是 IsLast
if !chunks[3].IsLast {
t.Error("last chunk should be IsLast")
// 验证最后一个是 Final整轮结束
if !chunks[3].Final {
t.Error("last chunk should be Final")
}
if chunks[3].Audio != nil {
t.Error("last chunk Audio should be nil")
@@ -110,7 +113,7 @@ func TestOpenAIService_SynthesizeStream_APIError(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")
@@ -119,17 +122,17 @@ func TestOpenAIService_SynthesizeStream_APIError(t *testing.T) {
t.Fatalf("SynthesizeStream() error: %v", err)
}
// 应该只有一个 IsLast chunk音频被跳过
// 应该只有一个 Final chunk音频被跳过
var chunks []Chunk
for c := range ch {
chunks = append(chunks, c)
}
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1 (IsLast only)", len(chunks))
t.Fatalf("got %d chunks, want 1 (Final only)", len(chunks))
}
if !chunks[0].IsLast {
t.Error("chunk should be IsLast")
if !chunks[0].Final {
t.Error("chunk should be Final")
}
}
@@ -142,7 +145,7 @@ func TestOpenAIService_SynthesizeStream_Timeout(t *testing.T) {
defer srv.Close()
// 1 秒超时
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 1, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 1, 30, zap.NewNop().Sugar())
textStream := sendSentences("很长的句子")
@@ -159,12 +162,12 @@ func TestOpenAIService_SynthesizeStream_Timeout(t *testing.T) {
chunks = append(chunks, c)
}
// 超时后音频被跳过,只有 IsLast
// 超时后音频被跳过,只有 Final
if len(chunks) != 1 {
t.Fatalf("got %d chunks, want 1", len(chunks))
}
if !chunks[0].IsLast {
t.Error("chunk should be IsLast")
if !chunks[0].Final {
t.Error("chunk should be Final")
}
}
@@ -177,7 +180,7 @@ func TestOpenAIService_SynthesizeStream_EmptyText(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
// 空句子应该被跳过
textStream := sendSentences("", "你好", "")
@@ -197,7 +200,7 @@ func TestOpenAIService_SynthesizeStream_EmptyText(t *testing.T) {
t.Errorf("API called %d times, want 1", callCount)
}
// 1 个音频 + 1 个 IsLast
// 1 个音频IsLast: true+ 1 个 Final
if len(chunks) != 2 {
t.Fatalf("got %d chunks, want 2", len(chunks))
}
@@ -210,7 +213,7 @@ func TestOpenAIService_SynthesizeStream_ContextCancelled(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
// 发送多个句子,但在第一个后取消
textStream := make(chan string, 3)
@@ -252,7 +255,7 @@ func TestOpenAIService_SynthesizeStream_PartialFailure(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("第一句", "第二句", "第三句")
@@ -266,12 +269,18 @@ func TestOpenAIService_SynthesizeStream_PartialFailure(t *testing.T) {
chunks = append(chunks, c)
}
// 2 个成功音频 + 1 个 IsLast(第二句被跳过)
// 2 个成功音频IsLast: true+ 1 个 Final(第二句被跳过)
if len(chunks) != 3 {
t.Fatalf("got %d chunks, want 3", len(chunks))
}
if !chunks[len(chunks)-1].IsLast {
t.Error("last chunk should be IsLast")
if !chunks[0].IsLast {
t.Error("first audio chunk should be IsLast")
}
if !chunks[1].IsLast {
t.Error("second audio chunk should be IsLast")
}
if !chunks[len(chunks)-1].Final {
t.Error("last chunk should be Final")
}
}
@@ -286,7 +295,7 @@ func TestOpenAIService_SynthesizeStream_CustomVoice(t *testing.T) {
})
defer srv.Close()
svc := NewOpenAIService("test-key", "", "alloy", srv.URL, 1.0, 5, zap.NewNop().Sugar())
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
textStream := sendSentences("你好")

View File

@@ -21,5 +21,6 @@ type Options struct {
// Chunk 一个音频片段。
type Chunk struct {
Audio []byte // MP3 音频数据(未 Base64 编码)
IsLast bool // 是否为最后一片
IsLast bool // 当前句子是否为最后一片(每句结束时为 true
Final bool // 整轮 TTS 是否结束(所有句子合成完毕后为 true此时 Audio 为 nil
}

View File

@@ -0,0 +1,244 @@
package api
import (
"errors"
"net/http"
"github.com/gin-gonic/gin"
"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 端点。
type AuthHandler struct {
authService auth.Service
tokenMgr *auth.TokenManager
}
// NewAuthHandler 创建 AuthHandler。
func NewAuthHandler(authService auth.Service, tokenMgr *auth.TokenManager) *AuthHandler {
return &AuthHandler{
authService: authService,
tokenMgr: tokenMgr,
}
}
// RegisterRoutes 注册认证相关路由到给定的路由组。
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup, limiter ratelimit.Limiter) {
authGroup := rg.Group("/auth")
{
// 注册和登录端点添加限流中间件(按 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)
}
}
// 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{
"code": apperr.CodeInvalidInput,
"message": "invalid request body",
})
return
}
if msg := validateCredentials(req.Username, req.Password); msg != "" {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": msg,
})
return
}
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{
"code": apperr.CodeInvalidInput,
"message": "invalid request body",
})
return
}
if msg := validateCredentials(req.Username, req.Password); msg != "" {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": msg,
})
return
}
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{
"code": apperr.CodeInvalidInput,
"message": "invalid request body",
})
return
}
if req.RefreshToken == "" {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": "refresh_token is required",
})
return
}
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 {
RefreshToken string `json:"refresh_token"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": "invalid request body",
})
return
}
if req.RefreshToken == "" {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": "refresh_token is required",
})
return
}
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",
})
return
}
log.Infow("logout success",
"user_id", userID)
c.JSON(http.StatusOK, gin.H{
"message": "logged out successfully",
})
}
// validateCredentials 校验用户名和密码格式。
// 返回空字符串表示校验通过,否则返回错误描述。
func validateCredentials(username, password string) string {
if len(username) > 64 {
return "username must not exceed 64 characters"
}
if len(password) < 8 || len(password) > 72 {
return "password must be 8-72 characters"
}
return ""
}
// handleAuthError 将 auth 层错误映射为 HTTP 响应。
func handleAuthError(c *gin.Context, err error) {
switch {
case errors.Is(err, auth.ErrUsernameTaken):
c.JSON(http.StatusConflict, gin.H{
"code": apperr.CodeUsernameTaken,
"message": "username already taken",
})
case errors.Is(err, auth.ErrInvalidCredentials):
c.JSON(http.StatusUnauthorized, gin.H{
"code": apperr.CodeInvalidCredentials,
"message": "invalid username or password",
})
case errors.Is(err, auth.ErrRefreshTokenUsed):
c.JSON(http.StatusUnauthorized, gin.H{
"code": apperr.CodeInvalidToken,
"message": "refresh token has been used or expired",
})
default:
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "internal server error",
})
}
}

View File

@@ -0,0 +1,324 @@
package api_test
import (
"bytes"
"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"
"github.com/hhs/camtalk/internal/api"
"github.com/hhs/camtalk/internal/auth"
)
// mockAuthService 实现 auth.Service 接口,用于 API 测试。
type mockAuthService struct {
RegisterFunc func(ctx context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error)
LoginFunc func(ctx context.Context, req auth.LoginRequest) (*auth.AuthResponse, error)
RefreshFunc func(ctx context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error)
LogoutFunc func(ctx context.Context, userID, refreshToken string) error
}
func (m *mockAuthService) Register(ctx context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error) {
return m.RegisterFunc(ctx, req)
}
func (m *mockAuthService) Login(ctx context.Context, req auth.LoginRequest) (*auth.AuthResponse, error) {
return m.LoginFunc(ctx, req)
}
func (m *mockAuthService) Refresh(ctx context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error) {
return m.RefreshFunc(ctx, req)
}
func (m *mockAuthService) Logout(ctx context.Context, userID, refreshToken string) error {
return m.LogoutFunc(ctx, userID, refreshToken)
}
// newTestRouter 创建带 AuthHandler 路由的测试 Gin 引擎。
func newTestRouter(svc auth.Service) *gin.Engine {
gin.SetMode(gin.TestMode)
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"), nil) // 测试时不启用限流
return r
}
// newTestRouterWithToken 创建带 AuthHandler 路由的测试引擎,同时返回 TokenManager 以便生成测试 token。
func newTestRouterWithToken(svc auth.Service) (*gin.Engine, *auth.TokenManager) {
gin.SetMode(gin.TestMode)
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"), nil) // 测试时不启用限流
return r, tm
}
func sampleAuthResponse() *auth.AuthResponse {
return &auth.AuthResponse{
User: auth.UserResponse{
ID: "user-123",
Username: "alice",
},
AccessToken: "access-token",
RefreshToken: "refresh-token",
}
}
// --- Register ---
func TestRegister_Success(t *testing.T) {
svc := &mockAuthService{
RegisterFunc: func(_ context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error) {
assert.Equal(t, "alice", req.Username)
assert.Equal(t, "password123", req.Password)
return sampleAuthResponse(), nil
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "password123"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusCreated, w.Code)
var resp auth.AuthResponse
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, "alice", resp.User.Username)
assert.NotEmpty(t, resp.AccessToken)
}
func TestRegister_InvalidInput_EmptyBody(t *testing.T) {
svc := &mockAuthService{}
r := newTestRouter(svc)
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "INVALID_INPUT")
}
func TestRegister_InvalidInput_UsernameTooShort(t *testing.T) {
svc := &mockAuthService{}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RegisterRequest{Username: "ab", Password: "password123"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "username must be 3-64 characters")
}
func TestRegister_InvalidInput_PasswordTooShort(t *testing.T) {
svc := &mockAuthService{}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "short"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "password must be 8-72 characters")
}
func TestRegister_UsernameTaken(t *testing.T) {
svc := &mockAuthService{
RegisterFunc: func(_ context.Context, _ auth.RegisterRequest) (*auth.AuthResponse, error) {
return nil, auth.ErrUsernameTaken
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "password123"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusConflict, w.Code)
assert.Contains(t, w.Body.String(), "USERNAME_TAKEN")
}
// --- Login ---
func TestLogin_Success(t *testing.T) {
svc := &mockAuthService{
LoginFunc: func(_ context.Context, req auth.LoginRequest) (*auth.AuthResponse, error) {
assert.Equal(t, "alice", req.Username)
assert.Equal(t, "password123", req.Password)
return sampleAuthResponse(), nil
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.LoginRequest{Username: "alice", Password: "password123"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp auth.AuthResponse
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, "alice", resp.User.Username)
}
func TestLogin_InvalidCredentials(t *testing.T) {
svc := &mockAuthService{
LoginFunc: func(_ context.Context, _ auth.LoginRequest) (*auth.AuthResponse, error) {
return nil, auth.ErrInvalidCredentials
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.LoginRequest{Username: "alice", Password: "wrong-password"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code)
assert.Contains(t, w.Body.String(), "INVALID_CREDENTIALS")
}
// --- Refresh ---
func TestRefresh_Success(t *testing.T) {
svc := &mockAuthService{
RefreshFunc: func(_ context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error) {
assert.Equal(t, "some-refresh-token", req.RefreshToken)
return sampleAuthResponse(), nil
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: "some-refresh-token"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestRefresh_MissingToken(t *testing.T) {
svc := &mockAuthService{}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: ""})
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "refresh_token is required")
}
func TestRefresh_UsedToken(t *testing.T) {
svc := &mockAuthService{
RefreshFunc: func(_ context.Context, _ auth.RefreshRequest) (*auth.AuthResponse, error) {
return nil, auth.ErrRefreshTokenUsed
},
}
r := newTestRouter(svc)
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: "used-token"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code)
assert.Contains(t, w.Body.String(), "INVALID_TOKEN")
}
// --- Logout ---
func TestLogout_Success(t *testing.T) {
logoutCalled := false
svc := &mockAuthService{
LogoutFunc: func(_ context.Context, userID, refreshToken string) error {
assert.Equal(t, "user-123", userID)
assert.Equal(t, "refresh-token-to-revoke", refreshToken)
logoutCalled = true
return nil
},
}
r, tm := newTestRouterWithToken(svc)
// 生成有效 token
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
body, _ := json.Marshal(map[string]string{"refresh_token": "refresh-token-to-revoke"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.True(t, logoutCalled)
assert.Contains(t, w.Body.String(), "logged out successfully")
}
func TestLogout_MissingAuth(t *testing.T) {
svc := &mockAuthService{}
r := newTestRouter(svc)
body, _ := json.Marshal(map[string]string{"refresh_token": "some-token"})
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code)
}
func TestLogout_MissingRefreshToken(t *testing.T) {
svc := &mockAuthService{}
r, tm := newTestRouterWithToken(svc)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
body, _ := json.Marshal(map[string]string{"refresh_token": ""})
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "refresh_token is required")
}

View File

@@ -0,0 +1,358 @@
// Package api 提供 REST API 处理函数。
package api
import (
"errors"
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"github.com/hhs/camtalk/internal/auth"
apperr "github.com/hhs/camtalk/internal/errors"
"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 端点。
type ConversationHandler struct {
sessionMgr session.Manager
tokenMgr *auth.TokenManager
msgRepo store.MessageRepository // 可选,为 nil 时 fallback 到内存查询
}
// NewConversationHandler 创建 ConversationHandler。
// msgRepo 可选,为 nil 时消息查询走内存。
func NewConversationHandler(sessionMgr session.Manager, tokenMgr *auth.TokenManager, msgRepo store.MessageRepository) *ConversationHandler {
return &ConversationHandler{
sessionMgr: sessionMgr,
tokenMgr: tokenMgr,
msgRepo: msgRepo,
}
}
// RegisterRoutes 注册对话相关路由到给定的路由组。所有端点需要认证。
func (h *ConversationHandler) RegisterRoutes(rg *gin.RouterGroup) {
conv := rg.Group("/conversations", auth.AuthMiddleware(h.tokenMgr))
{
conv.GET("", h.List)
conv.POST("", h.Create)
conv.GET("/:id", h.Get)
conv.PATCH("/:id", h.UpdateTitle)
conv.DELETE("/:id", h.Delete)
conv.GET("/:id/messages", h.GetMessages)
}
}
// 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"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page <= 0 {
page = 1
}
if size <= 0 || size > 100 {
size = 20
}
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",
})
return
}
c.JSON(http.StatusOK, gin.H{
"conversations": summaries,
"total": total,
"page": page,
"size": size,
})
}
// CreateConversationRequest POST /api/conversations 请求体。
type CreateConversationRequest struct {
Config *models.SessionConfig `json:"config,omitempty"`
}
// 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
_ = c.ShouldBindJSON(&req)
cfg := models.DefaultConfig()
if req.Config != nil {
cfg = *req.Config
}
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",
})
return
}
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",
})
return
}
log.Infow("conversation created",
"conversation_id", sess.ID,
"user_id", userID)
c.JSON(http.StatusCreated, gin.H{
"id": sess.ID,
"title": sess.Title,
"created_at": sess.CreatedAt,
"updated_at": sess.UpdatedAt,
})
}
// Get GET /api/conversations/:id — 获取对话详情。
func (h *ConversationHandler) Get(c *gin.Context) {
sessionID := c.Param("id")
sess, err := h.getSessionForUser(c, sessionID)
if err != nil {
return // getSessionForUser 已写入响应
}
c.JSON(http.StatusOK, gin.H{
"id": sess.ID,
"title": sess.Title,
"created_at": sess.CreatedAt,
"updated_at": sess.UpdatedAt,
"config": sess.Config,
})
}
// UpdateTitleRequest PATCH /api/conversations/:id 请求体。
type UpdateTitleRequest struct {
Title string `json:"title"`
}
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
if _, err := h.getSessionForUser(c, sessionID); err != nil {
return
}
var req UpdateTitleRequest
if err := c.ShouldBindJSON(&req); err != nil || req.Title == "" {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": "title is required",
})
return
}
if len([]rune(req.Title)) > 100 {
c.JSON(http.StatusBadRequest, gin.H{
"code": apperr.CodeInvalidInput,
"message": "title must be 100 characters or less",
})
return
}
if err := h.sessionMgr.UpdateTitle(c.Request.Context(), sessionID, req.Title); err != nil {
if errors.Is(err, session.ErrSessionNotFound) {
c.JSON(http.StatusNotFound, gin.H{
"code": apperr.CodeSessionNotFound,
"message": "conversation not found",
})
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",
})
return
}
c.JSON(http.StatusOK, gin.H{
"message": "title updated",
})
}
// Delete DELETE /api/conversations/:id — 删除对话。
func (h *ConversationHandler) Delete(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
if _, err := h.getSessionForUser(c, sessionID); err != nil {
return
}
if err := h.sessionMgr.Destroy(c.Request.Context(), sessionID); err != nil {
if errors.Is(err, session.ErrSessionNotFound) {
c.JSON(http.StatusNotFound, gin.H{
"code": apperr.CodeSessionNotFound,
"message": "conversation not found",
})
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",
})
return
}
c.Status(http.StatusNoContent)
}
// GetMessages GET /api/conversations/:id/messages — 获取对话消息列表。
//
// 查询参数:
// - limit: 返回消息数量上限,默认 50
// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
func (h *ConversationHandler) GetMessages(c *gin.Context) {
log := trace.FromContext(c.Request.Context())
sessionID := c.Param("id")
// 先校验归属
if _, err := h.getSessionForUser(c, sessionID); err != nil {
return
}
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "50"))
if limit <= 0 || limit > 200 {
limit = 50
}
beforeID, _ := strconv.ParseInt(c.DefaultQuery("before", "0"), 10, 64)
// 优先从 PostgreSQL 查询(支持持久化后的全量历史)
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",
})
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,
})
return
}
// fallback从内存查询
allMessages, err := h.sessionMgr.GetHistory(c.Request.Context(), sessionID, 0)
if err != nil {
if errors.Is(err, session.ErrSessionNotFound) {
c.JSON(http.StatusNotFound, gin.H{
"code": apperr.CodeSessionNotFound,
"message": "conversation not found",
})
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",
})
return
}
total := len(allMessages)
// beforeID > 0 时表示偏移量(兼容旧接口语义)
if beforeID > 0 && int(beforeID) <= total {
allMessages = allMessages[:beforeID]
}
// 取最后 limit 条
start := len(allMessages) - limit
if start < 0 {
start = 0
}
messages := allMessages[start:]
if messages == nil {
messages = []models.Message{}
}
c.JSON(http.StatusOK, gin.H{
"messages": messages,
"total": total,
})
}
// getSessionForUser 获取会话并校验当前用户是否有权限访问。
// 返回 404而非 403以避免信息泄露。
func (h *ConversationHandler) getSessionForUser(c *gin.Context, sessionID string) (*models.Session, error) {
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
if err != nil {
if errors.Is(err, session.ErrSessionNotFound) {
c.JSON(http.StatusNotFound, gin.H{
"code": apperr.CodeSessionNotFound,
"message": "conversation not found",
})
} else {
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "internal server error",
})
}
return nil, err
}
userID := c.GetString(auth.ContextKeyUserID)
if sess.UserID != userID {
c.JSON(http.StatusNotFound, gin.H{
"code": apperr.CodeSessionNotFound,
"message": "conversation not found",
})
return nil, errors.New("forbidden")
}
return sess, nil
}

View File

@@ -0,0 +1,575 @@
package api_test
import (
"bytes"
"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"
"github.com/hhs/camtalk/internal/api"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/session"
)
// mockSessionManager 实现 session.Manager 接口,用于 ConversationHandler 测试。
type mockSessionManager struct {
CreateFunc func(ctx context.Context, userID string, config models.SessionConfig) (string, error)
GetFunc func(ctx context.Context, sessionID string) (*models.Session, error)
UpdateConfigFunc func(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error
UpdateTitleFunc func(ctx context.Context, sessionID string, title string) error
ListByUserFunc func(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error)
GetHistoryFunc func(ctx context.Context, sessionID string, limit int) ([]models.Message, error)
AppendMessageFunc func(ctx context.Context, sessionID string, msg models.Message) error
SetActiveRequestFunc func(ctx context.Context, sessionID string, requestID string) error
GetActiveRequestIDFunc func(ctx context.Context, sessionID string) (string, error)
ClearActiveRequestFunc func(ctx context.Context, sessionID string) error
TouchFunc func(ctx context.Context, sessionID string) error
DestroyFunc func(ctx context.Context, sessionID string) error
ActiveCountFunc func() int
}
func (m *mockSessionManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
return m.CreateFunc(ctx, userID, config)
}
func (m *mockSessionManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
return m.GetFunc(ctx, sessionID)
}
func (m *mockSessionManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
return m.UpdateConfigFunc(ctx, sessionID, patch)
}
func (m *mockSessionManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
return m.UpdateTitleFunc(ctx, sessionID, title)
}
func (m *mockSessionManager) ListByUser(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) {
return m.ListByUserFunc(ctx, userID, page, size)
}
func (m *mockSessionManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
return m.GetHistoryFunc(ctx, sessionID, limit)
}
func (m *mockSessionManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
return m.AppendMessageFunc(ctx, sessionID, msg)
}
func (m *mockSessionManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
return m.SetActiveRequestFunc(ctx, sessionID, requestID)
}
func (m *mockSessionManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
return m.GetActiveRequestIDFunc(ctx, sessionID)
}
func (m *mockSessionManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
return m.ClearActiveRequestFunc(ctx, sessionID)
}
func (m *mockSessionManager) Touch(ctx context.Context, sessionID string) error {
return m.TouchFunc(ctx, sessionID)
}
func (m *mockSessionManager) Destroy(ctx context.Context, sessionID string) error {
return m.DestroyFunc(ctx, sessionID)
}
func (m *mockSessionManager) ActiveCount() int {
return m.ActiveCountFunc()
}
// newConvTestRouter 创建带 ConversationHandler 路由的测试引擎,同时返回 TokenManager。
func newConvTestRouter(mgr session.Manager) (*gin.Engine, *auth.TokenManager) {
gin.SetMode(gin.TestMode)
r := gin.New()
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
h := api.NewConversationHandler(mgr, tm, nil)
h.RegisterRoutes(r.Group("/api"))
return r, tm
}
// --- List ---
func TestConversationList_Success(t *testing.T) {
now := time.Now()
mgr := &mockSessionManager{
ListByUserFunc: func(_ context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) {
assert.Equal(t, "user-123", userID)
assert.Equal(t, 1, page)
assert.Equal(t, 20, size)
return []session.ConversationSummary{
{ID: "sess-1", Title: "对话一", MessageCount: 3, UpdatedAt: now},
{ID: "sess-2", Title: "对话二", MessageCount: 1, UpdatedAt: now.Add(-time.Hour)},
}, 2, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, float64(2), resp["total"])
convs := resp["conversations"].([]interface{})
assert.Len(t, convs, 2)
}
func TestConversationList_WithPagination(t *testing.T) {
mgr := &mockSessionManager{
ListByUserFunc: func(_ context.Context, _ string, page, size int) ([]session.ConversationSummary, int, error) {
assert.Equal(t, 2, page)
assert.Equal(t, 10, size)
return []session.ConversationSummary{}, 0, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations?page=2&size=10", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestConversationList_MissingAuth(t *testing.T) {
mgr := &mockSessionManager{}
r, _ := newConvTestRouter(mgr)
req := httptest.NewRequest(http.MethodGet, "/api/conversations", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusUnauthorized, w.Code)
}
// --- Create ---
func TestConversationCreate_Success(t *testing.T) {
createdID := "new-session-id"
now := time.Now()
mgr := &mockSessionManager{
CreateFunc: func(_ context.Context, userID string, cfg models.SessionConfig) (string, error) {
assert.Equal(t, "user-123", userID)
return createdID, nil
},
GetFunc: func(_ context.Context, sessionID string) (*models.Session, error) {
assert.Equal(t, createdID, sessionID)
return &models.Session{
ID: createdID,
UserID: "user-123",
Title: models.DefaultSessionTitle,
CreatedAt: now,
UpdatedAt: now,
Config: models.DefaultConfig(),
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodPost, "/api/conversations", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusCreated, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, createdID, resp["id"])
assert.Equal(t, models.DefaultSessionTitle, resp["title"])
}
func TestConversationCreate_WithConfig(t *testing.T) {
mgr := &mockSessionManager{
CreateFunc: func(_ context.Context, _ string, cfg models.SessionConfig) (string, error) {
assert.False(t, cfg.TTSEnabled)
assert.Equal(t, "high", cfg.DetailLevel)
return "sess-1", nil
},
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{
ID: "sess-1",
UserID: "user-123",
Title: models.DefaultSessionTitle,
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
body, _ := json.Marshal(api.CreateConversationRequest{
Config: &models.SessionConfig{TTSEnabled: false, DetailLevel: "high", Language: "zh-CN"},
})
req := httptest.NewRequest(http.MethodPost, "/api/conversations", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusCreated, w.Code)
}
// --- Get ---
func TestConversationGet_Success(t *testing.T) {
now := time.Now()
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, sessionID string) (*models.Session, error) {
assert.Equal(t, "sess-1", sessionID)
return &models.Session{
ID: "sess-1",
UserID: "user-123",
Title: "我的对话",
CreatedAt: now,
UpdatedAt: now,
Config: models.DefaultConfig(),
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, "我的对话", resp["title"])
}
func TestConversationGet_NotFound(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return nil, session.ErrSessionNotFound
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/nonexistent", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
assert.Contains(t, w.Body.String(), "SESSION_NOT_FOUND")
}
func TestConversationGet_Forbidden(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
// 会话属于另一个用户
return &models.Session{
ID: "sess-1",
UserID: "other-user",
Title: "他人对话",
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
// 返回 404 而非 403避免信息泄露
assert.Equal(t, http.StatusNotFound, w.Code)
assert.Contains(t, w.Body.String(), "SESSION_NOT_FOUND")
}
// --- UpdateTitle ---
func TestConversationUpdateTitle_Success(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
UpdateTitleFunc: func(_ context.Context, sessionID, title string) error {
assert.Equal(t, "sess-1", sessionID)
assert.Equal(t, "新标题", title)
return nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
body, _ := json.Marshal(api.UpdateTitleRequest{Title: "新标题"})
req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Contains(t, w.Body.String(), "title updated")
}
func TestConversationUpdateTitle_EmptyTitle(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
body, _ := json.Marshal(api.UpdateTitleRequest{Title: ""})
req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "title is required")
}
func TestConversationUpdateTitle_TooLong(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
longTitle := ""
for i := 0; i < 101; i++ {
longTitle += "测"
}
body, _ := json.Marshal(api.UpdateTitleRequest{Title: longTitle})
req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, w.Body.String(), "title must be 100 characters or less")
}
// --- Delete ---
func TestConversationDelete_Success(t *testing.T) {
destroyCalled := false
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
DestroyFunc: func(_ context.Context, sessionID string) error {
assert.Equal(t, "sess-1", sessionID)
destroyCalled = true
return nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodDelete, "/api/conversations/sess-1", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusNoContent, w.Code)
assert.True(t, destroyCalled)
}
func TestConversationDelete_Forbidden(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "other-user"}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodDelete, "/api/conversations/sess-1", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
// --- GetMessages ---
func TestConversationGetMessages_Success(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
GetHistoryFunc: func(_ context.Context, sessionID string, limit int) ([]models.Message, error) {
assert.Equal(t, "sess-1", sessionID)
assert.Equal(t, 0, limit) // 获取全量
return []models.Message{
{Role: "user", Content: "你好"},
{Role: "assistant", Content: "你好!有什么可以帮助你的吗?"},
{Role: "user", Content: "今天天气怎么样?"},
{Role: "assistant", Content: "今天天气不错!"},
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, float64(4), resp["total"])
msgs := resp["messages"].([]interface{})
assert.Len(t, msgs, 4)
}
func TestConversationGetMessages_WithLimit(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
GetHistoryFunc: func(_ context.Context, _ string, _ int) ([]models.Message, error) {
return []models.Message{
{Role: "user", Content: "消息1"},
{Role: "assistant", Content: "回复1"},
{Role: "user", Content: "消息2"},
{Role: "assistant", Content: "回复2"},
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages?limit=2", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
msgs := resp["messages"].([]interface{})
assert.Len(t, msgs, 2)
}
func TestConversationGetMessages_WithBefore(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "user-123"}, nil
},
GetHistoryFunc: func(_ context.Context, _ string, _ int) ([]models.Message, error) {
return []models.Message{
{Role: "user", Content: "消息1"},
{Role: "assistant", Content: "回复1"},
{Role: "user", Content: "消息2"},
{Role: "assistant", Content: "回复2"},
}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages?before=2&limit=10", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp map[string]interface{}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
// before=2 表示取 index 0..1,共 2 条
msgs := resp["messages"].([]interface{})
assert.Len(t, msgs, 2)
}
func TestConversationGetMessages_Forbidden(t *testing.T) {
mgr := &mockSessionManager{
GetFunc: func(_ context.Context, _ string) (*models.Session, error) {
return &models.Session{ID: "sess-1", UserID: "other-user"}, nil
},
}
r, tm := newConvTestRouter(mgr)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages", nil)
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}

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)
@@ -36,8 +39,10 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
cfg = *req.Config
}
sessionID, err := h.sessionMgr.Create(c.Request.Context(), cfg)
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

@@ -0,0 +1,134 @@
package auth
import (
"crypto/sha256"
"encoding/hex"
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/google/uuid"
)
// 自定义错误。
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"`
TokenType string `json:"token_type"`
jwt.RegisteredClaims
}
// TokenManager JWT 令牌管理器。
type TokenManager struct {
secret []byte
accessTTL time.Duration
refreshTTL time.Duration
}
// NewTokenManager 创建 TokenManager。
// secret: JWT 签名密钥accessTTL/refreshTTL: 令牌有效期。
func NewTokenManager(secret string, accessTTL, refreshTTL time.Duration) *TokenManager {
return &TokenManager{
secret: []byte(secret),
accessTTL: accessTTL,
refreshTTL: refreshTTL,
}
}
// GeneratePair 生成 access + refresh 令牌对。
func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh string, err error) {
now := time.Now()
// access token
accessClaims := &Claims{
UserID: userID,
Username: username,
TokenType: TokenTypeAccess,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)),
IssuedAt: jwt.NewNumericDate(now),
Issuer: "camtalk",
},
}
accessTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims)
access, err = accessTkn.SignedString(tm.secret)
if err != nil {
return "", "", err
}
// refresh token含唯一 token_id 用于 DB 关联)
tokenID := uuid.New().String()
refreshClaims := &Claims{
UserID: userID,
Username: username,
TokenType: TokenTypeRefresh,
RegisteredClaims: jwt.RegisteredClaims{
ID: tokenID,
ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)),
IssuedAt: jwt.NewNumericDate(now),
Issuer: "camtalk",
},
}
refreshTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims)
refresh, err = refreshTkn.SignedString(tm.secret)
return
}
// ValidateAccess 校验 access token 并返回 Claims。
func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) {
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) {
claims, err := tm.validate(tokenStr)
if err != nil {
return nil, err
}
if claims.TokenType != TokenTypeRefresh {
return nil, ErrInvalidToken
}
return claims, nil
}
// validate 解析并校验 JWT。
func (tm *TokenManager) validate(tokenStr string) (*Claims, error) {
token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, ErrInvalidToken
}
return tm.secret, nil
})
if err != nil {
return nil, ErrInvalidToken
}
claims, ok := token.Claims.(*Claims)
if !ok || !token.Valid {
return nil, ErrInvalidToken
}
return claims, nil
}
// HashToken 对 token 做 SHA256 哈希,用于 DB 存储。
func HashToken(token string) string {
h := sha256.Sum256([]byte(token))
return hex.EncodeToString(h[:])
}

View File

@@ -0,0 +1,165 @@
package auth
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestGeneratePair_ReturnsNonEmptyTokens(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)
assert.NotEmpty(t, access)
assert.NotEmpty(t, refresh)
assert.NotEqual(t, access, refresh)
}
func TestValidateAccess_ValidToken(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)
claims, err := tm.ValidateAccess(access)
require.NoError(t, err)
assert.Equal(t, "user-123", claims.UserID)
assert.Equal(t, "alice", claims.Username)
assert.Equal(t, "camtalk", claims.Issuer)
}
func TestValidateRefresh_ValidToken(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)
claims, err := tm.ValidateRefresh(refresh)
require.NoError(t, err)
assert.Equal(t, "user-123", claims.UserID)
assert.Equal(t, "alice", claims.Username)
assert.NotEmpty(t, claims.ID) // refresh token 应含唯一 ID
}
func TestValidateAccess_ExpiredToken(t *testing.T) {
// 使用极短的 TTL
tm := NewTokenManager("test-secret-key", -1*time.Second, -1*time.Second)
access, _, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm.ValidateAccess(access)
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_WrongSecret(t *testing.T) {
tm1 := NewTokenManager("secret-1", 15*time.Minute, 7*24*time.Hour)
tm2 := NewTokenManager("secret-2", 15*time.Minute, 7*24*time.Hour)
access, _, err := tm1.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm2.ValidateAccess(access)
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_InvalidFormat(t *testing.T) {
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
_, err := tm.ValidateAccess("not-a-valid-token")
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestValidateAccess_EmptyString(t *testing.T) {
tm := NewTokenManager("test-secret-key", 15*time.Minute, 7*24*time.Hour)
_, err := tm.ValidateAccess("")
assert.ErrorIs(t, err, ErrInvalidToken)
}
func TestHashToken_Deterministic(t *testing.T) {
hash1 := HashToken("some-token-value")
hash2 := HashToken("some-token-value")
assert.Equal(t, hash1, hash2)
assert.Len(t, hash1, 64) // SHA256 hex = 64 chars
}
func TestHashToken_DifferentInputsDifferentHashes(t *testing.T) {
hash1 := HashToken("token-a")
hash2 := HashToken("token-b")
assert.NotEqual(t, hash1, hash2)
}
func TestValidateRefresh_ExpiredToken(t *testing.T) {
tm := NewTokenManager("test-secret-key", -1*time.Minute, -1*time.Minute)
_, refresh, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
_, err = tm.ValidateRefresh(refresh)
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
tm := NewTokenManager("test-secret-key", accessTTL, refreshTTL)
before := time.Now()
access, refresh, err := tm.GeneratePair("user-123", "alice")
require.NoError(t, err)
after := time.Now()
// 校验 access token 有效期范围
accessClaims, err := tm.ValidateAccess(access)
require.NoError(t, err)
assert.True(t, accessClaims.ExpiresAt.Time.After(before.Add(accessTTL).Add(-1*time.Second)))
assert.True(t, accessClaims.ExpiresAt.Time.Before(after.Add(accessTTL).Add(1*time.Second)))
// 校验 refresh token 有效期范围
refreshClaims, err := tm.ValidateRefresh(refresh)
require.NoError(t, err)
assert.True(t, refreshClaims.ExpiresAt.Time.After(before.Add(refreshTTL).Add(-1*time.Second)))
assert.True(t, refreshClaims.ExpiresAt.Time.Before(after.Add(refreshTTL).Add(1*time.Second)))
}

View File

@@ -0,0 +1,70 @@
package auth
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/hhs/camtalk/internal/trace"
)
// contextKey 用于在 Gin context 中存储 Claims 的 key。
const (
ContextKeyUserID = "user_id"
ContextKeyUsername = "username"
)
// AuthMiddleware 返回 Gin 中间件,从 Authorization: Bearer <token> 提取并校验 JWT。
// 校验成功后将 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",
})
return
}
// 提取 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",
})
return
}
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",
})
return
}
// 将用户信息写入 context
c.Set(ContextKeyUserID, claims.UserID)
c.Set(ContextKeyUsername, claims.Username)
c.Next()
}
}

View File

@@ -0,0 +1,19 @@
package auth
import "golang.org/x/crypto/bcrypt"
const bcryptCost = 10
// HashPassword 使用 bcrypt 对密码进行哈希。
func HashPassword(password string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcryptCost)
if err != nil {
return "", err
}
return string(hash), nil
}
// CheckPassword 校验密码与哈希是否匹配。
func CheckPassword(hashedPassword, password string) error {
return bcrypt.CompareHashAndPassword([]byte(hashedPassword), []byte(password))
}

View File

@@ -0,0 +1,224 @@
package auth
import (
"context"
"errors"
"time"
"github.com/hhs/camtalk/internal/store"
)
// 自定义业务错误。
var (
ErrUsernameTaken = errors.New("username already taken")
ErrInvalidCredentials = errors.New("invalid username or password")
ErrRefreshTokenUsed = errors.New("refresh token has been used or expired")
)
// RegisterRequest 注册请求。
type RegisterRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
// LoginRequest 登录请求。
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
// RefreshRequest 刷新令牌请求。
type RefreshRequest struct {
RefreshToken string `json:"refresh_token"`
}
// AuthResponse 认证响应。
type AuthResponse struct {
User UserResponse `json:"user"`
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
}
// UserResponse 用户信息响应。
type UserResponse struct {
ID string `json:"id"`
Username string `json:"username"`
CreatedAt time.Time `json:"created_at"`
}
// Service 认证业务接口。
type Service interface {
Register(ctx context.Context, req RegisterRequest) (*AuthResponse, error)
Login(ctx context.Context, req LoginRequest) (*AuthResponse, error)
Refresh(ctx context.Context, req RefreshRequest) (*AuthResponse, error)
Logout(ctx context.Context, userID, refreshToken string) error
}
// authService 认证服务实现。
type authService struct {
tokenMgr *TokenManager
userRepo store.UserRepository
}
// NewAuthService 创建认证服务。
func NewAuthService(tokenMgr *TokenManager, userRepo store.UserRepository) Service {
return &authService{
tokenMgr: tokenMgr,
userRepo: userRepo,
}
}
// Register 用户注册。
func (s *authService) Register(ctx context.Context, req RegisterRequest) (*AuthResponse, error) {
// 检查用户名是否已存在
_, err := s.userRepo.FindByUsername(ctx, req.Username)
if err == nil {
return nil, ErrUsernameTaken
}
if !errors.Is(err, store.ErrUserNotFound) {
return nil, err
}
// 哈希密码
hash, err := HashPassword(req.Password)
if err != nil {
return nil, err
}
// 创建用户
userID, err := s.userRepo.Create(ctx, req.Username, hash)
if err != nil {
if errors.Is(err, store.ErrUsernameTaken) {
return nil, ErrUsernameTaken
}
return nil, err
}
// 生成令牌对
access, refresh, err := s.tokenMgr.GeneratePair(userID, req.Username)
if err != nil {
return nil, err
}
// 保存 refresh token hash 到 DB
if err := s.saveRefreshToken(ctx, userID, refresh); err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: userID,
Username: req.Username,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Login 用户登录。
func (s *authService) Login(ctx context.Context, req LoginRequest) (*AuthResponse, error) {
user, err := s.userRepo.FindByUsername(ctx, req.Username)
if err != nil {
if errors.Is(err, store.ErrUserNotFound) {
return nil, ErrInvalidCredentials
}
return nil, err
}
// 校验密码
if err := CheckPassword(user.PasswordHash, req.Password); err != nil {
return nil, ErrInvalidCredentials
}
// 生成令牌对
access, refresh, err := s.tokenMgr.GeneratePair(user.ID, user.Username)
if err != nil {
return nil, err
}
// 保存 refresh token hash
if err := s.saveRefreshToken(ctx, user.ID, refresh); err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: user.ID,
Username: user.Username,
CreatedAt: user.CreatedAt,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Refresh 刷新令牌Refresh Token Rotation
func (s *authService) Refresh(ctx context.Context, req RefreshRequest) (*AuthResponse, error) {
// 校验 refresh token
claims, err := s.tokenMgr.ValidateRefresh(req.RefreshToken)
if err != nil {
return nil, ErrRefreshTokenUsed
}
tokenHash := HashToken(req.RefreshToken)
// 查找 DB 中的 token hash确认未被使用
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
}
// 确认 token 归属的用户与 claims 一致
if userID != claims.UserID {
return nil, ErrRefreshTokenUsed
}
// 删除旧 refresh tokenrotation
_ = s.userRepo.DeleteRefreshToken(ctx, tokenHash)
// 生成新的令牌对
access, refresh, err := s.tokenMgr.GeneratePair(claims.UserID, claims.Username)
if err != nil {
return nil, err
}
// 保存新 refresh token
if err := s.saveRefreshToken(ctx, claims.UserID, refresh); err != nil {
return nil, err
}
// 查用户信息
user, err := s.userRepo.FindByID(ctx, claims.UserID)
if err != nil {
return nil, err
}
return &AuthResponse{
User: UserResponse{
ID: user.ID,
Username: user.Username,
CreatedAt: user.CreatedAt,
},
AccessToken: access,
RefreshToken: refresh,
}, nil
}
// Logout 登出,删除 refresh token。
func (s *authService) Logout(ctx context.Context, userID, refreshToken string) error {
tokenHash := HashToken(refreshToken)
return s.userRepo.DeleteRefreshToken(ctx, tokenHash)
}
// saveRefreshToken 将 refresh token 的 hash 保存到 DB。
func (s *authService) saveRefreshToken(ctx context.Context, userID, refreshToken string) error {
tokenHash := HashToken(refreshToken)
expiresAt := time.Now().Add(s.tokenMgr.refreshTTL)
return s.userRepo.SaveRefreshToken(ctx, userID, tokenHash, expiresAt)
}

View File

@@ -0,0 +1,235 @@
package auth_test
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/store"
)
// newTestService 创建测试用的 AuthService + MemUserRepository。
func newTestService(t *testing.T) (auth.Service, *store.MemUserRepository) {
t.Helper()
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
repo := store.NewMemUserRepository()
svc := auth.NewAuthService(tm, repo)
return svc, repo
}
// --- Register ---
func TestRegister_Success(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
resp, err := svc.Register(ctx, auth.RegisterRequest{
Username: "alice",
Password: "password123",
})
require.NoError(t, err)
assert.NotEmpty(t, resp.User.ID)
assert.Equal(t, "alice", resp.User.Username)
assert.NotEmpty(t, resp.AccessToken)
assert.NotEmpty(t, resp.RefreshToken)
}
func TestRegister_DuplicateUsername(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
_, err := svc.Register(ctx, auth.RegisterRequest{
Username: "alice",
Password: "password123",
})
require.NoError(t, err)
// 同名再次注册
_, err = svc.Register(ctx, auth.RegisterRequest{
Username: "alice",
Password: "another-password",
})
assert.ErrorIs(t, err, auth.ErrUsernameTaken)
}
// --- Login ---
func TestLogin_Success(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
// 先注册
_, err := svc.Register(ctx, auth.RegisterRequest{
Username: "bob",
Password: "password123",
})
require.NoError(t, err)
// 登录
resp, err := svc.Login(ctx, auth.LoginRequest{
Username: "bob",
Password: "password123",
})
require.NoError(t, err)
assert.Equal(t, "bob", resp.User.Username)
assert.NotEmpty(t, resp.AccessToken)
assert.NotEmpty(t, resp.RefreshToken)
}
func TestLogin_WrongPassword(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
_, err := svc.Register(ctx, auth.RegisterRequest{
Username: "bob",
Password: "password123",
})
require.NoError(t, err)
_, err = svc.Login(ctx, auth.LoginRequest{
Username: "bob",
Password: "wrong-password",
})
assert.ErrorIs(t, err, auth.ErrInvalidCredentials)
}
func TestLogin_UserNotFound(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
_, err := svc.Login(ctx, auth.LoginRequest{
Username: "nonexistent",
Password: "password123",
})
assert.ErrorIs(t, err, auth.ErrInvalidCredentials)
}
// --- Refresh ---
func TestRefresh_Success(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
// 注册
regResp, err := svc.Register(ctx, auth.RegisterRequest{
Username: "charlie",
Password: "password123",
})
require.NoError(t, err)
// 刷新
refreshResp, err := svc.Refresh(ctx, auth.RefreshRequest{
RefreshToken: regResp.RefreshToken,
})
require.NoError(t, err)
assert.Equal(t, "charlie", refreshResp.User.Username)
assert.NotEmpty(t, refreshResp.AccessToken)
assert.NotEmpty(t, refreshResp.RefreshToken)
// 新旧 refresh token 应不同rotation
assert.NotEqual(t, regResp.RefreshToken, refreshResp.RefreshToken)
}
func TestRefresh_UsedTokenFails(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
regResp, err := svc.Register(ctx, auth.RegisterRequest{
Username: "charlie",
Password: "password123",
})
require.NoError(t, err)
// 第一次刷新
_, err = svc.Refresh(ctx, auth.RefreshRequest{
RefreshToken: regResp.RefreshToken,
})
require.NoError(t, err)
// 用旧 token 再次刷新 → 应失败
_, err = svc.Refresh(ctx, auth.RefreshRequest{
RefreshToken: regResp.RefreshToken,
})
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
}
func TestRefresh_InvalidToken(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
_, err := svc.Refresh(ctx, auth.RefreshRequest{
RefreshToken: "completely-invalid-token",
})
assert.ErrorIs(t, err, auth.ErrRefreshTokenUsed)
}
// --- Logout ---
func TestLogout_Success(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
regResp, err := svc.Register(ctx, auth.RegisterRequest{
Username: "dave",
Password: "password123",
})
require.NoError(t, err)
// 登出
err = svc.Logout(ctx, regResp.User.ID, regResp.RefreshToken)
require.NoError(t, err)
// 登出后 refresh token 应失效
_, err = svc.Refresh(ctx, auth.RefreshRequest{
RefreshToken: regResp.RefreshToken,
})
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,8 +2,7 @@ package config
import (
"fmt"
"os"
"strings"
"path/filepath"
"github.com/joho/godotenv"
"github.com/spf13/viper"
@@ -11,12 +10,21 @@ import (
// Config 应用配置。
type Config struct {
App AppConfig `mapstructure:"app"`
Server ServerConfig `mapstructure:"server"`
Redis RedisConfig `mapstructure:"redis"`
AI AIConfig `mapstructure:"ai"`
Storage StorageConfig `mapstructure:"storage"`
Log LogConfig `mapstructure:"log"`
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 会话管理配置。
type SessionConfig struct {
TTL int `mapstructure:"ttl"` // 会话过期时间(分钟)
MaxHistory int `mapstructure:"max_history"` // 对话历史上限(条)
}
type AppConfig struct {
@@ -25,10 +33,14 @@ type AppConfig struct {
}
type ServerConfig struct {
Host string `mapstructure:"host"`
Port int `mapstructure:"port"`
ReadTimeout int `mapstructure:"read_timeout"`
WriteTimeout int `mapstructure:"write_timeout"`
Host string `mapstructure:"host"`
Port int `mapstructure:"port"`
ReadTimeout int `mapstructure:"read_timeout"`
WriteTimeout int `mapstructure:"write_timeout"`
HeartbeatInterval int `mapstructure:"heartbeat_interval"` // 心跳检查间隔(秒)
HeartbeatTimeout int `mapstructure:"heartbeat_timeout"` // 心跳超时(秒)
ShutdownTimeout int `mapstructure:"shutdown_timeout"` // 优雅关闭超时(秒)
AllowedOrigins []string `mapstructure:"allowed_origins"` // CORS 允许的来源,空表示允许所有
}
// Addr 返回 host:port 地址。
@@ -49,112 +61,212 @@ type AIConfig struct {
}
type STTConfig struct {
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Endpoint string `mapstructure:"endpoint"`
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Endpoint string `mapstructure:"endpoint"`
Timeout int `mapstructure:"timeout"` // STT 超时(秒)
HTTPClientTimeout int `mapstructure:"http_client_timeout"` // HTTP 客户端超时(秒)
}
type LLMConfig struct {
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Endpoint string `mapstructure:"endpoint"`
Timeout int `mapstructure:"timeout"`
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Endpoint string `mapstructure:"endpoint"`
Timeout int `mapstructure:"timeout"`
HTTPClientTimeout int `mapstructure:"http_client_timeout"` // HTTP 客户端超时(秒)
}
type TTSConfig struct {
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Voice string `mapstructure:"voice"`
Speed float64 `mapstructure:"speed"`
Endpoint string `mapstructure:"endpoint"`
Timeout int `mapstructure:"timeout"`
Provider string `mapstructure:"provider"`
APIKey string `mapstructure:"api_key"`
Model string `mapstructure:"model"`
Voice string `mapstructure:"voice"`
Speed float64 `mapstructure:"speed"`
Endpoint string `mapstructure:"endpoint"`
Timeout int `mapstructure:"timeout"`
HTTPClientTimeout int `mapstructure:"http_client_timeout"` // HTTP 客户端超时(秒)
OutputFormat string `mapstructure:"output_format"` // 输出格式mp3/wav
SampleRate int `mapstructure:"sample_rate"` // 输出采样率
}
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"`
}
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml
func Load() (*Config, error) {
// AuthConfig 认证配置
type AuthConfig struct {
JWTSecret string `mapstructure:"jwt_secret"` // JWT 签名密钥,必须通过环境变量 CAMTALK_AUTH_JWT_SECRET 设置
AccessTTL int `mapstructure:"access_ttl"` // Access Token 过期时间(分钟),默认 15
RefreshTTL int `mapstructure:"refresh_ttl"` // Refresh Token 过期时间(分钟),默认 100807天
}
// RateLimitConfig 限流配置。
type RateLimitConfig struct {
Enabled bool `mapstructure:"enabled"`
Query BucketConfig `mapstructure:"query"`
Login BucketConfig `mapstructure:"login"`
Register BucketConfig `mapstructure:"register"`
}
// BucketConfig 令牌桶配置。
type BucketConfig struct {
Capacity int `mapstructure:"capacity"` // 桶容量(突发上限)
Rate float64 `mapstructure:"rate"` // 每秒填充令牌数
}
// Load 加载配置。优先级:环境变量 > config.{env}.yaml > config.yaml > 默认值。
// 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(filepath.Join(workDir, "config")) // 配置文件在 config/ 目录下
v.AddConfigPath(workDir) // 兼容旧路径
// 默认值
v.SetDefault("app.env", "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("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.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.tts.provider", "openai")
v.SetDefault("ai.tts.model", "tts-1")
v.SetDefault("ai.tts.voice", "alloy")
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("storage.driver", "memory")
v.SetDefault("log.level", "info")
v.SetDefault("log.format", "console")
// 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() // 文件不存在也不报错
}
// 加载 .env 文件(不覆盖已有环境变量
// 按优先级尝试:当前目录、上级目录(兼容从 backend/ 或项目根目录启动)
_ = godotenv.Load()
_ = godotenv.Load("../.env")
// 环境变量覆盖
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

@@ -14,6 +14,12 @@ const (
CodeSTTError = "STT_ERROR"
CodeTTSError = "TTS_ERROR"
CodeInternalError = "INTERNAL_ERROR"
// 认证相关错误码
CodeUsernameTaken = "USERNAME_TAKEN"
CodeInvalidCredentials = "INVALID_CREDENTIALS"
CodeInvalidToken = "INVALID_TOKEN"
CodeInvalidInput = "INVALID_INPUT"
)
// Sender 定义发送 WS 错误消息的接口,便于测试 mock。

View File

@@ -5,7 +5,10 @@ import "time"
// Session 会话。
type Session struct {
ID string `json:"session_id"`
UserID string `json:"user_id,omitempty"`
Title string `json:"title"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
Config SessionConfig `json:"config"`
}
@@ -14,11 +17,15 @@ type SessionConfig struct {
TTSEnabled bool `json:"tts_enabled"`
DetailLevel string `json:"detail_level"` // "low" | "high"
Language string `json:"language"`
Scenario string `json:"scenario"` // 情景 ID如 "free_chat"、"interviewer"
}
// DefaultSessionTitle 默认会话标题。
const DefaultSessionTitle = "新对话"
// DefaultConfig 默认会话配置。
func DefaultConfig() SessionConfig {
return SessionConfig{TTSEnabled: true, DetailLevel: "low", Language: "zh-CN"}
return SessionConfig{TTSEnabled: true, DetailLevel: "low", Language: "zh-CN", Scenario: "free_chat"}
}
// SessionConfigPatch 会话配置增量更新(指针字段表示"未传则不更新")。
@@ -26,6 +33,7 @@ type SessionConfigPatch struct {
TTSEnabled *bool `json:"tts_enabled,omitempty"`
DetailLevel *string `json:"detail_level,omitempty"`
Language *string `json:"language,omitempty"`
Scenario *string `json:"scenario,omitempty"`
}
// Apply 将 patch 中的非 nil 字段覆盖到 cfg。
@@ -39,6 +47,18 @@ func (p SessionConfigPatch) Apply(cfg *SessionConfig) {
if p.Language != nil {
cfg.Language = *p.Language
}
if p.Scenario != nil {
cfg.Scenario = *p.Scenario
}
}
// User 用户。
type User struct {
ID string `json:"id"`
Username string `json:"username"`
PasswordHash string `json:"-"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// Message 对话消息。
@@ -55,6 +75,7 @@ type WsQuery struct {
RequestID string `json:"request_id"`
Image string `json:"image"` // base64
Audio string `json:"audio"` // base64
Text string `json:"text"` // 用户手动输入的文本(有值时跳过 STT
MimeType string `json:"mime_type"` // 默认 "audio/pcm"
}
@@ -65,6 +86,7 @@ type WsConfig struct {
TTSEnabled *bool `json:"tts_enabled,omitempty"`
DetailLevel *string `json:"detail_level,omitempty"`
Language *string `json:"language,omitempty"`
Scenario *string `json:"scenario,omitempty"`
} `json:"payload"`
}
@@ -111,7 +133,8 @@ type WsTTSAudio struct {
RequestID string `json:"request_id"`
Audio string `json:"audio"` // base64
MimeType string `json:"mime_type"` // "audio/mp3" 或 "audio/pcm"
IsLast bool `json:"is_last"`
IsLast bool `json:"is_last"` // 当前句子的音频是否完整(每句结束时为 true
Final bool `json:"final"` // 整轮 TTS 是否结束(所有句子合成完毕后为 true
}
// WsError 服务端 error 消息。

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,328 +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/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 上报
}
// New 创建 Pipeline 实例。
func New(
sttService stt.Service,
llmService llm.Service,
ttsService tts.Service,
sessionMgr session.Manager,
model string,
) *Pipeline {
return &Pipeline{
sttService: sttService,
llmService: llmService,
ttsService: ttsService,
sessionMgr: sessionMgr,
model: model,
}
}
// 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()
// 解码音频数据
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 != "" {
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: STT 语音识别
log.Infow("开始语音识别", "request_id", req.RequestID)
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)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "STT_ERROR",
Message: "语音识别失败",
})
return err
}
// 发送 STT 结果
if err := sender.SendSTTResult(models.WsSTTResult{
Type: "stt_result",
RequestID: req.RequestID,
Text: sttResult,
IsFinal: true,
}); err != nil {
log.Errorw("发送 STT 结果失败", "error", err)
}
// 追加用户消息到历史
p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "user",
Content: sttResult,
})
// Step 2+3: LLM 流式推理 + TTS 并行合成
log.Infow("开始 LLM 推理", "request_id", req.RequestID)
llmReq := llm.Request{
Image: image,
Text: sttResult,
History: history,
Language: sess.Config.Language,
}
llmStream, err := p.llmService.ChatStream(ctx, llmReq)
if err != nil {
log.Errorw("LLM 流式推理启动失败", "error", err)
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "LLM_ERROR",
Message: "LLM 推理失败",
})
return err
}
// 创建句子切分器
sentenceCh := make(chan string, 4)
splitter := NewSplitter(sentenceCh)
// 并行LLM 消费 + TTS 合成
var wg sync.WaitGroup
var fullText string
var ttsErr error
// goroutine 1: 消费 LLM token + 句子切分
wg.Add(1)
go func() {
defer wg.Done()
defer close(sentenceCh)
fullText = 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()
if err := sender.SendLLMDone(models.WsLLMDone{
Type: "llm_done",
RequestID: req.RequestID,
FullText: fullText,
Model: p.model,
LatencyMs: latency,
}); 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 并进行句子切分。
func (p *Pipeline) consumeLLMStream(
ctx context.Context,
stream <-chan llm.Chunk,
requestID string,
sender Sender,
splitter *Splitter,
) string {
log := logger.Log
var fullText strings.Builder
for chunk := range stream {
// 检查上下文是否已取消
select {
case <-ctx.Done():
log.Infow("LLM 流被中断", "request_id", requestID)
return fullText.String()
default:
}
if chunk.Done {
// 流结束
if chunk.TokensUsed != nil {
log.Infow("LLM 用量统计",
"request_id", requestID,
"prompt_tokens", chunk.TokensUsed.Prompt,
"completion_tokens", chunk.TokensUsed.Completion,
"total_tokens", chunk.TokensUsed.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()
}
// 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: "alloy",
Speed: 1.0,
OutputFmt: "mp3",
SampleRate: 24000,
})
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,
}); err != nil {
log.Errorw("发送 tts_audio 失败", "error", err)
}
}
return nil
}

View File

@@ -1,661 +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/logger"
"github.com/hhs/camtalk/internal/models"
)
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, config models.SessionConfig) (string, error) {
args := m.Called(ctx, config)
return args.String(0), args.Error(1)
}
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, "gpt-4o")
// 执行
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, "gpt-4o")
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, "gpt-4o")
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, "gpt-4o")
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, "gpt-4o")
// 创建可取消的上下文
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, "gpt-4o")
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, "gpt-4o")
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, "gpt-4o")
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

@@ -4,6 +4,7 @@ package session
import (
"context"
"errors"
"time"
"github.com/hhs/camtalk/internal/models"
)
@@ -11,11 +12,21 @@ import (
// ErrSessionNotFound 会话不存在或已过期。
var ErrSessionNotFound = errors.New("session not found")
// ConversationSummary 对话摘要(列表展示用)。
type ConversationSummary struct {
ID string `json:"id"`
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"`
}
// Manager 会话管理器接口。
// WebSocket Handler 通过此接口操作会话,不直接接触存储层。
type Manager interface {
// Create 创建新会话,返回 session ID。
Create(ctx context.Context, config models.SessionConfig) (string, error)
// Create 创建新会话,返回 session ID。userID 为空表示匿名会话。
Create(ctx context.Context, userID string, config models.SessionConfig) (string, error)
// Get 获取会话(含 config。不存在返回 ErrSessionNotFound。
Get(ctx context.Context, sessionID string) (*models.Session, error)
@@ -23,6 +34,12 @@ type Manager interface {
// UpdateConfig 更新会话配置config 消息触发)。
UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error
// UpdateTitle 更新会话标题。
UpdateTitle(ctx context.Context, sessionID string, title string) error
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error)
// GetHistory 获取最近 N 轮对话历史(供 Orchestrator 构建 LLM 上下文)。
GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error)

View File

@@ -2,6 +2,8 @@ package session
import (
"context"
"encoding/json"
"sort"
"sync"
"time"
@@ -9,6 +11,7 @@ import (
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/store"
)
const (
@@ -32,11 +35,31 @@ type MemoryManager struct {
ttl time.Duration
maxHistory int
stopCleaner chan struct{}
msgRepo store.MessageRepository // 可选消息持久化Write-Through
sessRepo store.SessionRepository // 可选会话持久化Write-Through
}
// Option MemoryManager 的函数式选项。
type Option func(*MemoryManager)
// WithMessageRepository 注入消息持久化仓库,启用 Write-Through 模式。
func WithMessageRepository(repo store.MessageRepository) Option {
return func(m *MemoryManager) {
m.msgRepo = repo
}
}
// WithSessionRepository 注入会话持久化仓库,启用会话元数据 Write-Through 模式。
func WithSessionRepository(repo store.SessionRepository) Option {
return func(m *MemoryManager) {
m.sessRepo = repo
}
}
// NewMemoryManager 创建内存版 SessionManager。
// ttl 为会话过期时间maxHistory 为对话历史上限0 表示使用默认值 20
func NewMemoryManager(ttl time.Duration, maxHistory int) *MemoryManager {
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。
func NewMemoryManager(ttl time.Duration, maxHistory int, opts ...Option) *MemoryManager {
if ttl <= 0 {
ttl = defaultTTL
}
@@ -51,6 +74,10 @@ func NewMemoryManager(ttl time.Duration, maxHistory int) *MemoryManager {
stopCleaner: make(chan struct{}),
}
for _, opt := range opts {
opt(m)
}
// 启动后台清理 goroutine每分钟清除过期会话。
go m.cleanLoop()
@@ -95,58 +122,246 @@ func (m *MemoryManager) isExpired(entry *sessionEntry) bool {
return time.Since(entry.lastActive) > m.ttl
}
// Create 创建新会话。
func (m *MemoryManager) Create(_ context.Context, config models.SessionConfig) (string, error) {
// Create 创建新会话。userID 为空表示匿名会话。
func (m *MemoryManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
m.mu.Lock()
defer m.mu.Unlock()
id := uuid.New().String()
now := time.Now()
m.sessions[id] = &sessionEntry{
session: models.Session{
ID: id,
UserID: userID,
Title: models.DefaultSessionTitle,
CreatedAt: now,
UpdatedAt: now,
Config: config,
},
history: make([]models.Message, 0),
lastActive: now,
}
m.mu.Unlock()
logger.Log.Debugw("session created", "session", id)
// Write-Through异步写 PG使用 Background context避免 HTTP 请求结束后 context 被取消)
if m.sessRepo != nil {
go func() {
cfgJSON, _ := json.Marshal(config)
if err := m.sessRepo.Save(context.Background(), store.SessionRecord{
ID: id, UserID: userID, Title: models.DefaultSessionTitle,
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
}); err != nil {
logger.Log.Warnw("persist session failed", "session", id, "error", err)
}
}()
}
logger.Log.Debugw("session created", "session", id, "user_id", userID)
return id, nil
}
// Get 获取会话。
func (m *MemoryManager) Get(_ context.Context, sessionID string) (*models.Session, error) {
// Get 获取会话。内存中不存在时,尝试从 PG 加载(透明恢复)。
func (m *MemoryManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
m.mu.RLock()
defer m.mu.RUnlock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
return nil, ErrSessionNotFound
if ok && !m.isExpired(entry) {
sess := entry.session
m.mu.RUnlock()
return &sess, nil
}
m.mu.RUnlock()
// 内存未命中,尝试从 PG 加载
if m.sessRepo != nil {
rec, err := m.sessRepo.FindByID(ctx, sessionID)
if err != nil {
return nil, ErrSessionNotFound
}
sess := m.recordToSession(rec)
// 加载到内存(含消息历史)
if m.msgRepo != nil {
_ = m.LoadSessionFromRepo(ctx, sess)
} else {
_ = m.LoadSession(sess, nil)
}
return sess, nil
}
sess := entry.session // 复制一份返回
return &sess, nil
return nil, ErrSessionNotFound
}
// UpdateConfig 更新会话配置。
func (m *MemoryManager) UpdateConfig(_ context.Context, sessionID string, patch models.SessionConfigPatch) error {
func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
m.mu.Lock()
defer m.mu.Unlock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound
}
patch.Apply(&entry.session.Config)
entry.lastActive = time.Now()
cfg := entry.session.Config
m.mu.Unlock()
// Write-Through异步更新 PG使用 Background context
if m.sessRepo != nil {
go func() {
cfgJSON, _ := json.Marshal(cfg)
if err := m.sessRepo.UpdateConfig(context.Background(), sessionID, cfgJSON); err != nil {
logger.Log.Warnw("update session config in DB failed", "session", sessionID, "error", err)
}
}()
}
logger.Log.Debugw("session config updated", "session", sessionID)
return nil
}
// UpdateTitle 更新会话标题。
func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
m.mu.Lock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound
}
entry.session.Title = title
entry.session.UpdatedAt = time.Now()
entry.lastActive = time.Now()
m.mu.Unlock()
// Write-Through异步更新 PG使用 Background context
if m.sessRepo != nil {
go func() {
if err := m.sessRepo.UpdateTitle(context.Background(), sessionID, title); err != nil {
logger.Log.Warnw("update session title in DB failed", "session", sessionID, "error", err)
}
}()
}
logger.Log.Debugw("session title updated", "session", sessionID, "title", title)
return nil
}
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
// 若配置了 SessionRepository从 PG 查询(包含内存中已过期的会话)。
// 若配置了 MessageRepository消息统计从 PostgreSQL 聚合查询(更准确)。
func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
if page <= 0 {
page = 1
}
if size <= 0 {
size = 20
}
// 优先从 PG 查询会话列表(包含已过期的会话)
if m.sessRepo != nil {
recs, total, err := m.sessRepo.FindByUser(ctx, userID, page, size)
if err != nil {
logger.Log.Warnw("list sessions from DB failed, falling back to in-memory", "error", err)
return m.listByUserFromMemory(ctx, userID, page, size)
}
list := make([]ConversationSummary, 0, len(recs))
var sessionIDs []string
for _, rec := range recs {
list = append(list, ConversationSummary{
ID: rec.ID,
Title: rec.Title,
CreatedAt: rec.CreatedAt,
UpdatedAt: rec.UpdatedAt,
})
sessionIDs = append(sessionIDs, rec.ID)
}
// 用内存中的消息数填充
m.mu.RLock()
for i := range list {
if entry, ok := m.sessions[list[i].ID]; ok {
list[i].MessageCount = len(entry.history)
if len(entry.history) > 0 {
list[i].LastMessage = entry.history[len(entry.history)-1].Content
}
}
}
m.mu.RUnlock()
// 从 PG 获取更准确的消息统计
if m.msgRepo != nil && len(sessionIDs) > 0 {
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
for i := range list {
if s, ok := stats[list[i].ID]; ok {
list[i].LastMessage = s.LastMessage
list[i].MessageCount = s.MessageCount
}
}
}
}
return list, total, nil
}
// fallback纯内存查询
return m.listByUserFromMemory(ctx, userID, page, size)
}
// listByUserFromMemory 从内存中获取用户的对话列表(无 PG 时的 fallback
func (m *MemoryManager) listByUserFromMemory(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
m.mu.RLock()
var list []ConversationSummary
var sessionIDs []string
for _, entry := range m.sessions {
if entry.session.UserID != userID || m.isExpired(entry) {
continue
}
summary := ConversationSummary{
ID: entry.session.ID,
Title: entry.session.Title,
CreatedAt: entry.session.CreatedAt,
UpdatedAt: entry.lastActive,
}
summary.MessageCount = len(entry.history)
if len(entry.history) > 0 {
summary.LastMessage = entry.history[len(entry.history)-1].Content
}
list = append(list, summary)
sessionIDs = append(sessionIDs, entry.session.ID)
}
m.mu.RUnlock()
// 从 PG 获取更准确的消息统计
if m.msgRepo != nil && len(sessionIDs) > 0 {
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
for i := range list {
if s, ok := stats[list[i].ID]; ok {
list[i].LastMessage = s.LastMessage
list[i].MessageCount = s.MessageCount
}
}
}
}
sort.Slice(list, func(i, j int) bool {
return list[i].UpdatedAt.After(list[j].UpdatedAt)
})
total := len(list)
start := (page - 1) * size
if start >= total {
return []ConversationSummary{}, total, nil
}
end := start + size
if end > total {
end = total
}
return list[start:end], total, nil
}
// GetHistory 获取最近 N 轮对话历史。
func (m *MemoryManager) GetHistory(_ context.Context, sessionID string, limit int) ([]models.Message, error) {
m.mu.RLock()
@@ -168,26 +383,116 @@ func (m *MemoryManager) GetHistory(_ context.Context, sessionID string, limit in
}
// AppendMessage 追加一条对话消息,同时刷新 TTL。
// 若配置了 MessageRepository消息会异步写入 PostgreSQLWrite-Through
func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg models.Message) error {
m.mu.Lock()
defer m.mu.Unlock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound
}
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 条
if len(entry.history) > m.maxHistory {
entry.history = entry.history[len(entry.history)-m.maxHistory:]
}
entry.lastActive = time.Now()
now := time.Now()
entry.lastActive = now
entry.session.UpdatedAt = now
// 复制标题(释放锁后安全使用)
persistTitle := entry.session.Title
m.mu.Unlock()
// 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 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)
}
}
}()
}
return nil
}
// generateTitle 从首条消息生成对话标题(取前 20 个字符)。
func generateTitle(firstMessage string) string {
runes := []rune(firstMessage)
if len(runes) > 20 {
return string(runes[:20]) + "…"
}
return firstMessage
}
// LoadSession 从外部存储加载会话到内存热存储。
// 用于 conversation_id 恢复场景WS 连接时会话不在内存中,从 PostgreSQL 加载。
// 若会话已在内存中,返回 nil幂等
func (m *MemoryManager) LoadSession(sess *models.Session, messages []models.Message) error {
m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.sessions[sess.ID]; ok {
return nil // 已在内存中,无需重复加载
}
m.sessions[sess.ID] = &sessionEntry{
session: *sess,
history: messages,
lastActive: time.Now(),
}
logger.Log.Debugw("session loaded from DB", "session", sess.ID, "messages", len(messages))
return nil
}
// LoadSessionFromRepo 从 MessageRepository 加载会话消息并注册到内存。
// 适用于已注入 MessageRepository 的场景,调用方只需传入 session 元数据。
func (m *MemoryManager) LoadSessionFromRepo(ctx context.Context, sess *models.Session) error {
if m.msgRepo == nil {
return m.LoadSession(sess, nil)
}
// 从冷存储加载全部消息limit=0 表示全量)
stored, err := m.msgRepo.GetMessages(ctx, sess.ID, 0, 0)
if err != nil {
return err
}
messages := make([]models.Message, len(stored))
for i, s := range stored {
messages[i] = models.Message{Role: s.Role, Content: s.Content}
}
return m.LoadSession(sess, messages)
}
// SetActiveRequest 标记当前正在处理的请求 ID。
func (m *MemoryManager) SetActiveRequest(_ context.Context, sessionID string, requestID string) error {
m.mu.Lock()
@@ -246,15 +551,26 @@ func (m *MemoryManager) Touch(_ context.Context, sessionID string) error {
}
// Destroy 显式销毁会话。
func (m *MemoryManager) Destroy(_ context.Context, sessionID string) error {
func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.sessions[sessionID]; !ok {
m.mu.Unlock()
return ErrSessionNotFound
}
delete(m.sessions, sessionID)
m.mu.Unlock()
// Write-Through异步删除 PG使用 Background context
if m.sessRepo != nil {
go func() {
if err := m.sessRepo.Delete(context.Background(), sessionID); err != nil {
logger.Log.Warnw("delete session from DB failed", "session", sessionID, "error", err)
}
}()
}
logger.Log.Debugw("session destroyed", "session", sessionID)
return nil
}
@@ -273,3 +589,19 @@ func (m *MemoryManager) ActiveCount() int {
}
return count
}
// recordToSession 将 store.SessionRecord 转换为 models.Session。
func (m *MemoryManager) recordToSession(rec *store.SessionRecord) *models.Session {
cfg := models.DefaultConfig()
if len(rec.Config) > 0 {
_ = json.Unmarshal(rec.Config, &cfg)
}
return &models.Session{
ID: rec.ID,
UserID: rec.UserID,
Title: rec.Title,
CreatedAt: rec.CreatedAt,
UpdatedAt: rec.UpdatedAt,
Config: cfg,
}
}

View File

@@ -19,7 +19,7 @@ func TestCreateAndGet(t *testing.T) {
ctx := context.Background()
config := models.DefaultConfig()
id, err := m.Create(ctx, config)
id, err := m.Create(ctx, "", config)
if err != nil {
t.Fatalf("Create: %v", err)
}
@@ -56,7 +56,7 @@ func TestExpire(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
// 未过期时应能获取
_, err := m.Get(ctx, id)
@@ -78,7 +78,7 @@ func TestDestroy(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
if err := m.Destroy(ctx, id); err != nil {
t.Fatalf("Destroy: %v", err)
@@ -106,7 +106,7 @@ func TestAppendMessageAndGetHistory(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
msgs := []models.Message{
{Role: "user", Content: "你好"},
@@ -138,7 +138,7 @@ func TestGetHistoryLimit(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
for i := 0; i < 10; i++ {
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
@@ -159,7 +159,7 @@ func TestHistoryLimit(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
// 插入超过上限的消息
for i := 0; i < 10; i++ {
@@ -180,7 +180,7 @@ func TestUpdateConfig(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
ttsEnabled := false
detailLevel := "high"
@@ -211,7 +211,7 @@ func TestActiveRequest(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
// 初始应为空
reqID, err := m.GetActiveRequestID(ctx, id)
@@ -246,7 +246,7 @@ func TestTouchRefreshesTTL(t *testing.T) {
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, models.DefaultConfig())
id, _ := m.Create(ctx, "", models.DefaultConfig())
// 50ms 后 Touch应重置 TTL
time.Sleep(50 * time.Millisecond)
@@ -278,9 +278,204 @@ func TestActiveCount(t *testing.T) {
t.Errorf("initial ActiveCount = %d, want 0", m.ActiveCount())
}
m.Create(ctx, models.DefaultConfig())
m.Create(ctx, models.DefaultConfig())
m.Create(ctx, "", models.DefaultConfig())
m.Create(ctx, "", models.DefaultConfig())
if m.ActiveCount() != 2 {
t.Errorf("ActiveCount = %d, want 2", m.ActiveCount())
}
}
func TestCreateWithUserID(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
id, err := m.Create(ctx, "user-123", models.DefaultConfig())
if err != nil {
t.Fatalf("Create: %v", err)
}
sess, err := m.Get(ctx, id)
if err != nil {
t.Fatalf("Get: %v", err)
}
if sess.UserID != "user-123" {
t.Errorf("UserID = %q, want %q", sess.UserID, "user-123")
}
if sess.Title != models.DefaultSessionTitle {
t.Errorf("Title = %q, want %q", sess.Title, models.DefaultSessionTitle)
}
if sess.UpdatedAt.IsZero() {
t.Error("UpdatedAt should not be zero")
}
}
func TestUpdateTitle(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
if err := m.UpdateTitle(ctx, id, "自定义标题"); err != nil {
t.Fatalf("UpdateTitle: %v", err)
}
sess, _ := m.Get(ctx, id)
if sess.Title != "自定义标题" {
t.Errorf("Title = %q, want %q", sess.Title, "自定义标题")
}
}
func TestUpdateTitleNotFound(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
err := m.UpdateTitle(ctx, "nonexistent", "标题")
if err != ErrSessionNotFound {
t.Errorf("UpdateTitle nonexistent: err = %v, want ErrSessionNotFound", err)
}
}
func TestAutoTitleOnFirstMessage(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
// 首条 user 消息应自动更新标题
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "你好世界"})
sess, _ := m.Get(ctx, id)
if sess.Title != "你好世界" {
t.Errorf("Title = %q, want %q", sess.Title, "你好世界")
}
}
func TestAutoTitleLongMessage(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
// 超过 20 字符的消息应截断
longMsg := "这是一条很长很长很长很长很长很长很长很长的消息"
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: longMsg})
sess, _ := m.Get(ctx, id)
expected := string([]rune(longMsg)[:20]) + "…"
if sess.Title != expected {
t.Errorf("Title = %q, want %q", sess.Title, expected)
}
}
func TestAutoTitleNotOverwritten(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
// 首条消息设置标题
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "第一条消息"})
// 第二条消息不应覆盖已有的标题
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "第二条消息"})
sess, _ := m.Get(ctx, id)
if sess.Title != "第一条消息" {
t.Errorf("Title = %q, want %q", sess.Title, "第一条消息")
}
}
func TestListByUser(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
// 创建两个用户的不同会话
id1, _ := m.Create(ctx, "user-1", models.DefaultConfig())
m.AppendMessage(ctx, id1, models.Message{Role: "user", Content: "会话1"})
id2, _ := m.Create(ctx, "user-1", models.DefaultConfig())
m.AppendMessage(ctx, id2, models.Message{Role: "user", Content: "会话2"})
m.Create(ctx, "user-2", models.DefaultConfig()) // 其他用户的会话
list, total, err := m.ListByUser(ctx, "user-1", 1, 10)
if err != nil {
t.Fatalf("ListByUser: %v", err)
}
if total != 2 {
t.Errorf("total = %d, want 2", total)
}
if len(list) != 2 {
t.Fatalf("len = %d, want 2", len(list))
}
// 按 UpdatedAt 降序id2 应在前
if list[0].ID != id2 {
t.Errorf("list[0].ID = %q, want %q", list[0].ID, id2)
}
if list[0].Title != "会话2" {
t.Errorf("list[0].Title = %q, want %q", list[0].Title, "会话2")
}
if list[0].LastMessage != "会话2" {
t.Errorf("list[0].LastMessage = %q, want %q", list[0].LastMessage, "会话2")
}
}
func TestListByUserPagination(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
// 创建 5 个会话
for i := 0; i < 5; i++ {
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
}
// 第 1 页,每页 2 条
list, total, _ := m.ListByUser(ctx, "user-1", 1, 2)
if total != 5 {
t.Errorf("total = %d, want 5", total)
}
if len(list) != 2 {
t.Errorf("page 1 len = %d, want 2", len(list))
}
// 第 2 页
list, _, _ = m.ListByUser(ctx, "user-1", 2, 2)
if len(list) != 2 {
t.Errorf("page 2 len = %d, want 2", len(list))
}
// 第 3 页(最后一页)
list, _, _ = m.ListByUser(ctx, "user-1", 3, 2)
if len(list) != 1 {
t.Errorf("page 3 len = %d, want 1", len(list))
}
// 超出范围的页
list, _, _ = m.ListByUser(ctx, "user-1", 10, 2)
if len(list) != 0 {
t.Errorf("out of range page len = %d, want 0", len(list))
}
}
func TestListByUserEmpty(t *testing.T) {
m := NewMemoryManager(30*time.Minute, 20)
defer m.Stop()
ctx := context.Background()
list, total, err := m.ListByUser(ctx, "no-such-user", 1, 10)
if err != nil {
t.Fatalf("ListByUser: %v", err)
}
if total != 0 {
t.Errorf("total = %d, want 0", total)
}
if len(list) != 0 {
t.Errorf("len = %d, want 0", len(list))
}
}

View File

@@ -10,14 +10,16 @@ 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 实现。
// 数据结构:
// - session:{id}:meta → Hash会话元数据
// - session:{id}:history → List对话历史
// - user:{id}:sessions → Set用户会话索引
type RedisManager struct {
rdb *redis.Client
ttl time.Duration
@@ -35,37 +37,59 @@ func NewRedisManager(rdb *redis.Client, ttl time.Duration, maxHistory int) *Redi
return &RedisManager{rdb: rdb, ttl: ttl, maxHistory: maxHistory}
}
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
// Ping 检查 Redis 连接是否正常。
func (m *RedisManager) Ping(ctx context.Context) error {
return m.rdb.Ping(ctx).Err()
}
// Create 创建新会话。
func (m *RedisManager) Create(ctx context.Context, config models.SessionConfig) (string, error) {
id := uuidNew()
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) {
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()
// 写入 meta Hash
pipe.HSet(ctx, metaKey(id), map[string]interface{}{
"session_id": id,
"config.tts_enabled": strconv.FormatBool(config.TTSEnabled),
meta := map[string]interface{}{
"session_id": id,
"user_id": userID,
"title": models.DefaultSessionTitle,
"config.tts_enabled": strconv.FormatBool(config.TTSEnabled),
"config.detail_level": config.DetailLevel,
"config.language": config.Language,
"created_at": now.Format(time.RFC3339),
"last_active": now.Format(time.RFC3339),
"active_request_id": "",
})
"config.language": config.Language,
"created_at": now.Format(time.RFC3339),
"updated_at": now.Format(time.RFC3339),
"last_active": now.Format(time.RFC3339),
"active_request_id": "",
}
pipe.HSet(ctx, metaKey(id), meta)
pipe.Expire(ctx, metaKey(id), m.ttl)
// 初始化空 history List
pipe.RPush(ctx, histKey(id), placeholderHistoryMark)
pipe.Expire(ctx, histKey(id), m.ttl)
// 如果有 userID添加到用户会话索引
if userID != "" {
pipe.SAdd(ctx, userSessKey(userID), id)
pipe.Expire(ctx, userSessKey(userID), m.ttl)
}
if _, err := pipe.Exec(ctx); err != nil {
return "", fmt.Errorf("redis create session: %w", err)
}
logger.Log.Debugw("redis session created", "session", id)
log := trace.FromContext(ctx)
log.Debugw("redis session created", "session_id", id, "user_id", userID)
return id, nil
}
@@ -74,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 {
@@ -83,13 +110,17 @@ func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Sessi
}
sess := &models.Session{
ID: vals["session_id"],
ID: vals["session_id"],
UserID: vals["user_id"],
Title: vals["title"],
}
sess.CreatedAt, _ = time.Parse(time.RFC3339, vals["created_at"])
sess.UpdatedAt, _ = time.Parse(time.RFC3339, vals["updated_at"])
sess.Config.TTSEnabled, _ = strconv.ParseBool(vals["config.tts_enabled"])
sess.Config.DetailLevel = vals["config.detail_level"]
sess.Config.Language = vals["config.language"]
log.Debugw("redis session retrieved", "session_id", sessionID)
return sess, nil
}
@@ -104,8 +135,10 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
return ErrSessionNotFound
}
now := time.Now().UTC().Format(time.RFC3339)
fields := map[string]interface{}{
"last_active": time.Now().UTC().Format(time.RFC3339),
"last_active": now,
"updated_at": now,
}
if patch.TTSEnabled != nil {
fields["config.tts_enabled"] = strconv.FormatBool(*patch.TTSEnabled)
@@ -123,10 +156,117 @@ 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
}
// UpdateTitle 更新会话标题。
func (m *RedisManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
exists, err := m.rdb.Exists(ctx, metaKey(sessionID)).Result()
if err != nil {
return fmt.Errorf("redis check session: %w", err)
}
if exists == 0 {
return ErrSessionNotFound
}
now := time.Now().UTC().Format(time.RFC3339)
if err := m.rdb.HSet(ctx, metaKey(sessionID), "title", title, "updated_at", now, "last_active", now).Err(); err != nil {
return fmt.Errorf("redis update title: %w", err)
}
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
log := trace.FromContext(ctx)
log.Debugw("redis session title updated", "session_id", sessionID, "title", title)
return nil
}
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
func (m *RedisManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
if page <= 0 {
page = 1
}
if size <= 0 {
size = 20
}
// 从用户会话索引获取所有 session ID
sessionIDs, err := m.rdb.SMembers(ctx, userSessKey(userID)).Result()
if err != nil {
return nil, 0, fmt.Errorf("redis list user sessions: %w", err)
}
// 收集有效的会话摘要
var list []ConversationSummary
for _, sid := range sessionIDs {
vals, err := m.rdb.HGetAll(ctx, metaKey(sid)).Result()
if err != nil || len(vals) == 0 {
continue
}
updatedAt, _ := time.Parse(time.RFC3339, vals["updated_at"])
lastActive, _ := time.Parse(time.RFC3339, vals["last_active"])
// 检查是否过期
if time.Since(lastActive) > m.ttl {
continue
}
// 获取最后一条消息
lastMsg := ""
msgCount := 0
raws, err := m.rdb.LRange(ctx, histKey(sid), 0, 0).Result()
if err == nil && len(raws) > 0 && raws[0] != placeholderHistoryMark {
var msg models.Message
if json.Unmarshal([]byte(raws[0]), &msg) == nil {
lastMsg = msg.Content
}
}
// 获取消息总数(减去占位符)
totalLen, err := m.rdb.LLen(ctx, histKey(sid)).Result()
if err == nil {
msgCount = int(totalLen)
if msgCount > 0 {
msgCount-- // 减去占位符
}
}
list = append(list, ConversationSummary{
ID: vals["session_id"],
Title: vals["title"],
LastMessage: lastMsg,
MessageCount: msgCount,
UpdatedAt: updatedAt,
})
}
// 按 UpdatedAt 降序排序
for i := 0; i < len(list); i++ {
for j := i + 1; j < len(list); j++ {
if list[j].UpdatedAt.After(list[i].UpdatedAt) {
list[i], list[j] = list[j], list[i]
}
}
}
total := len(list)
// 分页
start := (page - 1) * size
if start >= total {
return []ConversationSummary{}, total, nil
}
end := start + size
if end > total {
end = total
}
return list[start:end], total, nil
}
// GetHistory 获取最近 N 轮对话历史。
func (m *RedisManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
// 检查会话是否存在
@@ -155,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)
@@ -193,8 +337,18 @@ func (m *RedisManager) AppendMessage(ctx context.Context, sessionID string, msg
// 刷新 TTL
pipe.Expire(ctx, histKey(sessionID), m.ttl)
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
// 更新 last_active
pipe.HSet(ctx, metaKey(sessionID), "last_active", time.Now().UTC().Format(time.RFC3339))
now := time.Now().UTC().Format(time.RFC3339)
// 更新 last_active 和 updated_at
pipe.HSet(ctx, metaKey(sessionID), "last_active", now, "updated_at", now)
// 自动更新标题:首条 user 消息时,如果标题为默认值
if msg.Role == "user" {
title, _ := m.rdb.HGet(ctx, metaKey(sessionID), "title").Result()
if title == models.DefaultSessionTitle {
pipe.HSet(ctx, metaKey(sessionID), "title", generateTitle(msg.Content))
}
}
if _, err := pipe.Exec(ctx); err != nil {
return fmt.Errorf("redis append message: %w", err)
@@ -280,6 +434,9 @@ func (m *RedisManager) Touch(ctx context.Context, sessionID string) error {
// Destroy 显式销毁会话。
func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
// 先获取 user_id 以便清理索引
userID, _ := m.rdb.HGet(ctx, metaKey(sessionID), "user_id").Result()
deleted, err := m.rdb.Del(ctx, metaKey(sessionID), histKey(sessionID)).Result()
if err != nil {
return fmt.Errorf("redis destroy session: %w", err)
@@ -288,7 +445,13 @@ func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
return ErrSessionNotFound
}
logger.Log.Debugw("redis session destroyed", "session", sessionID)
// 清理用户会话索引
if userID != "" {
m.rdb.SRem(ctx, userSessKey(userID), 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

@@ -0,0 +1,17 @@
package store
import (
"context"
"github.com/jackc/pgx/v5/pgxpool"
)
// NewPostgresPool 创建 PostgreSQL 连接池。
func NewPostgresPool(ctx context.Context, dsn string) (*pgxpool.Pool, error) {
cfg, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, err
}
cfg.MaxConns = 10
return pgxpool.NewWithConfig(ctx, cfg)
}

View File

@@ -0,0 +1,50 @@
package store
import (
"context"
"errors"
"time"
"github.com/hhs/camtalk/internal/models"
)
var (
// ErrMessageNotFound 消息不存在。
ErrMessageNotFound = errors.New("message not found")
)
// MessageRepository 消息持久化接口。
type MessageRepository interface {
// SaveMessage 保存一条消息。
SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error
// GetMessages 获取会话的消息列表(分页,按 created_at 升序)。
// beforeID 为 0 时从最新开始查询。
GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error)
// GetLastMessage 获取会话的最后一条消息。
GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error)
// GetMessageCount 获取会话的消息总数。
GetMessageCount(ctx context.Context, sessionID string) (int, error)
// GetSessionMessageStats 批量查询多个会话的消息统计last_message + message_count
// 返回的 map key 为 sessionID仅包含有消息的会话。
GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error)
}
// SessionMessageStats 单个会话的消息统计SQL 聚合查询结果)。
type SessionMessageStats struct {
LastMessage string
MessageCount int
}
// StoredMessage 持久化消息模型store 层)。
type StoredMessage struct {
ID int64 `json:"id"`
SessionID string `json:"-"`
Role string `json:"role"`
Content string `json:"content"`
TokensUsed int `json:"tokens_used"`
CreatedAt time.Time `json:"created_at"`
}

View File

@@ -0,0 +1,198 @@
package store
import (
"context"
"errors"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/trace"
)
// PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。
type PgMessageRepository struct {
pool *pgxpool.Pool
}
// NewPgMessageRepository 创建 PgMessageRepository。
func NewPgMessageRepository(pool *pgxpool.Pool) *PgMessageRepository {
return &PgMessageRepository{pool: pool}
}
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,
)
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
}
var rows []StoredMessage
var err error
if beforeID > 0 {
rows, err = r.queryMessages(ctx,
`SELECT id, session_id, role, content, tokens_used, created_at
FROM messages
WHERE session_id = $1 AND id < $2
ORDER BY id DESC
LIMIT $3`,
sessionID, beforeID, limit,
)
} else {
rows, err = r.queryMessages(ctx,
`SELECT id, session_id, role, content, tokens_used, created_at
FROM messages
WHERE session_id = $1
ORDER BY id DESC
LIMIT $2`,
sessionID, limit,
)
}
if err != nil {
log.Errorw("get messages failed", "session_id", sessionID, "error", err)
return nil, err
}
// 反转为升序
for i, j := 0, len(rows)-1; i < j; i, j = i+1, j-1 {
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()
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
FROM messages
WHERE session_id = $1
ORDER BY id DESC
LIMIT 1`,
sessionID,
).Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt)
if errors.Is(err, pgx.ErrNoRows) {
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
}
rows, err := r.pool.Query(ctx,
`WITH stats AS (
SELECT session_id, COUNT(*) AS cnt
FROM messages
WHERE session_id = ANY($1)
GROUP BY session_id
),
last_msg AS (
SELECT DISTINCT ON (session_id) session_id, content
FROM messages
WHERE session_id = ANY($1)
ORDER BY session_id, id DESC
)
SELECT s.session_id, s.cnt, COALESCE(lm.content, '')
FROM stats s
LEFT JOIN last_msg lm ON lm.session_id = s.session_id`,
sessionIDs,
)
if err != nil {
log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err)
return nil, err
}
defer rows.Close()
result := make(map[string]SessionMessageStats)
for rows.Next() {
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

@@ -0,0 +1,80 @@
package store
import (
"context"
"fmt"
"io/fs"
"sort"
"strings"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/logger"
)
// RunMigrations 从给定的 fs.FS 中读取 *.up.sql 文件并按版本号顺序执行。
// 已执行过的版本会跳过(通过 schema_migrations 表记录)。
func RunMigrations(ctx context.Context, pool *pgxpool.Pool, fsys fs.FS) error {
// 确保 schema_migrations 表存在
if _, err := pool.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
)`); err != nil {
return fmt.Errorf("create schema_migrations table: %w", err)
}
// 收集所有 *.up.sql 文件
entries, err := fs.ReadDir(fsys, ".")
if err != nil {
return fmt.Errorf("read migrations dir: %w", err)
}
var files []string
for _, e := range entries {
if !e.IsDir() && strings.HasSuffix(e.Name(), ".up.sql") {
files = append(files, e.Name())
}
}
sort.Strings(files)
for _, name := range files {
// 从文件名提取版本号,如 "001_users.up.sql" → 1
var version int
if _, err := fmt.Sscanf(name, "%d_", &version); err != nil {
return fmt.Errorf("parse version from %s: %w", name, err)
}
// 检查是否已执行
var exists bool
if err := pool.QueryRow(ctx,
`SELECT EXISTS(SELECT 1 FROM schema_migrations WHERE version = $1)`, version,
).Scan(&exists); err != nil {
return fmt.Errorf("check migration version %d: %w", version, err)
}
if exists {
logger.Log.Debugw("migration already applied", "version", version, "file", name)
continue
}
// 读取并执行
content, err := fs.ReadFile(fsys, name)
if err != nil {
return fmt.Errorf("read migration %s: %w", name, err)
}
if _, err := pool.Exec(ctx, string(content)); err != nil {
return fmt.Errorf("execute migration %s: %w", name, err)
}
// 记录已执行
if _, err := pool.Exec(ctx,
`INSERT INTO schema_migrations (version) VALUES ($1)`, version,
); err != nil {
return fmt.Errorf("record migration %d: %w", version, err)
}
logger.Log.Infow("migration applied", "version", version, "file", name)
}
return nil
}

View File

@@ -0,0 +1,47 @@
package store
import (
"context"
"errors"
"time"
)
var (
// ErrSessionNotFound 会话不存在。
ErrSessionNotFound = errors.New("session not found")
)
// SessionRepository 会话持久化接口。
type SessionRepository interface {
// Save 创建或更新会话UPSERT
Save(ctx context.Context, s SessionRecord) error
// FindByID 根据 ID 查询会话。
FindByID(ctx context.Context, id string) (*SessionRecord, error)
// FindByUser 查询用户的会话列表(分页,按 updated_at 降序)。
// 返回 (列表, 总数, error)。
FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error)
// UpdateTitle 更新会话标题。
UpdateTitle(ctx context.Context, id string, title string) error
// UpdateConfig 更新会话配置。
UpdateConfig(ctx context.Context, id string, configJSON []byte) error
// Touch 刷新 updated_at。
Touch(ctx context.Context, id string) error
// Delete 删除会话。
Delete(ctx context.Context, id string) error
}
// SessionRecord 持久化会话模型store 层)。
type SessionRecord struct {
ID string
UserID string
Title string
Config []byte // JSON 编码的 SessionConfig
CreatedAt time.Time
UpdatedAt time.Time
}

View File

@@ -0,0 +1,189 @@
package store
import (
"context"
"errors"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/trace"
)
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
type PgSessionRepository struct {
pool *pgxpool.Pool
}
// NewPgSessionRepository 创建 PgSessionRepository。
func NewPgSessionRepository(pool *pgxpool.Pool) *PgSessionRepository {
return &PgSessionRepository{pool: pool}
}
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)
ON CONFLICT (id) DO UPDATE SET
title = EXCLUDED.title,
config = EXCLUDED.config,
updated_at = EXCLUDED.updated_at`,
s.ID, s.UserID, s.Title, s.Config, s.CreatedAt, s.UpdatedAt,
)
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
FROM sessions WHERE id = $1`, id,
).Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt)
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
}
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
}
if size <= 0 {
size = 20
}
offset := (page - 1) * size
// 查询总数
var total int
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
}
// 查询列表
rows, err := r.pool.Query(ctx,
`SELECT id, user_id, title, config, created_at, updated_at
FROM sessions
WHERE user_id = $1
ORDER BY updated_at DESC
LIMIT $2 OFFSET $3`,
userID, size, offset,
)
if err != nil {
log.Errorw("find user sessions failed", "user_id", userID, "error", err)
return nil, 0, err
}
defer rows.Close()
var list []SessionRecord
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

@@ -0,0 +1,46 @@
package store
import (
"context"
"errors"
"time"
)
var (
ErrUserNotFound = errors.New("user not found")
ErrUsernameTaken = errors.New("username already taken")
ErrRefreshTokenNotFound = errors.New("refresh token not found")
)
// UserRepository 用户持久化接口。
type UserRepository interface {
// Create 创建用户,返回生成的 ID。
Create(ctx context.Context, username, passwordHash string) (string, error)
// FindByUsername 按用户名查找,不存在返回 ErrUserNotFound。
FindByUsername(ctx context.Context, username string) (*User, error)
// FindByID 按 ID 查找,不存在返回 ErrUserNotFound。
FindByID(ctx context.Context, id string) (*User, error)
// SaveRefreshToken 保存 refresh token hash。
SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error
// FindRefreshToken 按 token hash 查找,返回 user_id。不存在返回 ErrRefreshTokenNotFound。
FindRefreshToken(ctx context.Context, tokenHash string) (string, error)
// DeleteRefreshToken 按 token hash 删除。
DeleteRefreshToken(ctx context.Context, tokenHash string) error
// DeleteUserRefreshTokens 删除用户的所有 refresh token登出所有设备
DeleteUserRefreshTokens(ctx context.Context, userID string) error
}
// User 用户数据模型store 层)。
type User struct {
ID string
Username string
PasswordHash string
CreatedAt time.Time
UpdatedAt time.Time
}

View File

@@ -0,0 +1,120 @@
package store
import (
"context"
"sync"
"time"
"github.com/google/uuid"
)
// MemUserRepository 基于内存的 UserRepository 实现(测试用)。
type MemUserRepository struct {
mu sync.RWMutex
users map[string]*User // id -> user
byUsername map[string]string // username -> id
refreshTokens map[string]string // tokenHash -> userID
tokenExpiry map[string]time.Time // tokenHash -> expiresAt
}
// NewMemUserRepository 创建 MemUserRepository。
func NewMemUserRepository() *MemUserRepository {
return &MemUserRepository{
users: make(map[string]*User),
byUsername: make(map[string]string),
refreshTokens: make(map[string]string),
tokenExpiry: make(map[string]time.Time),
}
}
func (r *MemUserRepository) Create(_ context.Context, username, passwordHash string) (string, error) {
r.mu.Lock()
defer r.mu.Unlock()
if _, exists := r.byUsername[username]; exists {
return "", ErrUsernameTaken
}
id := uuid.New().String()
now := time.Now()
user := &User{
ID: id,
Username: username,
PasswordHash: passwordHash,
CreatedAt: now,
UpdatedAt: now,
}
r.users[id] = user
r.byUsername[username] = id
return id, nil
}
func (r *MemUserRepository) FindByUsername(_ context.Context, username string) (*User, error) {
r.mu.RLock()
defer r.mu.RUnlock()
id, ok := r.byUsername[username]
if !ok {
return nil, ErrUserNotFound
}
u := r.users[id]
copy := *u
return &copy, nil
}
func (r *MemUserRepository) FindByID(_ context.Context, id string) (*User, error) {
r.mu.RLock()
defer r.mu.RUnlock()
u, ok := r.users[id]
if !ok {
return nil, ErrUserNotFound
}
copy := *u
return &copy, nil
}
func (r *MemUserRepository) SaveRefreshToken(_ context.Context, userID, tokenHash string, expiresAt time.Time) error {
r.mu.Lock()
defer r.mu.Unlock()
r.refreshTokens[tokenHash] = userID
r.tokenExpiry[tokenHash] = expiresAt
return nil
}
func (r *MemUserRepository) FindRefreshToken(_ context.Context, tokenHash string) (string, error) {
r.mu.RLock()
defer r.mu.RUnlock()
userID, ok := r.refreshTokens[tokenHash]
if !ok {
return "", ErrRefreshTokenNotFound
}
if time.Now().After(r.tokenExpiry[tokenHash]) {
return "", ErrRefreshTokenNotFound
}
return userID, nil
}
func (r *MemUserRepository) DeleteRefreshToken(_ context.Context, tokenHash string) error {
r.mu.Lock()
defer r.mu.Unlock()
delete(r.refreshTokens, tokenHash)
delete(r.tokenExpiry, tokenHash)
return nil
}
func (r *MemUserRepository) DeleteUserRefreshTokens(_ context.Context, userID string) error {
r.mu.Lock()
defer r.mu.Unlock()
for hash, uid := range r.refreshTokens {
if uid == userID {
delete(r.refreshTokens, hash)
delete(r.tokenExpiry, hash)
}
}
return nil
}

View File

@@ -0,0 +1,147 @@
package store
import (
"context"
"errors"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/trace"
)
// PgUserRepository 基于 PostgreSQL 的 UserRepository 实现。
type PgUserRepository struct {
pool *pgxpool.Pool
}
// NewPgUserRepository 创建 PgUserRepository。
func NewPgUserRepository(pool *pgxpool.Pool) *PgUserRepository {
return &PgUserRepository{pool: pool}
}
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`,
username,
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.CreatedAt, &u.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
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`,
id,
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.CreatedAt, &u.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
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,
)
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()`,
tokenHash,
).Scan(&userID)
if errors.Is(err, pgx.ErrNoRows) {
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,
)
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,
)
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,185 @@
package store
import (
"context"
"errors"
"testing"
"time"
)
// newUserRepo 返回一个可测试的 UserRepository 实现。
// 如需测试 Pg 实现,可在此替换为连接真实 DB 的版本。
func newUserRepo() UserRepository {
return NewMemUserRepository()
}
func TestUserRepository_Create(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
id, err := repo.Create(ctx, "alice", "hash123")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
if id == "" {
t.Fatal("expected non-empty ID")
}
// 重复用户名应返回 ErrUsernameTaken
_, err = repo.Create(ctx, "alice", "hash456")
if !errors.Is(err, ErrUsernameTaken) {
t.Fatalf("expected ErrUsernameTaken, got %v", err)
}
}
func TestUserRepository_FindByUsername(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
_, err := repo.Create(ctx, "bob", "hash_bob")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
user, err := repo.FindByUsername(ctx, "bob")
if err != nil {
t.Fatalf("FindByUsername failed: %v", err)
}
if user.Username != "bob" {
t.Fatalf("expected username bob, got %s", user.Username)
}
if user.PasswordHash != "hash_bob" {
t.Fatalf("expected password hash hash_bob, got %s", user.PasswordHash)
}
// 不存在的用户
_, err = repo.FindByUsername(ctx, "nobody")
if !errors.Is(err, ErrUserNotFound) {
t.Fatalf("expected ErrUserNotFound, got %v", err)
}
}
func TestUserRepository_FindByID(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
id, err := repo.Create(ctx, "charlie", "hash_charlie")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
user, err := repo.FindByID(ctx, id)
if err != nil {
t.Fatalf("FindByID failed: %v", err)
}
if user.ID != id {
t.Fatalf("expected ID %s, got %s", id, user.ID)
}
if user.Username != "charlie" {
t.Fatalf("expected username charlie, got %s", user.Username)
}
// 不存在的 ID
_, err = repo.FindByID(ctx, "nonexistent-uuid")
if !errors.Is(err, ErrUserNotFound) {
t.Fatalf("expected ErrUserNotFound, got %v", err)
}
}
func TestUserRepository_RefreshToken(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
userID, err := repo.Create(ctx, "dave", "hash_dave")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
tokenHash := "abc123hash"
expiresAt := time.Now().Add(7 * 24 * time.Hour)
// 保存 token
if err := repo.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
t.Fatalf("SaveRefreshToken failed: %v", err)
}
// 查找 token
foundUserID, err := repo.FindRefreshToken(ctx, tokenHash)
if err != nil {
t.Fatalf("FindRefreshToken failed: %v", err)
}
if foundUserID != userID {
t.Fatalf("expected userID %s, got %s", userID, foundUserID)
}
// 不存在的 token
_, err = repo.FindRefreshToken(ctx, "nonexistent")
if !errors.Is(err, ErrRefreshTokenNotFound) {
t.Fatalf("expected ErrRefreshTokenNotFound, got %v", err)
}
// 删除 token
if err := repo.DeleteRefreshToken(ctx, tokenHash); err != nil {
t.Fatalf("DeleteRefreshToken failed: %v", err)
}
_, err = repo.FindRefreshToken(ctx, tokenHash)
if !errors.Is(err, ErrRefreshTokenNotFound) {
t.Fatalf("expected ErrRefreshTokenNotFound after delete, got %v", err)
}
}
func TestUserRepository_DeleteUserRefreshTokens(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
userID, err := repo.Create(ctx, "eve", "hash_eve")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
// 保存多个 token
for i := 0; i < 3; i++ {
tokenHash := "token_" + string(rune('a'+i))
expiresAt := time.Now().Add(7 * 24 * time.Hour)
if err := repo.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
t.Fatalf("SaveRefreshToken failed: %v", err)
}
}
// 删除用户所有 token
if err := repo.DeleteUserRefreshTokens(ctx, userID); err != nil {
t.Fatalf("DeleteUserRefreshTokens failed: %v", err)
}
// 验证全部删除
for i := 0; i < 3; i++ {
tokenHash := "token_" + string(rune('a'+i))
_, err := repo.FindRefreshToken(ctx, tokenHash)
if !errors.Is(err, ErrRefreshTokenNotFound) {
t.Fatalf("expected ErrRefreshTokenNotFound for token_%c, got %v", 'a'+i, err)
}
}
}
func TestUserRepository_ExpiredRefreshToken(t *testing.T) {
repo := newUserRepo()
ctx := context.Background()
userID, err := repo.Create(ctx, "frank", "hash_frank")
if err != nil {
t.Fatalf("Create failed: %v", err)
}
tokenHash := "expired_token"
expiresAt := time.Now().Add(-1 * time.Hour) // 已过期
if err := repo.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
t.Fatalf("SaveRefreshToken failed: %v", err)
}
// 过期 token 应返回 ErrRefreshTokenNotFound
_, err = repo.FindRefreshToken(ctx, tokenHash)
if !errors.Is(err, ErrRefreshTokenNotFound) {
t.Fatalf("expected ErrRefreshTokenNotFound for expired token, got %v", err)
}
}

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,25 +11,45 @@ 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"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源
// newUpgrader 根据配置创建 WebSocket upgrader
func newUpgrader(cfg *config.Config) websocket.Upgrader {
allowedOrigins := cfg.Server.AllowedOrigins
return websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
if len(allowedOrigins) == 0 {
return true // 未配置则允许所有来源(开发模式)
}
origin := r.Header.Get("Origin")
for _, o := range allowedOrigins {
if o == origin || o == "*" {
return true
}
}
return false
},
}
}
// 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 包调用)。
@@ -75,25 +96,77 @@ func (w *WSClient) SendError(err models.WsError) error {
}
// ServeWS 处理 WebSocket 升级请求。
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator) 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
return func(c *gin.Context) {
serveWS(c, sessionMgr, orch)
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr, limiter, scenarioRepo)
}
}
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator) {
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter, scenarioRepo store.UserScenarioRepository) {
// --- JWT 认证upgrade 前完成,失败直接返回 HTTP 错误) ---
token := c.Query("token")
if token == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
return
}
claims, err := tokenMgr.ValidateAccess(token)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
return
}
userID := claims.UserID
username := claims.Username
// --- conversation_id 处理upgrade 前校验归属) ---
conversationID := c.Query("conversation_id")
if conversationID != "" {
sess, err := sessionMgr.Get(c.Request.Context(), conversationID)
if err != nil || sess.UserID != userID {
c.JSON(http.StatusUnauthorized, gin.H{"error": "SESSION_NOT_FOUND"})
return
}
}
// 生成连接级 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()
// 创建会话
sessionID, err := sessionMgr.Create(context.Background(), models.DefaultConfig())
if err != nil {
logger.Log.Errorw("create session failed", "error", err)
return
// 创建或复用会话
var sessionID string
if conversationID != "" {
sessionID = conversationID
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 {
log := trace.FromContext(ctx)
log.Errorw("create session failed", "error", err)
return
}
ctx = trace.WithSessionID(ctx, sessionID)
}
client := &Client{
@@ -108,9 +181,10 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_ = client.SendJSON(models.WsConnected{
Type: "connected",
SessionID: sessionID,
ServerVersion: "0.1.0",
ServerVersion: version,
})
logger.Log.Infow("client connected", "session", sessionID)
log := trace.FromContext(ctx)
log.Infow("client connected", "user_id", userID, "username", username)
// 心跳检测
lastPong := time.Now()
@@ -122,13 +196,14 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
// 启动心跳检查 goroutine
done := make(chan struct{})
go func() {
ticker := time.NewTicker(30 * time.Second)
ticker := time.NewTicker(heartbeatInterval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if time.Since(lastPong) > 60*time.Second {
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
if time.Since(lastPong) > heartbeatTimeout {
log := trace.FromContext(ctx)
log.Warnw("heartbeat timeout")
conn.Close()
return
}
@@ -143,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
}
@@ -159,6 +235,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
switch envelope.Type {
case "ping":
lastPong = time.Now() // 刷新心跳计时器
_ = client.SendJSON(models.WsPong{Type: "pong"})
case "query":
@@ -167,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, 20)
// 创建可取消的 context
ctx, cancel := context.WithCancel(context.Background())
processCtx, cancel := context.WithCancel(queryCtx)
client.mu.Lock()
client.cancelFuncs[msg.RequestID] = cancel
client.mu.Unlock()
@@ -203,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)
}
}()
@@ -219,15 +310,72 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
TTSEnabled: msg.Payload.TTSEnabled,
DetailLevel: msg.Payload.DetailLevel,
Language: msg.Payload.Language,
Scenario: msg.Payload.Scenario,
}
if err := client.sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
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)
@@ -255,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

@@ -2,17 +2,20 @@ package ws
import (
"encoding/base64"
"net/http"
"net/http/httptest"
"strings"
"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"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/orchestrator"
@@ -45,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 {
@@ -131,19 +133,29 @@ func (m *MockOrchestrator) ProcessQuery(
// --- 测试辅助函数 ---
// setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。
// 返回的 wsURL 已包含有效 token可直接连接。
func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, string) {
t.Helper()
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
t.Cleanup(func() { sessionMgr.Stop() })
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
r := gin.New()
r.GET("/ws", ServeWS(sessionMgr, orch))
cfg := &config.Config{
App: config.AppConfig{Version: "test"},
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
Session: config.SessionConfig{MaxHistory: 20},
}
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr, nil, nil))
srv := httptest.NewServer(r)
// 构造 WebSocket URL
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
// 生成有效 token 并构造 WebSocket URL
token, _, err := tokenMgr.GeneratePair("test-user", "testuser")
require.NoError(t, err)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
return srv, wsURL
}
@@ -181,7 +193,7 @@ func TestWS_Connected(t *testing.T) {
msg := readJSON(t, conn)
assert.Equal(t, "connected", msg["type"])
assert.NotEmpty(t, msg["session_id"])
assert.Equal(t, "0.1.0", msg["server_version"])
assert.Equal(t, "test", msg["server_version"])
}
// TestWS_PingPong 验证 ping/pong 心跳。
@@ -209,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)
@@ -320,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")
@@ -561,3 +573,156 @@ func TestWS_QueryWithTTSDisabled(t *testing.T) {
err = conn.ReadJSON(&extra)
assert.Error(t, err, "不应有额外消息")
}
// --- 认证测试辅助 ---
// setupTestServerEx 创建测试服务器,返回 tokenMgr 和 sessionMgr 以便测试控制。
func setupTestServerEx(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, *auth.TokenManager, *session.MemoryManager) {
t.Helper()
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
t.Cleanup(func() { sessionMgr.Stop() })
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
r := gin.New()
cfg := &config.Config{
App: config.AppConfig{Version: "test"},
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
Session: config.SessionConfig{MaxHistory: 20},
}
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr, nil, nil))
srv := httptest.NewServer(r)
return srv, tokenMgr, sessionMgr
}
// httpGet 发送 HTTP GET 并返回状态码。
func httpGet(t *testing.T, url string) int {
t.Helper()
resp, err := http.Get(url)
require.NoError(t, err)
resp.Body.Close()
return resp.StatusCode
}
// --- 认证测试用例 ---
// TestWS_AuthMissingToken 验证无 token 时返回 401。
func TestWS_AuthMissingToken(t *testing.T) {
srv, _, _ := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
httpURL := srv.URL + "/ws"
status := httpGet(t, httpURL)
assert.Equal(t, http.StatusUnauthorized, status)
}
// TestWS_AuthInvalidToken 验证无效 token 时返回 401。
func TestWS_AuthInvalidToken(t *testing.T) {
srv, _, _ := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
httpURL := srv.URL + "/ws?token=invalid-token"
status := httpGet(t, httpURL)
assert.Equal(t, http.StatusUnauthorized, status)
}
// TestWS_AuthExpiredToken 验证过期 token 时返回 401。
func TestWS_AuthExpiredToken(t *testing.T) {
// 创建一个 access TTL 极短的 tokenMgr
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
defer sessionMgr.Stop()
tokenMgr := auth.NewTokenManager("test-secret", -1*time.Minute, 7*24*time.Hour) // 已过期
r := gin.New()
cfg := &config.Config{
App: config.AppConfig{Version: "test"},
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
Session: config.SessionConfig{MaxHistory: 20},
}
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr, nil, nil))
srv := httptest.NewServer(r)
defer srv.Close()
token, _, err := tokenMgr.GeneratePair("test-user", "testuser")
require.NoError(t, err)
httpURL := srv.URL + "/ws?token=" + token
status := httpGet(t, httpURL)
assert.Equal(t, http.StatusUnauthorized, status)
}
// TestWS_AuthValidToken 验证有效 token 能成功建立 WS 连接。
func TestWS_AuthValidToken(t *testing.T) {
srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
token, _, err := tokenMgr.GeneratePair("user-1", "alice")
require.NoError(t, err)
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
conn := connectWS(t, wsURL)
msg := readJSON(t, conn)
assert.Equal(t, "connected", msg["type"])
assert.NotEmpty(t, msg["session_id"])
}
// TestWS_AuthConversationIDResume 验证通过 conversation_id 恢复已有对话。
func TestWS_AuthConversationIDResume(t *testing.T) {
srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
userID := "user-1"
// 先创建一个属于该用户的 session
ctx := context.Background()
sessionID, err := sessionMgr.Create(ctx, userID, models.DefaultConfig())
require.NoError(t, err)
token, _, err := tokenMgr.GeneratePair(userID, "alice")
require.NoError(t, err)
// 带 conversation_id 连接
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") +
"/ws?token=" + token + "&conversation_id=" + sessionID
conn := connectWS(t, wsURL)
msg := readJSON(t, conn)
assert.Equal(t, "connected", msg["type"])
assert.Equal(t, sessionID, msg["session_id"], "应复用已有 session")
}
// TestWS_AuthConversationIDNotFound 验证 conversation_id 不存在时返回 401。
func TestWS_AuthConversationIDNotFound(t *testing.T) {
srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
token, _, err := tokenMgr.GeneratePair("user-1", "alice")
require.NoError(t, err)
httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=nonexistent-id"
status := httpGet(t, httpURL)
assert.Equal(t, http.StatusUnauthorized, status)
}
// TestWS_AuthConversationIDOwnership 验证 conversation_id 不属于当前用户时返回 401。
func TestWS_AuthConversationIDOwnership(t *testing.T) {
srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{})
defer srv.Close()
ctx := context.Background()
// user-A 创建 session
sessionID, err := sessionMgr.Create(ctx, "user-A", models.DefaultConfig())
require.NoError(t, err)
// user-B 尝试连接该 session
token, _, err := tokenMgr.GeneratePair("user-B", "bob")
require.NoError(t, err)
httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=" + sessionID
status := httpGet(t, httpURL)
assert.Equal(t, http.StatusUnauthorized, status, "非 owner 访问应返回 401")
}

View File

@@ -0,0 +1,5 @@
-- 删除 Refresh Token 表(自动删除相关索引)
DROP TABLE IF EXISTS refresh_tokens;
-- 删除用户表(自动删除相关索引)
DROP TABLE IF EXISTS users;

View File

@@ -0,0 +1,42 @@
-- 用户表
CREATE TABLE IF NOT EXISTS users (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
username VARCHAR(64) NOT NULL UNIQUE,
password_hash VARCHAR(255) NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
);
-- 用户名索引(用于登录查询)
CREATE INDEX IF NOT EXISTS idx_users_username ON users(username);
-- 表和列注释
COMMENT ON TABLE users IS '用户表,存储系统所有注册用户的基本信息';
COMMENT ON COLUMN users.id IS '用户唯一标识符 (UUID)';
COMMENT ON COLUMN users.username IS '用户名,最大 64 字符,全局唯一';
COMMENT ON COLUMN users.password_hash IS '密码哈希值,使用 bcrypt 算法cost=10';
COMMENT ON COLUMN users.created_at IS '用户注册时间';
COMMENT ON COLUMN users.updated_at IS '用户信息最后更新时间';
-- Refresh Token 表
CREATE TABLE IF NOT EXISTS refresh_tokens (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
token_hash VARCHAR(64) NOT NULL UNIQUE,
expires_at TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
);
-- Token hash 索引(用于刷新验证)
CREATE INDEX IF NOT EXISTS idx_refresh_tokens_token_hash ON refresh_tokens(token_hash);
-- 用户 ID 索引(用于登出所有设备)
CREATE INDEX IF NOT EXISTS idx_refresh_tokens_user_id ON refresh_tokens(user_id);
-- 表和列注释
COMMENT ON TABLE refresh_tokens IS 'Refresh Token 表,用于 JWT 双 token 机制的长期身份验证';
COMMENT ON COLUMN refresh_tokens.id IS 'Token 唯一标识符 (UUID)';
COMMENT ON COLUMN refresh_tokens.user_id IS '所属用户 ID外键关联 users 表,用户删除时级联删除';
COMMENT ON COLUMN refresh_tokens.token_hash IS 'Token 哈希值,使用 SHA-256 算法,十六进制编码 (64 字符)';
COMMENT ON COLUMN refresh_tokens.expires_at IS 'Token 过期时间,默认有效期 7 天';
COMMENT ON COLUMN refresh_tokens.created_at IS 'Token 创建时间';

View File

@@ -0,0 +1 @@
DROP TABLE IF EXISTS messages;

View File

@@ -0,0 +1,28 @@
-- 消息表
CREATE TABLE IF NOT EXISTS messages (
id BIGSERIAL PRIMARY KEY,
session_id UUID NOT NULL,
role VARCHAR(10) NOT NULL,
content TEXT NOT NULL,
tokens_used INTEGER NOT NULL DEFAULT 0,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
CONSTRAINT check_tokens_non_negative CHECK (tokens_used >= 0)
);
-- 按会话查询消息(分页核心索引)
CREATE INDEX IF NOT EXISTS idx_messages_session_id_created_at
ON messages(session_id, created_at);
-- 按会话查询最后一条消息
CREATE INDEX IF NOT EXISTS idx_messages_session_id_id_desc
ON messages(session_id, id DESC);
-- 表和列注释
COMMENT ON TABLE messages IS '消息表,存储所有会话的消息记录';
COMMENT ON COLUMN messages.id IS '消息唯一标识符,自增序列';
COMMENT ON COLUMN messages.session_id IS '所属会话 ID关联 sessions 表';
COMMENT ON COLUMN messages.role IS '消息角色,可选值: ''user'' (用户), ''assistant'' (AI 助手), ''system'' (系统)';
COMMENT ON COLUMN messages.content IS '消息内容,无长度限制';
COMMENT ON COLUMN messages.tokens_used IS '消息消耗的 token 数量,用于计费统计';
COMMENT ON COLUMN messages.created_at IS '消息创建时间';

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