Compare commits
472 Commits
dc5e49d216
...
develop
| Author | SHA1 | Date | |
|---|---|---|---|
| b1c3958e9b | |||
| f9a01269d2 | |||
| 74f1a8eb34 | |||
| 4dd4138197 | |||
| 5ef059eb40 | |||
| 9523fe6650 | |||
| a219934f4b | |||
| 72a8d4803e | |||
| dcf53a3783 | |||
| b5ec6551bf | |||
| 79d321c78a | |||
| 57bd8c72b8 | |||
| d640b5b41b | |||
| f1b9e4fcc1 | |||
| 18aa5f1949 | |||
| a16628119a | |||
| 3888941f85 | |||
| 0a816049fa | |||
| 140a178993 | |||
| 31c1720c00 | |||
| c6093c01c1 | |||
| 3267f35edb | |||
| a309c269e0 | |||
| 22930f0080 | |||
| c00b8a83d9 | |||
| 36007babcb | |||
| f01c84d89e | |||
| c3808ff38c | |||
| 2740bea755 | |||
| 69b9694cfe | |||
| 70ee212ea2 | |||
| f08f547e76 | |||
| 03a19befe1 | |||
| c4b68b77d1 | |||
| f2a883bda8 | |||
| 4cc0713459 | |||
| 2180751a9b | |||
| 964c5c967e | |||
| dcd08d031c | |||
| 6bca4e0d47 | |||
| 881c3f9853 | |||
| 2881941be6 | |||
| 1b99d24fc1 | |||
| f9e68b96b9 | |||
| c58c6b59a5 | |||
| e24fb2cc19 | |||
| 7a705744d0 | |||
| 583c33727a | |||
| 6cbabb63bb | |||
| b4fbf8625b | |||
| 239f8f9877 | |||
| d430e6e5b2 | |||
| edc66625ba | |||
| 9dce107a84 | |||
| 1c7dd708a0 | |||
| d0e4bdaeec | |||
| 949a707e0f | |||
| 55f7f183a7 | |||
| 6b4b033df3 | |||
| 76d331c885 | |||
| ad700743ef | |||
| 6ab4776e08 | |||
| c17798ec67 | |||
| 8a43f4406a | |||
| 065673fae2 | |||
| 03c27e7790 | |||
| 0a59173476 | |||
| 939e43acd0 | |||
| 87c3e7a8dd | |||
| d6e9555a97 | |||
| 9c763ec12a | |||
| 1bab02ae84 | |||
| 89d7b7c17c | |||
| 34d498510e | |||
| 104b28efd3 | |||
| a02a8bc374 | |||
| 0d99d06f06 | |||
| 8b18953010 | |||
| 0108ef2064 | |||
| e0dc8272a5 | |||
| 11c3955bd6 | |||
| a66ab764d9 | |||
| 51117b43f6 | |||
| 492fb06c08 | |||
| 6967dd7b2e | |||
| eea5c07eaa | |||
| 1079e22699 | |||
| ad5d90e344 | |||
| c094fe0867 | |||
| 032de796c8 | |||
| 8b4acb3ce7 | |||
| 9e5f691056 | |||
| 03127aa01a | |||
| 99fcd6bc29 | |||
| d0f5f5c94d | |||
| 361c5d07d3 | |||
| 910e71b6f0 | |||
| 19645be04e | |||
| d53265755a | |||
| 252cdcc8e7 | |||
| 515d7ae034 | |||
| 311330cea1 | |||
| 7adf81c6e5 | |||
| ea00939c13 | |||
| b74fb3564d | |||
| 3b6226394b | |||
| a6df8c9131 | |||
| 023c834074 | |||
| 3e00e39e8a | |||
| 9ad486d117 | |||
| 6af26ffc91 | |||
| 7ee6918015 | |||
| 2fb23852ef | |||
| ed95ce56a8 | |||
| e151c5b665 | |||
| dffd8bd4a5 | |||
| 3140a660e4 | |||
| 23e0e22a12 | |||
| ba7c5ed5ea | |||
| cc7a333b6d | |||
| f8af2f0ccc | |||
| d6059ae397 | |||
| d4398e55fe | |||
| f5440c9e2d | |||
| 190f40908f | |||
| ddb90c58ef | |||
| 340c26b7a6 | |||
| ca07188eda | |||
| 15a50f855b | |||
| e37dc7074d | |||
| f86c2560cd | |||
| 90c5ad7724 | |||
| 45e2c4f37b | |||
| b7d0edb6da | |||
| 159bf278e6 | |||
| 542126df90 | |||
| e56600e408 | |||
| 04215e7a53 | |||
| 9ba4eb4825 | |||
| bc7eb7409c | |||
| 0a8a20f3f8 | |||
| e20f58b4ae | |||
| ab07e01adf | |||
| 0447fdacac | |||
| 12cf59f1e8 | |||
| 4f735e6d29 | |||
| 32aea44f3b | |||
| 4d51bcefb0 | |||
| fe5ac20a1f | |||
| 2a13ee9c89 | |||
| df35ff73b5 | |||
| 92e6636d05 | |||
| 12e2d24f37 | |||
| ae100c4a75 | |||
| 4a5905307c | |||
| 2a7d4c74d4 | |||
| 6bb4773e07 | |||
| 97e125234b | |||
| 39073b7673 | |||
| 85d47c1fc4 | |||
| 4ff8cec312 | |||
| 3572b867c0 | |||
| dfed964f76 | |||
| 4651ae185b | |||
| 1c65433d40 | |||
| 590c4592ea | |||
| 15cd157f45 | |||
| 970f10a274 | |||
| d78cdb509d | |||
| ea70d2efc6 | |||
| a2a28a9f56 | |||
| 898e30b526 | |||
| 87c54b80c0 | |||
| 7288f443f1 | |||
| 488b1e62ba | |||
| 50c84fce88 | |||
| ae09fba400 | |||
| f8a79a3b0b | |||
| ae9a27c300 | |||
| 1f7cd407a8 | |||
| 6f81212997 | |||
| 402ad8c949 | |||
| 52015fa6c6 | |||
| 97706ea197 | |||
| d53de4f33f | |||
| 3dc2015a91 | |||
| de78d60959 | |||
| 93d5a90495 | |||
| 5910e02a66 | |||
| 88ee1e2548 | |||
| b8fd6a9330 | |||
| c2c9a6d392 | |||
| 795ace75d6 | |||
| dd846876ba | |||
| 898d1474ec | |||
| 582f68f68b | |||
| 7d328f2552 | |||
| eb1b90445f | |||
| e5537aaa1e | |||
| 556666c046 | |||
| 4235c5cae0 | |||
| 41f393c09d | |||
| 04dac0b673 | |||
| 4c7430d6f4 | |||
| 765cb34019 | |||
| 9576884619 | |||
| 4ffd84510e | |||
| 4b731b5ac0 | |||
| fd5c7712f8 | |||
| f38fbf0527 | |||
| b10a508356 | |||
| c0b4eeda46 | |||
| e6481e0faa | |||
| a04275cc76 | |||
| dca37f3e48 | |||
| 16302af7d2 | |||
| 5d8cacf16d | |||
| dbfdf3c3e5 | |||
| 54454ae2d7 | |||
| 720c2e2b5b | |||
| c2bae4e3b7 | |||
| 838493145f | |||
| 943f36e9ec | |||
| 083ede8b6d | |||
| 14242eb896 | |||
| 6db5bef519 | |||
| 1ddd75d34f | |||
| 7ea6c0d4f2 | |||
| 97357013bb | |||
| 057a226610 | |||
| b60e513317 | |||
| 638a387550 | |||
| 9dc9025d10 | |||
| 92e47718bc | |||
| 6ce36e3ed8 | |||
| 6032608e49 | |||
| 668da26e23 | |||
| 44f4f3ce1a | |||
| c3118dd72a | |||
| f22243a2d2 | |||
| 8ff3c519ac | |||
| ffc9b4a320 | |||
| eba0607f94 | |||
| e89585e453 | |||
| e11934cf70 | |||
| 1f19435951 | |||
| 16105f5ac6 | |||
| 386c918591 | |||
| 966b30218a | |||
| d724a8e9f9 | |||
| 8360e31328 | |||
| 1109ab41f7 | |||
| 542720695d | |||
| 8de956719c | |||
| a1ae7740f1 | |||
| b57adf153b | |||
| c3a32ce276 | |||
| 7b745018c5 | |||
| f4515ce5e4 | |||
| 96f4bc7abb | |||
| dae5722945 | |||
| 3c5c4943e8 | |||
| 80c6b1b56e | |||
| 905b56640e | |||
| 2aa3c98ab6 | |||
| d01aeaba68 | |||
| f902e05e31 | |||
| 62baa656ee | |||
| ec8555d44b | |||
| 6487a8ecab | |||
| 9ff971fd89 | |||
| c01ee1d14c | |||
| 9aaed88cc1 | |||
| 15b437043f | |||
| fdea3aa9d1 | |||
| 430811ffc4 | |||
| f878988048 | |||
| e455af6b2d | |||
| 5211d14d50 | |||
| 5b8cbb025f | |||
| fca9108882 | |||
| 20597101f5 | |||
| 5e83531d54 | |||
| 9842b4a457 | |||
| ce13e5a048 | |||
| 36827edfea | |||
| af73e78aa3 | |||
| c9c174f400 | |||
| c0972856c6 | |||
| 0ff810a297 | |||
| bf1e24453c | |||
| 0edafbbf8a | |||
| 4b30e67c2e | |||
| 07211675e9 | |||
| 7a0443ebcc | |||
| fe053f94dc | |||
| 9433612c4b | |||
| 92f3f45c41 | |||
| 74cd6c7d2d | |||
| 8f8e9dea5f | |||
| 2bb72096de | |||
| 6b99336941 | |||
| ae91f877fe | |||
| f23a425f1e | |||
| 12eec6fc0e | |||
| a0769956f2 | |||
| 0f3bf432fd | |||
| 0ba2f9807c | |||
| 9379d3ef82 | |||
| cb0c747133 | |||
| 91fdc007f7 | |||
| 18c889b9b8 | |||
| ffcb11625a | |||
| da8818c3be | |||
| 443f613fbd | |||
| f46f8ffa5a | |||
| 159db37a60 | |||
| 15b1dd56fa | |||
| 35372438a5 | |||
| daa4d23895 | |||
| 9e813f0a08 | |||
| 97484735cf | |||
| a128ed1471 | |||
| e33b5650a5 | |||
| b358bf05b5 | |||
| 4ef1e9f08b | |||
| d32112e074 | |||
| 81e50de7f6 | |||
| 2f05b5fa2b | |||
| 304a45b0ad | |||
| e95b1603c1 | |||
| 7153a7adf6 | |||
| 73eae8a118 | |||
| 1a6325698a | |||
| 4ebac835de | |||
| d8bbdb1744 | |||
| da5c727d8c | |||
| 7b3f4706e1 | |||
| 65718a8460 | |||
| 84f5303632 | |||
| 1c74fdb894 | |||
| d75cdf95c8 | |||
| 7108d2a58b | |||
| 80e45bb8ff | |||
| 81c2b64e5f | |||
| dffc8aa538 | |||
| dfccc824be | |||
| 312e762aff | |||
| 01026d1846 | |||
| d8630a8f26 | |||
| 41cacaa740 | |||
| 51ed6aa563 | |||
| 563ca12790 | |||
| bfc896ed35 | |||
| f6a2348c04 | |||
| 24a6ce5295 | |||
| 96cf009228 | |||
| 544716fe27 | |||
| 8a966c29b7 | |||
| aae0629b2d | |||
| 03b3566822 | |||
| f1ce28966c | |||
| 63f8cc279d | |||
| 4a7cf89f29 | |||
| 22aefbc216 | |||
| c34be3a996 | |||
| 00abc6c1e1 | |||
| 32f3efbb7a | |||
| eb7ebecfc0 | |||
| b1e71be3d9 | |||
| 1d938055e8 | |||
| bf0bb83f7d | |||
| 37af0dbe06 | |||
| 70dfa1c8f8 | |||
| 5aa383cab9 | |||
| e967d89e7e | |||
| b8ed49be63 | |||
| 2db0e3b0b6 | |||
| ca09a1fb72 | |||
| 6e0c67e1cb | |||
| a6ce6d9c4f | |||
| a7d679d04f | |||
| c31ecc46fc | |||
| 322061c387 | |||
| 4458ee82b2 | |||
| 420871cd7e | |||
| 03a05b89a0 | |||
| e3eb4f264b | |||
| a404fb97f9 | |||
| 3d3de828fc | |||
| 4dd89be79f | |||
| b0a7ce885e | |||
| 213aa67e9f | |||
| 89386ac00d | |||
| a666b82e64 | |||
| 5eeb592bf0 | |||
| 836a61dfd2 | |||
| 00f10088d0 | |||
| 2a235a94b2 | |||
| 3c665cde5d | |||
| 0d9d3db73b | |||
| 188782fa66 | |||
| f764b3a7d8 | |||
| 0227822120 | |||
| f17c3c24a7 | |||
| 6b98f35abc | |||
| 9ce9d8c9f6 | |||
| 0b2bbb827d | |||
| 5afde31a9e | |||
| 0e1f34d7c9 | |||
| a6a6a16185 | |||
| f07916ab0f | |||
| b7e803fa90 | |||
| 093dc26fa0 | |||
| 4aa9f31e57 | |||
| dd64ed23ee | |||
| 085e19038b | |||
| 314554ded7 | |||
| 19535e6a37 | |||
| 8e43fb134e | |||
| c9c697fe64 | |||
| c160e787a6 | |||
| 041dee319b | |||
| 29e186c062 | |||
| acc7dfce43 | |||
| 450c2c66c3 | |||
| 90f4b907a7 | |||
| 96b4c4899a | |||
| 078cfe6278 | |||
| 8f9b0b881e | |||
| 6fcbf023b4 | |||
| 22784bf421 | |||
| 33effe161a | |||
| fe4dedec71 | |||
| 26f87ce35a | |||
| ba6d565c0f | |||
| f172e6c468 | |||
| 6b8bd3c554 | |||
| 688e13a491 | |||
| 991ae4834c | |||
| da87ec776b | |||
| c523f53724 | |||
| bce3b73480 | |||
| d5d584c93b | |||
| b76a71c128 | |||
| 2e165a459c | |||
| 7e2537104d | |||
| f5536005f6 | |||
| e18a2004a0 | |||
| 2c6489b9a6 | |||
| 8dc59d2c09 | |||
| 7e3762b589 | |||
| 0cafe94f3a | |||
| 8033807f7c | |||
| 6a47c4dfbb | |||
| 0ef7aa657d | |||
| 46cf35967e | |||
| 360b2a4f50 | |||
| ee557eaa3d | |||
| 87b6407e2e | |||
| 2e222cbe73 | |||
| d5ee6deeb9 | |||
| 734c93a866 | |||
| 31dfabd882 | |||
| bab67a8749 | |||
| 5d70bc8862 | |||
| c0279e94c7 | |||
| 9ccd4d6238 | |||
| c68657bf30 | |||
| 5e89ab01bd | |||
| 60fc3ac5d1 | |||
| 0a0eaf7eb9 |
@@ -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 ./...
|
||||
35
.gitea/workflows/deploy.yml
Normal file
35
.gitea/workflows/deploy.yml
Normal 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
|
||||
@@ -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
|
||||
8
.gitignore
vendored
8
.gitignore
vendored
@@ -4,6 +4,7 @@ frontend/dist/
|
||||
|
||||
# ---- 后端 ----
|
||||
backend/bin/
|
||||
backend/server
|
||||
|
||||
# ---- 环境变量 ----
|
||||
.env
|
||||
@@ -18,5 +19,12 @@ Thumbs.db
|
||||
.idea/
|
||||
.vscode/
|
||||
|
||||
# ---- Playwright MCP ----
|
||||
.playwright-mcp/
|
||||
|
||||
# ---- Obsidian ----
|
||||
.obsidian/
|
||||
.claudian/
|
||||
修改过程笔记/
|
||||
学习复盘/
|
||||
docs/follow-up/
|
||||
|
||||
176
CLAUDE.md
176
CLAUDE.md
@@ -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-4o(LLM)、Deepgram(STT)、OpenAI TTS。仅通过 Go 网关访问,浏览器不直连。
|
||||
**AI 编排流水线**(Eino Graph 7 节点 DAG):`STT → History → ChatModel → Msg2Str → Splitter → TTS → Done`。LLM token 通过 Callback 实时推送,TTS 逐句并行合成。
|
||||
|
||||
**关键模式**:LLM 文本流和 TTS 音频流并行推送给客户端,以最小化感知延迟。
|
||||
**会话存储**(TieredManager):L1 Memory → L2 Redis → L3 PostgreSQL 三级存储,30 分钟 TTL,Redis 故障自动降级。
|
||||
|
||||
**存储**:冷热分离 —— Redis 存实时会话状态,PostgreSQL 存对话历史和用量统计(MVP 后引入)。Repository 接口模式(`HistoryRepository`、`UsageRepository`),MVP 用内存实现。
|
||||
**鉴权**:JWT 双 token 轮转(Access 120min + Refresh 7d),重放攻击检测(DB hash 校验),Redis 缓存装饰器。
|
||||
|
||||
## 技术栈
|
||||
|
||||
| 层级 | 技术 |
|
||||
前端:React 18 + TypeScript + Vite,VAD(@ricky0123/vad-web),ONNX Runtime,国际化(zh-CN / en-US / ja-JP)
|
||||
后端:Go 1.25+, Gin, WebSocket, Viper, Zap, CloudWeGo Eino Graph
|
||||
AI:DashScope qwen3-vl-plus, MiMo ASR/TTS(可切换 Deepgram/OpenAI TTS)
|
||||
存储:PostgreSQL 15 + Redis 7
|
||||
CI/CD:Gitea Actions(`.gitea/workflows/deploy.yml`),push main/v2 自动构建部署到自建 aliyun runner
|
||||
前端测试:**暂无**(package.json 无 test 脚本,无测试框架配置)
|
||||
|
||||
## 配置体系
|
||||
|
||||
配置优先级:**环境变量 > `config.{APP_ENV}.yaml` > `config.yaml` > 代码默认值**
|
||||
|
||||
配置文件位于 `backend/config/`:
|
||||
- `config.yaml` — 基础配置(dev 默认值)
|
||||
- `config.dev.yaml` — 开发环境覆盖(可选)
|
||||
- `config.prod.yaml` — 生产环境覆盖(可选)
|
||||
|
||||
环境切换:`APP_ENV=dev|prod`(dev 默认,prod 启用限流 + 严格 CORS + Release 模式)
|
||||
|
||||
敏感信息(API Key、JWT Secret、数据库密码)**只能通过环境变量或 `.env` 文件注入**,不写入 YAML 配置文件。核心环境变量(参考 `backend/.env.example`):
|
||||
|
||||
| 变量 | 说明 |
|
||||
|------|------|
|
||||
| 前端 | React 18, TypeScript, Vite, 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 Key(DashScope) |
|
||||
| `CAMTALK_AI_STT_API_KEY` | STT API Key(MiMo/Deepgram) |
|
||||
| `CAMTALK_AI_TTS_API_KEY` | TTS API Key(MiMo/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` — AuthProvider(JWT 自动刷新 + 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
559
README.md
@@ -1,2 +1,561 @@
|
||||
# CamTalk
|
||||
|
||||
<div align="center">
|
||||
|
||||
**多模态实时 AI 视觉对话助手**
|
||||
|
||||
用户通过摄像头和麦克风与 AI 交互,AI 理解视觉场景和语音输入后,以文字和语音形式给出自然回应
|
||||
|
||||
[](https://opensource.org/licenses/MIT)
|
||||
[](https://go.dev/)
|
||||
[](https://react.dev/)
|
||||
[](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** 重启浏览器
|
||||
|
||||
 -->
|
||||
|
||||
## ✨ 核心特性
|
||||
|
||||
- 🎥 **多模态理解**:摄像头视觉 + 麦克风语音双输入,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>声明式 DAG,Stream 模式,Callback AOP</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>STT</b></td>
|
||||
<td>MiMo ASR(默认)/ Deepgram</td>
|
||||
<td>实时语音识别,多语言支持</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>LLM</b></td>
|
||||
<td>DashScope qwen3-vl-plus</td>
|
||||
<td>多模态推理(通过 eino-ext OpenAI 接入)</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>TTS</b></td>
|
||||
<td>MiMo TTS(默认)/ OpenAI TTS</td>
|
||||
<td>自然语音合成</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>存储</b></td>
|
||||
<td>PostgreSQL 15 + Redis 7</td>
|
||||
<td>三级存储架构:Memory → Redis → PG</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>认证</b></td>
|
||||
<td>JWT (HS256) + bcrypt</td>
|
||||
<td>双 token 轮转 + Refresh Token Rotation</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>配置</b></td>
|
||||
<td>Viper + godotenv</td>
|
||||
<td>YAML + .env + 环境变量覆盖</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><b>日志</b></td>
|
||||
<td>Zap</td>
|
||||
<td>高性能结构化日志 + Trace ID 追踪</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
## 📁 项目结构
|
||||
|
||||
```
|
||||
CamTalk/
|
||||
├── 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 # AuthProvider(JWT 自动刷新)
|
||||
│ │ ├── 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 API(Auth/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_DSN(PostgreSQL 连接字符串)
|
||||
# - CAMTALK_AI_LLM_API_KEY(DashScope API Key)
|
||||
# - CAMTALK_AI_STT_API_KEY(MiMo/Deepgram API Key)
|
||||
# - CAMTALK_AI_TTS_API_KEY(MiMo/OpenAI API Key)
|
||||
```
|
||||
|
||||
#### 3. 启动后端
|
||||
|
||||
```bash
|
||||
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 逐 token(Callback) |
|
||||
| `llm_done` | LLM 推理完成 | Done 节点执行 |
|
||||
| `tts_audio` | TTS 音频片段 | TTS 节点逐句合成 |
|
||||
| `error` | 错误通知 | 任意节点失败 |
|
||||
| `pong` | 心跳响应 | 响应 `ping` |
|
||||
|
||||
**心跳机制**:
|
||||
- 客户端每 30 秒发送 `ping`
|
||||
- 服务端 60 秒无消息自动断连
|
||||
- 断连后自动重连(指数退避 1s → 30s)
|
||||
|
||||
完整协议定义见 [docs/02-接口文档.md](docs/02-接口文档.md)
|
||||
|
||||
## 🔐 认证体系
|
||||
|
||||
CamTalk 采用 **JWT 双 token 轮转 + Refresh Token Rotation** 安全机制:
|
||||
|
||||
### 双 Token 设计
|
||||
|
||||
| Token | 有效期 | 存储位置 | 用途 |
|
||||
|-------|-------|---------|------|
|
||||
| `access_token` | 120 分钟 | 前端内存(推荐)/ localStorage | 访问受保护资源 |
|
||||
| `refresh_token` | 7 天 | httpOnly Cookie(推荐)/ localStorage | 刷新 access_token |
|
||||
|
||||
### Refresh Token Rotation
|
||||
|
||||
每次刷新 token 时:
|
||||
1. 验证 `refresh_token` 签名和有效期
|
||||
2. 查询数据库中的 SHA256 哈希
|
||||
3. **如果哈希不存在** → 检测到 token 复用 → **吊销该用户所有 token**
|
||||
4. 删除旧 refresh_token,生成新 token pair
|
||||
5. 返回新 access_token + refresh_token
|
||||
|
||||
**防重放攻击**:旧 refresh_token 立即失效,复用时触发全局吊销,强制所有设备重新登录。
|
||||
|
||||
### REST API 端点
|
||||
|
||||
- `POST /api/auth/register` — 用户注册
|
||||
- `POST /api/auth/login` — 用户登录
|
||||
- `POST /api/auth/refresh` — 刷新 token
|
||||
- `POST /api/auth/logout` — 登出(需认证)
|
||||
- `GET /api/conversations` — 获取对话列表(需认证)
|
||||
- `POST /api/conversations` — 创建对话(需认证)
|
||||
- `GET /api/health` — 健康检查
|
||||
|
||||
详细设计见 [docs/10-鉴权体系.md](docs/10-鉴权体系.md)
|
||||
|
||||
## 💾 三级存储架构
|
||||
|
||||
**TieredManager** 实现会话状态的三级存储,平衡性能与可靠性:
|
||||
|
||||
```
|
||||
┌─────────────┐
|
||||
│ L1 Memory │ ← 微秒级读写,进程内缓存
|
||||
├─────────────┤
|
||||
│ L2 Redis │ ← 毫秒级访问,跨实例共享
|
||||
├─────────────┤
|
||||
│ L3 PostgreSQL│ ← 持久化存储,数据可靠性
|
||||
└─────────────┘
|
||||
```
|
||||
|
||||
**特性**:
|
||||
- ✅ **自动降级**:Redis 故障时自动切换到 Memory + PostgreSQL 模式
|
||||
- ✅ **灵活配置**:支持单级(Memory)、双级(Memory + PG)、完整三级
|
||||
- ✅ **TTL 管理**:会话默认 30 分钟过期,自动清理
|
||||
- ✅ **写穿透**:数据先写 L1,异步同步到 L2/L3
|
||||
|
||||
## 📊 数据库设计
|
||||
|
||||
系统使用 PostgreSQL 存储持久化数据:
|
||||
|
||||
### 核心表
|
||||
|
||||
| 表名 | 说明 | 关键字段 |
|
||||
|------|------|---------|
|
||||
| `users` | 用户账户 | `id (UUID)`, `username (UNIQUE)`, `password_hash (bcrypt)` |
|
||||
| `sessions` | 对话会话 | `id (UUID)`, `user_id (FK)`, `title`, `config (JSONB)` |
|
||||
| `messages` | 消息记录 | `id (BIGSERIAL)`, `session_id (FK)`, `role`, `content`, `tokens_used` |
|
||||
| `refresh_tokens` | 刷新令牌 | `token_hash (PK, SHA256)`, `user_id (FK)`, `expires_at` |
|
||||
|
||||
**关系**:`users 1:N sessions 1:N messages`,`users 1:N refresh_tokens`
|
||||
|
||||
**迁移管理**:使用嵌入式 SQL 文件(`backend/migrations/`),应用启动时自动执行。
|
||||
|
||||
## 🛡️ 安全特性
|
||||
|
||||
- 🔒 **密码安全**:bcrypt (cost=10) 哈希,自动生成盐值
|
||||
- 🔑 **Token 安全**:JWT HS256 签名,refresh_token SHA256 哈希存储
|
||||
- 🚫 **防重放攻击**:Refresh Token Rotation + 复用检测自动吊销
|
||||
- 🌐 **传输安全**:生产环境强制 HTTPS,开发环境 Vite proxy 同源代理
|
||||
- 🚦 **限流保护**:令牌桶算法(生产环境启用),防暴力破解
|
||||
- 🔍 **日志追踪**:全链路 Trace ID,请求/响应/错误统一记录
|
||||
|
||||
## 🌍 部署架构
|
||||
|
||||
```
|
||||
┌─────────────┐
|
||||
│ Nginx │ ← 反向代理(静态资源 + API + WebSocket)
|
||||
└──────┬──────┘
|
||||
│
|
||||
┌──────┴───────────────────┐
|
||||
│ Go Gateway 集群 │
|
||||
│ ├─ Gateway-1 │
|
||||
│ ├─ Gateway-2 │
|
||||
│ └─ Gateway-N │
|
||||
└───┬────────────┬─────────┘
|
||||
│ │
|
||||
┌───┴────┐ ┌───┴────────┐
|
||||
│ Redis │ │ PostgreSQL │
|
||||
└────────┘ └────────────┘
|
||||
│
|
||||
┌───┴────────────────────┐
|
||||
│ 外部 AI 服务 │
|
||||
│ ├─ DashScope (LLM) │
|
||||
│ ├─ MiMo (STT/TTS) │
|
||||
│ └─ Deepgram (可选) │
|
||||
└───────────────────────┘
|
||||
```
|
||||
|
||||
**跨域策略**:Nginx 统一反代前后端到同一域名,无跨域问题。
|
||||
|
||||
**水平扩展**:Gateway 无状态设计,会话状态存储在 Redis/PostgreSQL,支持多实例部署。
|
||||
|
||||
## 📖 文档
|
||||
|
||||
### 核心设计文档
|
||||
|
||||
| 文档 | 内容 |
|
||||
|------|------|
|
||||
| [01-架构设计](docs/01-架构设计.md) | 三层架构、技术栈、数据库设计、部署方案 |
|
||||
| [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
9
backend/.dockerignore
Normal file
@@ -0,0 +1,9 @@
|
||||
bin
|
||||
tmp
|
||||
.git
|
||||
.gitignore
|
||||
*.md
|
||||
.env*
|
||||
.vscode
|
||||
.idea
|
||||
vendor
|
||||
32
backend/.env.example
Normal file
32
backend/.env.example
Normal file
@@ -0,0 +1,32 @@
|
||||
# 运行环境
|
||||
# dev:本地开发环境(debug 日志、关闭限流、允许所有 CORS)
|
||||
# prod:生产环境(info 日志、启用限流、严格 CORS 白名单)
|
||||
# 本地开发保持 dev,生产部署会被 docker-compose.yml 覆盖为 prod
|
||||
APP_ENV=dev
|
||||
|
||||
# AI 服务 API Key
|
||||
CAMTALK_AI_STT_API_KEY=sk-your-stt-key
|
||||
CAMTALK_AI_LLM_API_KEY=sk-your-llm-key
|
||||
CAMTALK_AI_TTS_API_KEY=sk-your-tts-key
|
||||
|
||||
# JWT 认证
|
||||
CAMTALK_AUTH_JWT_SECRET=your-jwt-secret-here
|
||||
|
||||
# 三级存储配置
|
||||
# L2: Redis(热数据分布式会话层)
|
||||
CAMTALK_STORAGE_REDIS_ENABLED=true
|
||||
# 开发时填写远程服务器地址,部署时 docker-compose 会覆盖为容器内网地址
|
||||
CAMTALK_REDIS_ADDR=your-remote-server:6379
|
||||
CAMTALK_REDIS_PASSWORD=your-redis-password
|
||||
|
||||
# L3: PostgreSQL(冷数据持久化层)
|
||||
CAMTALK_STORAGE_PERSISTENCE_ENABLED=true
|
||||
CAMTALK_STORAGE_PERSISTENCE_DRIVER=postgres
|
||||
POSTGRES_USER=camtalk
|
||||
POSTGRES_PASSWORD=your-postgres-password
|
||||
# 开发时填写远程服务器地址,部署时 docker-compose 会覆盖为容器内网地址
|
||||
CAMTALK_STORAGE_DSN=postgres://camtalk:your-postgres-password@your-remote-server:5432/camtalk?sslmode=disable
|
||||
|
||||
# 可选覆盖(默认值见 config.yaml)
|
||||
# CAMTALK_SERVER_PORT=8080
|
||||
# CAMTALK_LOG_LEVEL=info
|
||||
15
backend/.gitignore
vendored
Normal file
15
backend/.gitignore
vendored
Normal file
@@ -0,0 +1,15 @@
|
||||
# 编译产物
|
||||
/server
|
||||
bin/
|
||||
|
||||
# 环境配置
|
||||
.env
|
||||
|
||||
# 临时文件
|
||||
tmp/
|
||||
|
||||
# IDE
|
||||
.idea/
|
||||
.vscode/
|
||||
*.swp
|
||||
*.swo
|
||||
38
backend/Dockerfile
Normal file
38
backend/Dockerfile
Normal 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"]
|
||||
@@ -1,40 +1,315 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"log"
|
||||
"context"
|
||||
"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/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/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() {
|
||||
r := gin.Default()
|
||||
// 加载配置(工作目录用于定位 .env 和 config.yaml)
|
||||
cfg, err := config.Load(".")
|
||||
if err != nil {
|
||||
panic("failed to load config: " + err.Error())
|
||||
}
|
||||
|
||||
// 初始化日志
|
||||
logger.Init(cfg.Log.Level, cfg.Log.Format)
|
||||
defer logger.Sync()
|
||||
|
||||
logger.Log.Infow("config loaded",
|
||||
"env", cfg.App.Env,
|
||||
"addr", cfg.Server.Addr(),
|
||||
)
|
||||
|
||||
// 初始化存储层(三级存储架构: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
|
||||
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 服务
|
||||
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,
|
||||
)
|
||||
|
||||
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" {
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
}
|
||||
|
||||
r := gin.New()
|
||||
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
|
||||
r.Use(trace.GinLogger()) // 第二层:记录请求
|
||||
r.Use(trace.GinRecovery()) // 第三层:panic 恢复
|
||||
|
||||
// REST API
|
||||
api := r.Group("/api")
|
||||
apiGroup := r.Group("/api")
|
||||
{
|
||||
api.GET("/health", healthHandler)
|
||||
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)
|
||||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter, userScenarioRepo))
|
||||
|
||||
log.Println("CamTalk gateway starting on :8080")
|
||||
if err := r.Run(":8080"); err != nil {
|
||||
log.Fatalf("failed to start server: %v", err)
|
||||
// HTTP Server
|
||||
srv := &http.Server{
|
||||
Addr: cfg.Server.Addr(),
|
||||
Handler: r,
|
||||
ReadTimeout: time.Duration(cfg.Server.ReadTimeout) * time.Second,
|
||||
WriteTimeout: time.Duration(cfg.Server.WriteTimeout) * time.Second,
|
||||
}
|
||||
|
||||
// Graceful shutdown
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||
defer stop()
|
||||
|
||||
go func() {
|
||||
logger.Log.Infow("server starting", "addr", srv.Addr)
|
||||
if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
logger.Log.Fatalw("listen failed", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
<-ctx.Done()
|
||||
logger.Log.Info("shutting down...")
|
||||
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(cfg.Server.ShutdownTimeout)*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if err := srv.Shutdown(shutdownCtx); err != nil {
|
||||
logger.Log.Errorw("shutdown error", "error", err)
|
||||
}
|
||||
logger.Log.Info("server stopped")
|
||||
}
|
||||
|
||||
// healthHandler 健康检查。
|
||||
func healthHandler(c *gin.Context) {
|
||||
c.JSON(200, gin.H{
|
||||
"status": "ok",
|
||||
"version": "0.1.0",
|
||||
"uptime": time.Since(startTime).String(),
|
||||
"active_sessions": 0, // TODO: 接入 Session Manager
|
||||
})
|
||||
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": version,
|
||||
"uptime_seconds": int(time.Since(startTime).Seconds()),
|
||||
"active_sessions": sessionMgr.ActiveCount(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
68
backend/config/config.dev.yaml
Normal file
68
backend/config/config.dev.yaml
Normal file
@@ -0,0 +1,68 @@
|
||||
# CamTalk 开发环境配置
|
||||
# 通过 APP_ENV=dev 加载此文件,覆盖 config.yaml 中的配置
|
||||
|
||||
server:
|
||||
host: "0.0.0.0"
|
||||
port: 8080
|
||||
heartbeat_interval: 30
|
||||
heartbeat_timeout: 60
|
||||
allowed_origins: [] # 开发环境允许所有来源
|
||||
|
||||
session:
|
||||
ttl: 30 # 开发环境会话较短,方便测试过期逻辑
|
||||
max_history: 20
|
||||
|
||||
ai:
|
||||
stt:
|
||||
provider: mimo # 与生产环境一致
|
||||
model: mimo-v2.5-asr
|
||||
endpoint: "https://api.xiaomimimo.com/v1"
|
||||
timeout: 10 # 开发环境超时较长,方便调试
|
||||
http_client_timeout: 30
|
||||
llm:
|
||||
provider: dashscope # 与生产环境一致
|
||||
model: qwen3-vl-plus
|
||||
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
timeout: 60 # 开发环境 LLM 超时较长
|
||||
http_client_timeout: 120
|
||||
tts:
|
||||
provider: mimo # 与生产环境一致
|
||||
model: mimo-v2.5-tts
|
||||
voice: mimo_default
|
||||
speed: 1.0
|
||||
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||
timeout: 10
|
||||
http_client_timeout: 30
|
||||
output_format: mp3
|
||||
sample_rate: 24000
|
||||
|
||||
storage:
|
||||
redis:
|
||||
enabled: true # 开发环境启用 Redis,测试三级存储
|
||||
persistence:
|
||||
enabled: true # 开发环境启用持久化
|
||||
|
||||
redis:
|
||||
addr: "localhost:6379" # 本地 Redis
|
||||
password: ""
|
||||
db: 0
|
||||
|
||||
auth:
|
||||
access_ttl: 120 # 开发环境 Access Token 2 小时,方便调试
|
||||
refresh_ttl: 10080 # 7 天
|
||||
|
||||
ratelimit:
|
||||
enabled: false # 开发环境关闭限流,方便测试
|
||||
query:
|
||||
capacity: 10
|
||||
rate: 0.2
|
||||
login:
|
||||
capacity: 5
|
||||
rate: 0.1
|
||||
register:
|
||||
capacity: 3
|
||||
rate: 0.05
|
||||
|
||||
log:
|
||||
level: debug # 开发环境 debug 日志
|
||||
format: console # 控制台格式,易读
|
||||
71
backend/config/config.prod.yaml
Normal file
71
backend/config/config.prod.yaml
Normal file
@@ -0,0 +1,71 @@
|
||||
# CamTalk 生产环境配置
|
||||
# 通过 APP_ENV=prod 加载此文件,覆盖 config.yaml 中的配置
|
||||
|
||||
server:
|
||||
host: "0.0.0.0"
|
||||
port: 8080
|
||||
read_timeout: 30
|
||||
write_timeout: 30
|
||||
shutdown_timeout: 15 # 生产环境优雅关闭时间稍长
|
||||
heartbeat_interval: 30
|
||||
heartbeat_timeout: 60
|
||||
|
||||
session:
|
||||
ttl: 60 # 生产环境会话 1 小时
|
||||
max_history: 20
|
||||
|
||||
ai:
|
||||
stt:
|
||||
provider: mimo # 生产环境推荐 MiMo,性价比高
|
||||
model: mimo-v2.5-asr
|
||||
endpoint: "https://api.xiaomimimo.com/v1"
|
||||
timeout: 5 # 生产环境严格超时控制
|
||||
http_client_timeout: 30
|
||||
llm:
|
||||
provider: dashscope # 生产环境推荐通义千问,稳定性好
|
||||
model: qwen3-vl-plus
|
||||
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
timeout: 30
|
||||
http_client_timeout: 60
|
||||
tts:
|
||||
provider: mimo # 生产环境推荐 MiMo TTS
|
||||
model: mimo-v2.5-tts
|
||||
voice: mimo_default
|
||||
speed: 1.0
|
||||
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||
timeout: 5
|
||||
http_client_timeout: 30
|
||||
output_format: mp3
|
||||
sample_rate: 24000
|
||||
|
||||
storage:
|
||||
redis:
|
||||
enabled: true # 生产环境必须启用 Redis
|
||||
persistence:
|
||||
enabled: true # 生产环境必须启用持久化
|
||||
driver: postgres
|
||||
|
||||
redis:
|
||||
addr: "redis:6379" # Docker Compose 内部服务名
|
||||
password: "" # 密码通过 CAMTALK_REDIS_PASSWORD 环境变量设置
|
||||
db: 0
|
||||
|
||||
auth:
|
||||
access_ttl: 120 # 生产环境 Access Token 2 小时
|
||||
refresh_ttl: 10080 # Refresh Token 7 天
|
||||
|
||||
ratelimit:
|
||||
enabled: true # 生产环境启用限流
|
||||
query:
|
||||
capacity: 10 # 允许突发 10 个请求
|
||||
rate: 0.2 # 每 5 秒恢复 1 个令牌
|
||||
login:
|
||||
capacity: 5 # 防暴力破解
|
||||
rate: 0.1 # 每 10 秒恢复 1 次
|
||||
register:
|
||||
capacity: 3 # 防批量注册
|
||||
rate: 0.05 # 每 20 秒恢复 1 次
|
||||
|
||||
log:
|
||||
level: info # 生产环境 info 级别
|
||||
format: json # JSON 格式,便于日志收集和分析
|
||||
80
backend/config/config.yaml
Normal file
80
backend/config/config.yaml
Normal file
@@ -0,0 +1,80 @@
|
||||
# CamTalk 后端配置
|
||||
|
||||
app:
|
||||
env: dev # dev / prod,可通过 APP_ENV 环境变量覆盖
|
||||
|
||||
server:
|
||||
host: "0.0.0.0"
|
||||
port: 8080
|
||||
read_timeout: 30 # 秒
|
||||
write_timeout: 30 # 秒
|
||||
shutdown_timeout: 10 # 优雅关闭超时(秒)
|
||||
heartbeat_interval: 30 # 心跳检查间隔(秒)
|
||||
heartbeat_timeout: 60 # 心跳超时断开(秒)
|
||||
allowed_origins: [] # CORS 白名单,空=允许所有
|
||||
|
||||
session:
|
||||
ttl: 30 # 会话过期时间(分钟)
|
||||
max_history: 20 # 对话历史上限(条)
|
||||
|
||||
ai:
|
||||
stt:
|
||||
provider: mimo # mimo / deepgram
|
||||
model: mimo-v2.5-asr
|
||||
endpoint: "https://api.xiaomimimo.com/v1"
|
||||
timeout: 5 # STT 请求超时(秒)
|
||||
http_client_timeout: 30 # HTTP 客户端超时(秒)
|
||||
llm:
|
||||
provider: dashscope # dashscope / openai
|
||||
model: qwen3-vl-plus
|
||||
endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
timeout: 30 # LLM 请求超时(秒)
|
||||
http_client_timeout: 60 # HTTP 客户端超时(秒)
|
||||
tts:
|
||||
provider: mimo # mimo / openai
|
||||
model: mimo-v2.5-tts
|
||||
voice: mimo_default
|
||||
speed: 1.0
|
||||
endpoint: "https://token-plan-cn.xiaomimimo.com/v1"
|
||||
timeout: 5 # TTS 请求超时(秒)
|
||||
http_client_timeout: 30 # HTTP 客户端超时(秒)
|
||||
output_format: mp3 # 输出格式:mp3 / wav
|
||||
sample_rate: 24000 # 输出采样率
|
||||
|
||||
storage:
|
||||
# 三级存储架构:L1 内存 → L2 Redis → L3 PostgreSQL
|
||||
redis:
|
||||
enabled: true # 是否启用 Redis(L2 热数据层)
|
||||
persistence:
|
||||
enabled: true # 是否启用持久化(L3 冷数据层)
|
||||
driver: postgres # postgres
|
||||
# dsn 通过环境变量 CAMTALK_STORAGE_DSN 设置
|
||||
|
||||
redis:
|
||||
addr: "localhost:6379"
|
||||
password: ""
|
||||
db: 0
|
||||
|
||||
auth:
|
||||
# jwt_secret 通过环境变量 CAMTALK_AUTH_JWT_SECRET 设置
|
||||
access_ttl: 120 # Access Token 过期时间(分钟)
|
||||
refresh_ttl: 10080 # Refresh Token 过期时间(分钟),7 天
|
||||
|
||||
ratelimit:
|
||||
enabled: false # 是否启用限流
|
||||
# WebSocket query 消息限流(核心,控制 AI 成本)
|
||||
query:
|
||||
capacity: 10 # 突发容量:允许连续发 10 个 query
|
||||
rate: 0.2 # 填充速率:每 5 秒补充 1 个令牌
|
||||
# REST API 登录限流(防暴力破解)
|
||||
login:
|
||||
capacity: 5 # 突发容量:允许连续 5 次登录尝试
|
||||
rate: 0.1 # 填充速率:每 10 秒补充 1 次
|
||||
# REST API 注册限流
|
||||
register:
|
||||
capacity: 3 # 突发容量:允许连续 3 次注册
|
||||
rate: 0.05 # 填充速率:每 20 秒补充 1 次
|
||||
|
||||
log:
|
||||
level: info # debug / info / warn / error
|
||||
format: console # console / json
|
||||
@@ -1,38 +1,85 @@
|
||||
module github.com/hhs/camtalk
|
||||
|
||||
go 1.23
|
||||
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/cloudwego/base64x v0.1.4 // indirect
|
||||
github.com/cloudwego/iasm v0.2.0 // 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.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
|
||||
github.com/go-playground/locales v0.14.1 // indirect
|
||||
github.com/go-playground/universal-translator v0.18.1 // indirect
|
||||
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/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.7 // 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/pelletier/go-toml/v2 v2.2.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
|
||||
github.com/spf13/pflag v1.0.10 // indirect
|
||||
github.com/stretchr/objx v0.5.2 // indirect
|
||||
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
|
||||
golang.org/x/arch v0.8.0 // indirect
|
||||
golang.org/x/crypto v0.23.0 // 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.11.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20230713183714-613f0c0eb8a1 // indirect
|
||||
golang.org/x/net v0.25.0 // indirect
|
||||
golang.org/x/sys v0.20.0 // indirect
|
||||
golang.org/x/text v0.15.0 // indirect
|
||||
golang.org/x/sync v0.17.0 // indirect
|
||||
golang.org/x/sys v0.30.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
|
||||
)
|
||||
|
||||
223
backend/go.sum
223
backend/go.sum
@@ -1,20 +1,60 @@
|
||||
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/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/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/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.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=
|
||||
@@ -23,71 +63,186 @@ github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJn
|
||||
github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY=
|
||||
github.com/go-playground/validator/v10 v10.20.0 h1:K9ISHbSaI0lyB2eWMPJo+kOS/FBExVwjEviJTixqxL8=
|
||||
github.com/go-playground/validator/v10 v10.20.0/go.mod h1:dbuPbCMFw/DrkbEynArYaCwl3amGuJotoKCe95atGMM=
|
||||
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
|
||||
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/google/go-cmp v0.5.5 h1:Khx7svrCpmxxtHBq5j2mp/xVjsi8hQMfNLvJFAlrGgU=
|
||||
github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
|
||||
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/klauspost/cpuid/v2 v2.2.7 h1:ZWSB3igEs+d0qvnxR/ZBzXVmxkgt8DdzP6m9pfuVLDM=
|
||||
github.com/klauspost/cpuid/v2 v2.2.7/go.mod h1:Lcz8mBdAVJIBVzewtcLocK12l3Y+JytZYpaMropDUws=
|
||||
github.com/knz/go-libedit v1.10.1/go.mod h1:MZTVkCWyz0oBc7JOWP3wNAzd002ZbM/5hgShxwh4x8M=
|
||||
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/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/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6Wq+LM=
|
||||
github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
|
||||
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=
|
||||
github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg=
|
||||
github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY=
|
||||
github.com/spf13/cast v1.10.0/go.mod h1:jNfB8QC9IA6ZuY2ZjDp0KtFO2LZZlg4S/7bzP6qqeHo=
|
||||
github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk=
|
||||
github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
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.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
|
||||
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
|
||||
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=
|
||||
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU=
|
||||
github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI=
|
||||
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=
|
||||
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=
|
||||
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.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/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
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.20.0 h1:Od9JTbYCk261bKm4M/mw7AklTlFYIa0bIp9BgSm1S8Y=
|
||||
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/text v0.15.0 h1:h1V/4gjBv8v9cjcR6+AR5+/cIYK5N/WAgiv4xlsEtAk=
|
||||
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
|
||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543 h1:E7g+9GITq07hpfrRu66IVDexMakfv52eLZ2CXBWiKr4=
|
||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
|
||||
golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
|
||||
golang.org/x/term v0.28.0 h1:/Ts8HFuMR2E6IP/jlo7QVLZHggjKQbhu/7H0LJFr3Gg=
|
||||
golang.org/x/term v0.28.0/go.mod h1:Sw/lC2IAUZ92udQNf3WodGtn4k/XoLyZoh8v/8uiwek=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
|
||||
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
|
||||
google.golang.org/protobuf v1.34.1 h1:9ddQBjfCyZPOHPUiPxpYESBLc+T8P3E+Vo4IbKZgFWg=
|
||||
google.golang.org/protobuf v1.34.1/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
|
||||
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
|
||||
gopkg.in/fsnotify.v1 v1.4.7/go.mod h1:Tz8NjZHkW78fSQdbUxIjBTcgA1z1m8ZHf0WmKUhAMys=
|
||||
gopkg.in/tomb.v1 v1.0.0-20141024135613-dd632973f1e7/go.mod h1:dt/ZhP58zS4L8KSrWDmTeBkI65Dw0HsyUHuEVlX15mw=
|
||||
gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
nullprogram.com/x/optparse v1.0.0/go.mod h1:KdyPE+Igbe0jQUrVfMqDMeJQIJZEuyV7pjYmp6pbG50=
|
||||
rsc.io/pdf v0.1.1/go.mod h1:n8OzWcQ6Sp37PL01nO98y4iUCRdTGarVfzxY20ICaU4=
|
||||
|
||||
38
backend/internal/ai/llm/llm.go
Normal file
38
backend/internal/ai/llm/llm.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
// Service 多模态大模型服务契约。
|
||||
type Service interface {
|
||||
// ChatStream 流式推理,返回增量文本的 channel。
|
||||
// 调用方必须消费 channel 直到 Done=true,否则需 cancel ctx 以释放连接。
|
||||
ChatStream(ctx context.Context, req Request) (<-chan Chunk, error)
|
||||
}
|
||||
|
||||
// Request 推理请求。
|
||||
type Request struct {
|
||||
Image []byte // JPEG 图片(已从 Base64 解码)
|
||||
Text string // 用户语音识别后的文本
|
||||
History []models.Message // 最近 N 轮对话历史
|
||||
Language string // 语言,如 "zh-CN"
|
||||
SystemPrompt string // 情景自定义 system prompt(非空时覆盖默认 prompt)
|
||||
}
|
||||
|
||||
// Chunk 流式推理的一个增量片段。
|
||||
type Chunk struct {
|
||||
Delta string // 增量文本
|
||||
Done bool // 是否结束
|
||||
TokensUsed *TokenUsage // 仅 Done=true 时有值
|
||||
Model string // 实际使用的模型名
|
||||
}
|
||||
|
||||
// TokenUsage 用量统计。
|
||||
type TokenUsage struct {
|
||||
Prompt int
|
||||
Completion int
|
||||
Total int
|
||||
}
|
||||
41
backend/internal/ai/llm/prompt.go
Normal file
41
backend/internal/ai/llm/prompt.go
Normal file
@@ -0,0 +1,41 @@
|
||||
package llm
|
||||
|
||||
import "strings"
|
||||
|
||||
// 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句话以内,除非用户要求详细说明。")
|
||||
} else {
|
||||
prompt.WriteString("You are a visual assistant. The user sees a scene through their camera and asks questions by voice. Answer concisely and naturally. If describing visual content, start with 'I see...'. Keep answers to 3-5 sentences unless the user asks for detail.")
|
||||
}
|
||||
|
||||
if detailLevel == "high" {
|
||||
if isChinese {
|
||||
prompt.WriteString("请提供更详细的视觉描述,包括颜色、位置、数量等细节。")
|
||||
} else {
|
||||
prompt.WriteString(" Provide detailed visual descriptions including colors, positions, quantities, and other details.")
|
||||
}
|
||||
}
|
||||
|
||||
return prompt.String()
|
||||
}
|
||||
148
backend/internal/ai/llm/scenarios.go
Normal file
148
backend/internal/ai/llm/scenarios.go
Normal file
@@ -0,0 +1,148 @@
|
||||
package llm
|
||||
|
||||
import "strings"
|
||||
|
||||
// scenarioPrompt 定义单个情景的多语言 system prompt 和首句引导。
|
||||
type scenarioPrompt struct {
|
||||
ZH string
|
||||
EN string
|
||||
JA string
|
||||
GreetingZH string // 首句引导(中文)
|
||||
GreetingEN string // 首句引导(英文)
|
||||
GreetingJA string // 首句引导(日文)
|
||||
}
|
||||
|
||||
// scenarioPrompts 预置情景 → prompt 映射表。
|
||||
// key 为情景 ID(与前端 Scenario.id 对齐)。
|
||||
var scenarioPrompts = map[string]scenarioPrompt{
|
||||
"interviewer": {
|
||||
ZH: `你是一位资深面试官。你通过摄像头观察面试者,并根据他们的背景和表现提出面试问题。
|
||||
|
||||
【角色定位】
|
||||
- 你是面试官,不是助手或顾问
|
||||
- 你的目标是评估候选人的能力
|
||||
- 保持专业、客观、礼貌
|
||||
|
||||
【交互规则】
|
||||
1. 每次只问一个问题,等用户回答后再追问
|
||||
2. 问题要有层次:自我介绍 → 专业问题 → 情景题 → 反问环节
|
||||
3. 对用户的回答给出简短点评(优点+不足),然后追问
|
||||
4. 如果摄像头能看到用户的环境,可以结合环境提出相关话题
|
||||
5. 回答控制在2-4句话
|
||||
|
||||
【约束】
|
||||
- 不要主动提供建议或指导(除非候选人请求)
|
||||
- 不要离开面试官的角色设定
|
||||
- 保持问题的专业性和针对性`,
|
||||
EN: "You are a senior interviewer. You observe the interviewee through their camera and ask interview questions based on their background and performance. Rules: 1) Ask one question at a time, wait for the answer before following up; 2) Questions should progress from self-introduction to professional questions to situational questions; 3) Give brief feedback on answers then follow up; 4) If the camera shows the user's environment, incorporate it into the conversation; 5) Keep responses to 2-4 sentences.",
|
||||
JA: "あなたはベテラン面接官です。カメラで面接者を見て、バックグラウンドと実績に基づいて面接質問をします。ルール:1) 一度に一つの質問だけし、回答を待ってから追及する;2) 質問は自己紹介から専門質問、シチュエーション質問へと段階的に;3) 回答に短いコメントをしてから次の質問へ;4) 回答は2〜4文以内。",
|
||||
GreetingZH: "你好!我是今天的面试官。让我们先从自我介绍开始,请简单介绍一下你自己和你应聘的岗位。",
|
||||
GreetingEN: "Hello! I'm your interviewer today. Let's start with a self-introduction. Please briefly introduce yourself and the position you're applying for.",
|
||||
GreetingJA: "こんにちは!本日の面接官です。まず自己紹介から始めましょう。あなた自身と応募職種について簡単に教えてください。",
|
||||
},
|
||||
"english_teacher": {
|
||||
ZH: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material (e.g., 'I can see a bookshelf behind you. What's your favorite book?'); 6) Keep responses to 3-5 sentences.",
|
||||
EN: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
||||
JA: "You are a friendly and patient English tutor. Speak in English with the user. Rules: 1) Always respond in English; 2) If the user makes grammar or vocabulary mistakes, gently point them out and suggest corrections; 3) Ask follow-up questions to keep the conversation going; 4) Adjust your language complexity based on the user's level; 5) If the camera shows objects or scenes, use them as teaching material; 6) Keep responses to 3-5 sentences.",
|
||||
GreetingZH: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||
GreetingEN: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||
GreetingJA: "Hi! I'm your English tutor. Let's practice English together! What would you like to talk about today?",
|
||||
},
|
||||
"debate": {
|
||||
ZH: `你是一位辩论赛对手。用户提出一个观点,你需要站在反方进行反驳。
|
||||
|
||||
【角色定位】
|
||||
- 你是辩论对手,不是评委或顾问
|
||||
- 你的目标是通过逻辑论证反驳对方观点
|
||||
- 保持理性、严谨、尊重对手
|
||||
|
||||
【交互规则】
|
||||
1. 逻辑严密,用事实和论据反驳,不要人身攻击
|
||||
2. 每次提出1-2个核心反驳点,并给出简要论据
|
||||
3. 如果用户论证有力,承认其合理性但仍要寻找突破口
|
||||
4. 适时提出反问,引导用户深入思考
|
||||
5. 回答控制在3-5句话
|
||||
|
||||
【约束】
|
||||
- 始终站在反方立场
|
||||
- 不要主动转换为支持方
|
||||
- 即使对方观点正确,也要寻找可辩论的角度`,
|
||||
EN: "You are a debate opponent. The user presents a viewpoint, and you argue against it. Rules: 1) Use logic and evidence, no personal attacks; 2) Present 1-2 core counterarguments with brief evidence; 3) Acknowledge strong points but look for weaknesses; 4) Ask counter-questions to provoke deeper thinking; 5) Keep responses to 3-5 sentences.",
|
||||
JA: "あなたはディベートの相手です。ユーザーが提示した观点に対して反論します。ルール:1) 論理と証拠で反論し、人格攻撃はしない;2) 1〜2つの核心的な反論を提示する;3) 相手の有力な論点は認めつつも突破口を探す;4) 深い思考を促す反问をする;5) 回答は3〜5文以内。",
|
||||
GreetingZH: "你好!我是你的辩论对手。请提出一个你坚信的观点,我会站在反方立场与你辩论,帮你锻炼逻辑思维。",
|
||||
GreetingEN: "Hello! I'm your debate opponent. Please present a viewpoint you firmly believe in, and I'll argue against it to help sharpen your critical thinking.",
|
||||
GreetingJA: "こんにちは!あなたのディベート相手です。あなたが信じる观点を提示してください。反対の立場から論じて、論理的思考を鍛えます。",
|
||||
},
|
||||
"interpreter": {
|
||||
ZH: "你是一名同声翻译员。将用户说的话实时翻译为目标语言。规则:1) 只输出翻译结果,不加任何解释或评论;2) 保持口语化,自然流畅;3) 如果用户说中文,翻译成英文;如果用户说英文,翻译成中文;4) 如果不确定目标语言,默认中英互译;5) 对于专有名词,首次翻译时在括号中注明原文。",
|
||||
EN: "You are a simultaneous interpreter. Translate what the user says in real-time. Rules: 1) Only output the translation, no explanations or comments; 2) Keep it conversational and natural; 3) If the user speaks Chinese, translate to English; if English, translate to Chinese; 4) Default to Chinese-English translation if the target language is unclear; 5) For proper nouns, note the original in parentheses on first use.",
|
||||
JA: "あなたは同時通訳者です。ユーザーの発言をリアルタイムで翻訳します。ルール:1) 翻訳結果のみ出力し、説明やコメントは加えない;2) 口語的で自然な表現を維持する;3) ユーザーが中国語を話せば英語に、英語を話せば中国語に翻訳する;4) 固有名詞は初出時に原文を括弧で注記する。",
|
||||
GreetingZH: "我是你的同声翻译。请开始说话,我会实时将中文翻译成英文,或将英文翻译成中文。",
|
||||
GreetingEN: "I'm your simultaneous interpreter. Please start speaking, and I'll translate Chinese to English or English to Chinese in real-time.",
|
||||
GreetingJA: "私はあなたの同時通訳者です。お話しください。中国語を英語に、または英語を中国語にリアルタイムで翻訳します。",
|
||||
},
|
||||
}
|
||||
|
||||
// GetScenarioPrompt 根据情景 ID 和语言获取对应的 system prompt。
|
||||
// 支持系统预置情景和用户自建情景。
|
||||
// customScenarios: 用户自建情景映射表(scenarioID → prompt),可为 nil
|
||||
// 返回空字符串表示无此情景(使用默认 prompt)。
|
||||
func GetScenarioPrompt(scenarioID, language string, customScenarios map[string]string) string {
|
||||
if scenarioID == "" || scenarioID == "free_chat" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// 1. 优先查找系统预置情景
|
||||
if p, ok := scenarioPrompts[scenarioID]; ok {
|
||||
switch {
|
||||
case strings.HasPrefix(language, "zh"):
|
||||
return p.ZH
|
||||
case strings.HasPrefix(language, "ja"):
|
||||
return p.JA
|
||||
default:
|
||||
return p.EN
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 查找用户自建情景
|
||||
if customScenarios != nil {
|
||||
if customPrompt, ok := customScenarios[scenarioID]; ok {
|
||||
return customPrompt
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 默认空字符串
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetScenarioGreeting 根据情景 ID 和语言获取对应的首句引导。
|
||||
// 支持系统预置情景和用户自建情景。
|
||||
// customGreetings: 用户自建情景的首句引导映射表(scenarioID → greeting),可为 nil
|
||||
// 返回空字符串表示无此情景或不需要引导(自由对话)。
|
||||
func GetScenarioGreeting(scenarioID, language string, customGreetings map[string]string) string {
|
||||
if scenarioID == "" || scenarioID == "free_chat" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// 1. 优先查找系统预置情景
|
||||
if p, ok := scenarioPrompts[scenarioID]; ok {
|
||||
switch {
|
||||
case strings.HasPrefix(language, "zh"):
|
||||
return p.GreetingZH
|
||||
case strings.HasPrefix(language, "ja"):
|
||||
return p.GreetingJA
|
||||
default:
|
||||
return p.GreetingEN
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 查找用户自建情景
|
||||
if customGreetings != nil {
|
||||
if customGreeting, ok := customGreetings[scenarioID]; ok {
|
||||
return customGreeting
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 默认空字符串
|
||||
return ""
|
||||
}
|
||||
142
backend/internal/ai/stt/deepgram.go
Normal file
142
backend/internal/ai/stt/deepgram.go
Normal file
@@ -0,0 +1,142 @@
|
||||
package stt
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// DeepgramService 基于 Deepgram WebSocket API 的语音识别实现。
|
||||
type DeepgramService struct {
|
||||
apiKey string
|
||||
model string
|
||||
endpoint string
|
||||
timeout time.Duration
|
||||
logger *zap.SugaredLogger
|
||||
}
|
||||
|
||||
// NewDeepgramService 创建 Deepgram STT 服务。
|
||||
// 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,
|
||||
}
|
||||
}
|
||||
|
||||
// deepgramResponse Deepgram WebSocket 响应。
|
||||
type deepgramResponse struct {
|
||||
Channel struct {
|
||||
Alternatives []struct {
|
||||
Transcript string `json:"transcript"`
|
||||
Confidence float64 `json:"confidence"`
|
||||
} `json:"alternatives"`
|
||||
} `json:"channel"`
|
||||
IsFinal bool `json:"is_final"`
|
||||
}
|
||||
|
||||
// Recognize 实现 stt.Service。通过 WebSocket 发送音频到 Deepgram,返回最终识别文本。
|
||||
func (d *DeepgramService) Recognize(ctx context.Context, audio []byte, opts Options) (string, error) {
|
||||
if len(audio) == 0 {
|
||||
return "", fmt.Errorf("stt: empty audio")
|
||||
}
|
||||
|
||||
// 构建 WebSocket URL,附带查询参数
|
||||
wsURL := d.buildURL(opts)
|
||||
|
||||
// 总超时
|
||||
ctx, cancel := context.WithTimeout(ctx, d.timeout)
|
||||
defer cancel()
|
||||
|
||||
// 建立 WebSocket 连接
|
||||
conn, _, err := websocket.DefaultDialer.DialContext(ctx, wsURL, http.Header{
|
||||
"Authorization": []string{"Token " + d.apiKey},
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("stt: connect deepgram: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// 发送音频数据(一次性)
|
||||
if err := conn.WriteMessage(websocket.BinaryMessage, audio); err != nil {
|
||||
return "", fmt.Errorf("stt: send audio: %w", err)
|
||||
}
|
||||
|
||||
// 发送 Close 消息通知服务端音频已发送完毕
|
||||
closeMsg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")
|
||||
_ = conn.WriteMessage(websocket.CloseMessage, closeMsg)
|
||||
|
||||
// 读取识别结果
|
||||
var transcript strings.Builder
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
// Close 帧是正常的结束信号
|
||||
if websocket.IsCloseError(err, websocket.CloseNormalClosure) {
|
||||
break
|
||||
}
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseNormalClosure) {
|
||||
break
|
||||
}
|
||||
return "", fmt.Errorf("stt: read response: %w", err)
|
||||
}
|
||||
|
||||
var resp deepgramResponse
|
||||
if err := json.Unmarshal(message, &resp); err != nil {
|
||||
d.logger.Warnw("stt: unmarshal response failed", "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
// 只累积 final 结果,跳过中间结果
|
||||
if resp.IsFinal && len(resp.Channel.Alternatives) > 0 {
|
||||
text := strings.TrimSpace(resp.Channel.Alternatives[0].Transcript)
|
||||
if text != "" {
|
||||
transcript.WriteString(text)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return strings.TrimSpace(transcript.String()), nil
|
||||
}
|
||||
|
||||
// buildURL 构建 Deepgram WebSocket URL,包含音频格式参数。
|
||||
func (d *DeepgramService) buildURL(opts Options) string {
|
||||
u, _ := url.Parse(d.endpoint)
|
||||
|
||||
encoding := opts.Encoding
|
||||
if encoding == "" {
|
||||
encoding = "pcm_s16le"
|
||||
}
|
||||
sampleRate := opts.SampleRate
|
||||
if sampleRate == 0 {
|
||||
sampleRate = 16000
|
||||
}
|
||||
language := opts.Language
|
||||
if language == "" {
|
||||
language = "zh-CN"
|
||||
}
|
||||
|
||||
q := u.Query()
|
||||
q.Set("encoding", encoding)
|
||||
q.Set("sample_rate", fmt.Sprintf("%d", sampleRate))
|
||||
q.Set("language", language)
|
||||
q.Set("model", d.model)
|
||||
q.Set("punctuate", "true")
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
return u.String()
|
||||
}
|
||||
234
backend/internal/ai/stt/mimo.go
Normal file
234
backend/internal/ai/stt/mimo.go
Normal 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
|
||||
}
|
||||
258
backend/internal/ai/stt/mimo_test.go
Normal file
258
backend/internal/ai/stt/mimo_test.go
Normal 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()
|
||||
|
||||
// 发送一个简单的有效 WAV(44 字节头 + 少量 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)
|
||||
}
|
||||
// 应该是有效 WAV(RIFF 头)
|
||||
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
|
||||
}
|
||||
16
backend/internal/ai/stt/stt.go
Normal file
16
backend/internal/ai/stt/stt.go
Normal file
@@ -0,0 +1,16 @@
|
||||
package stt
|
||||
|
||||
import "context"
|
||||
|
||||
// Service 语音识别服务契约。
|
||||
type Service interface {
|
||||
// Recognize 识别一段完整音频,返回最终文本。
|
||||
Recognize(ctx context.Context, audio []byte, opts Options) (string, error)
|
||||
}
|
||||
|
||||
// Options 语音识别参数。
|
||||
type Options struct {
|
||||
Encoding string // 音频编码,如 "pcm_s16le"
|
||||
SampleRate int // 采样率,如 16000
|
||||
Language string // 语言,如 "zh-CN"
|
||||
}
|
||||
214
backend/internal/ai/tts/mimo.go
Normal file
214
backend/internal/ai/tts/mimo.go
Normal 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
|
||||
}
|
||||
414
backend/internal/ai/tts/mimo_test.go
Normal file
414
backend/internal/ai/tts/mimo_test.go
Normal 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 次 API(3 个句子)
|
||||
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")
|
||||
}
|
||||
}
|
||||
156
backend/internal/ai/tts/openai.go
Normal file
156
backend/internal/ai/tts/openai.go
Normal file
@@ -0,0 +1,156 @@
|
||||
package tts
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// OpenAIService 基于 OpenAI TTS API 的语音合成实现。
|
||||
type OpenAIService struct {
|
||||
apiKey string
|
||||
model string
|
||||
voice string
|
||||
speed float64
|
||||
endpoint string
|
||||
timeout time.Duration
|
||||
logger *zap.SugaredLogger
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewOpenAIService 创建 OpenAI TTS 服务。
|
||||
// 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
|
||||
}
|
||||
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,
|
||||
voice: voice,
|
||||
speed: speed,
|
||||
endpoint: endpoint,
|
||||
timeout: timeout,
|
||||
logger: logger,
|
||||
client: &http.Client{Timeout: httpClientTimeout},
|
||||
}
|
||||
}
|
||||
|
||||
// ttsRequest OpenAI TTS API 请求。
|
||||
type ttsRequest struct {
|
||||
Model string `json:"model"`
|
||||
Input string `json:"input"`
|
||||
Voice string `json:"voice"`
|
||||
ResponseFormat string `json:"response_format"`
|
||||
Speed float64 `json:"speed"`
|
||||
}
|
||||
|
||||
// SynthesizeStream 实现 tts.Service。从 textStream 读取句子,逐句调用 OpenAI TTS API。
|
||||
func (o *OpenAIService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error) {
|
||||
voice := opts.Voice
|
||||
if voice == "" {
|
||||
voice = o.voice
|
||||
}
|
||||
speed := opts.Speed
|
||||
if speed <= 0 {
|
||||
speed = o.speed
|
||||
}
|
||||
|
||||
ch := make(chan Chunk, 4)
|
||||
go func() {
|
||||
defer close(ch)
|
||||
|
||||
for text := range textStream {
|
||||
if text == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
audio, err := o.synthesize(ctx, text, voice, speed)
|
||||
if err != nil {
|
||||
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: 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 调用 OpenAI TTS API 合成单个句子。
|
||||
func (o *OpenAIService) synthesize(ctx context.Context, text, voice string, speed float64) ([]byte, error) {
|
||||
// 单句超时
|
||||
ctx, cancel := context.WithTimeout(ctx, o.timeout)
|
||||
defer cancel()
|
||||
|
||||
body := ttsRequest{
|
||||
Model: o.model,
|
||||
Input: text,
|
||||
Voice: voice,
|
||||
ResponseFormat: "mp3",
|
||||
Speed: speed,
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tts: marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint+"/audio/speech", bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tts: create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+o.apiKey)
|
||||
|
||||
resp, err := o.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tts: send request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
errBody, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("tts: api error (status %d): %s", resp.StatusCode, string(errBody))
|
||||
}
|
||||
|
||||
// 读取整个 MP3 响应
|
||||
audio, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tts: read response: %w", err)
|
||||
}
|
||||
|
||||
return audio, nil
|
||||
}
|
||||
309
backend/internal/ai/tts/openai_test.go
Normal file
309
backend/internal/ai/tts/openai_test.go
Normal file
@@ -0,0 +1,309 @@
|
||||
package tts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// mockTTSServer 创建模拟 OpenAI TTS API 的 HTTP 服务器。
|
||||
func mockTTSServer(t *testing.T, handler http.HandlerFunc) *httptest.Server {
|
||||
t.Helper()
|
||||
return httptest.NewServer(handler)
|
||||
}
|
||||
|
||||
// sendSentences 向 channel 发送句子并关闭。
|
||||
func sendSentences(sentences ...string) <-chan string {
|
||||
ch := make(chan string, len(sentences))
|
||||
for _, s := range sentences {
|
||||
ch <- s
|
||||
}
|
||||
close(ch)
|
||||
return ch
|
||||
}
|
||||
|
||||
func TestOpenAIService_SynthesizeStream_Success(t *testing.T) {
|
||||
var callCount int32
|
||||
srv := mockTTSServer(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, "/audio/speech") {
|
||||
t.Errorf("path = %s, should contain /audio/speech", r.URL.Path)
|
||||
}
|
||||
auth := r.Header.Get("Authorization")
|
||||
if auth != "Bearer test-key" {
|
||||
t.Errorf("Authorization = %q, want %q", auth, "Bearer test-key")
|
||||
}
|
||||
|
||||
// 验证请求体
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
if !strings.Contains(string(body), "tts-1") {
|
||||
t.Errorf("request body should contain model tts-1")
|
||||
}
|
||||
|
||||
// 返回假 MP3 数据
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "fake-mp3-data")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
|
||||
|
||||
textStream := sendSentences("你好", "世界", "!")
|
||||
|
||||
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{
|
||||
Voice: "alloy", Speed: 1.0, 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 次 API(3 个句子)
|
||||
if atomic.LoadInt32(&callCount) != 3 {
|
||||
t.Errorf("API called %d times, want 3", callCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIService_SynthesizeStream_APIError(t *testing.T) {
|
||||
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, "internal error")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 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 TestOpenAIService_SynthesizeStream_Timeout(t *testing.T) {
|
||||
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(3 * time.Second)
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "late-mp3")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
// 1 秒超时
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 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 TestOpenAIService_SynthesizeStream_EmptyText(t *testing.T) {
|
||||
var callCount int32
|
||||
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
atomic.AddInt32(&callCount, 1)
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "mp3")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 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 TestOpenAIService_SynthesizeStream_ContextCancelled(t *testing.T) {
|
||||
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "mp3")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 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 TestOpenAIService_SynthesizeStream_PartialFailure(t *testing.T) {
|
||||
var callCount int32
|
||||
srv := mockTTSServer(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
|
||||
}
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "mp3-%d", n)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 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 TestOpenAIService_SynthesizeStream_CustomVoice(t *testing.T) {
|
||||
srv := mockTTSServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
if !strings.Contains(string(body), "nova") {
|
||||
t.Errorf("request body should contain voice 'nova', got: %s", string(body))
|
||||
}
|
||||
w.Header().Set("Content-Type", "audio/mpeg")
|
||||
fmt.Fprintf(w, "mp3")
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
svc := NewOpenAIService("test-key", "tts-1", "alloy", srv.URL, 1.0, 5, 30, zap.NewNop().Sugar())
|
||||
|
||||
textStream := sendSentences("你好")
|
||||
|
||||
ch, err := svc.SynthesizeStream(context.Background(), textStream, Options{Voice: "nova"})
|
||||
if err != nil {
|
||||
t.Fatalf("SynthesizeStream() error: %v", err)
|
||||
}
|
||||
|
||||
for range ch {
|
||||
}
|
||||
}
|
||||
26
backend/internal/ai/tts/tts.go
Normal file
26
backend/internal/ai/tts/tts.go
Normal file
@@ -0,0 +1,26 @@
|
||||
package tts
|
||||
|
||||
import "context"
|
||||
|
||||
// Service 语音合成服务契约。
|
||||
type Service interface {
|
||||
// SynthesizeStream 流式合成。
|
||||
// textStream 接收句子级文本(由 Orchestrator 的句子切分器产出),
|
||||
// 返回的 channel 持续输出 MP3 音频 chunk。
|
||||
SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error)
|
||||
}
|
||||
|
||||
// Options 合成参数。
|
||||
type Options struct {
|
||||
Voice string // "alloy" | "nova" | "shimmer" 等
|
||||
Speed float64 // 1.0 为正常语速
|
||||
OutputFmt string // "mp3" — 固定使用 MP3
|
||||
SampleRate int // 24000
|
||||
}
|
||||
|
||||
// Chunk 一个音频片段。
|
||||
type Chunk struct {
|
||||
Audio []byte // MP3 音频数据(未 Base64 编码)
|
||||
IsLast bool // 当前句子是否为最后一片(每句结束时为 true)
|
||||
Final bool // 整轮 TTS 是否结束(所有句子合成完毕后为 true,此时 Audio 为 nil)
|
||||
}
|
||||
244
backend/internal/api/auth.go
Normal file
244
backend/internal/api/auth.go
Normal 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",
|
||||
})
|
||||
}
|
||||
}
|
||||
324
backend/internal/api/auth_test.go
Normal file
324
backend/internal/api/auth_test.go
Normal 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")
|
||||
}
|
||||
358
backend/internal/api/conversation.go
Normal file
358
backend/internal/api/conversation.go
Normal 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
|
||||
}
|
||||
575
backend/internal/api/conversation_test.go
Normal file
575
backend/internal/api/conversation_test.go
Normal 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)
|
||||
}
|
||||
107
backend/internal/api/session.go
Normal file
107
backend/internal/api/session.go
Normal file
@@ -0,0 +1,107 @@
|
||||
// Package api 提供 REST API 处理函数。
|
||||
package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// SessionHandler 提供会话相关的 REST 端点。
|
||||
type SessionHandler struct {
|
||||
sessionMgr session.Manager
|
||||
}
|
||||
|
||||
// NewSessionHandler 创建 SessionHandler。
|
||||
func NewSessionHandler(sessionMgr session.Manager) *SessionHandler {
|
||||
return &SessionHandler{sessionMgr: sessionMgr}
|
||||
}
|
||||
|
||||
// CreateSessionRequest POST /api/sessions 请求体(所有字段可选)。
|
||||
type CreateSessionRequest struct {
|
||||
Config *models.SessionConfig `json:"config,omitempty"`
|
||||
}
|
||||
|
||||
// CreateSession POST /api/sessions — 创建新会话。
|
||||
func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
|
||||
var req CreateSessionRequest
|
||||
// 请求体可选,解析失败不报错(使用默认配置)
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
|
||||
cfg := models.DefaultConfig()
|
||||
if req.Config != nil {
|
||||
cfg = *req.Config
|
||||
}
|
||||
|
||||
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",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// 获取创建后的会话以返回 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",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("session created",
|
||||
"session_id", sess.ID)
|
||||
c.JSON(http.StatusCreated, gin.H{
|
||||
"session_id": sess.ID,
|
||||
"created_at": sess.CreatedAt,
|
||||
})
|
||||
}
|
||||
|
||||
// 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)
|
||||
if err != nil {
|
||||
if err == session.ErrSessionNotFound {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": "SESSION_NOT_FOUND",
|
||||
"message": "session not found or already expired",
|
||||
})
|
||||
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",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("session destroyed",
|
||||
"session_id", sessionID)
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册会话相关路由到给定的路由组。
|
||||
func (h *SessionHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
||||
rg.POST("/sessions", h.CreateSession)
|
||||
rg.DELETE("/sessions/:id", h.DestroySession)
|
||||
}
|
||||
208
backend/internal/api/user_scenario_handler.go
Normal file
208
backend/internal/api/user_scenario_handler.go
Normal file
@@ -0,0 +1,208 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
MaxScenariosPerUser = 20 // 每个用户最多 20 个自建情景
|
||||
MaxPromptLength = 2000 // Prompt 最大长度
|
||||
)
|
||||
|
||||
// UserScenarioHandler 用户情景 API Handler。
|
||||
type UserScenarioHandler struct {
|
||||
repo store.UserScenarioRepository
|
||||
}
|
||||
|
||||
// NewUserScenarioHandler 创建用户情景 Handler。
|
||||
func NewUserScenarioHandler(repo store.UserScenarioRepository) *UserScenarioHandler {
|
||||
return &UserScenarioHandler{repo: repo}
|
||||
}
|
||||
|
||||
// List 获取用户的所有自建情景。
|
||||
// GET /api/scenarios
|
||||
func (h *UserScenarioHandler) List(c *gin.Context) {
|
||||
userID, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||
return
|
||||
}
|
||||
|
||||
scenarios, err := h.repo.FindByUserID(c.Request.Context(), userID.(string))
|
||||
if err != nil {
|
||||
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "查询失败"})
|
||||
return
|
||||
}
|
||||
|
||||
if scenarios == nil {
|
||||
scenarios = []*models.UserScenario{}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, models.UserScenarioListResponse{
|
||||
Scenarios: scenarios,
|
||||
Total: len(scenarios),
|
||||
})
|
||||
}
|
||||
|
||||
// Create 创建用户情景。
|
||||
// POST /api/scenarios
|
||||
func (h *UserScenarioHandler) Create(c *gin.Context) {
|
||||
userID, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||
return
|
||||
}
|
||||
|
||||
var req models.CreateUserScenarioRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// 检查用户是否已达上限
|
||||
count, err := h.repo.CountByUserID(c.Request.Context(), userID.(string))
|
||||
if err != nil {
|
||||
logger.Log.Errorw("统计用户情景数量失败", "user_id", userID, "error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建失败"})
|
||||
return
|
||||
}
|
||||
if count >= MaxScenariosPerUser {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "已达创建上限(最多 20 个)"})
|
||||
return
|
||||
}
|
||||
|
||||
// 创建情景
|
||||
scenario := &models.UserScenario{
|
||||
UserID: userID.(string),
|
||||
Name: req.Name,
|
||||
Icon: req.Icon,
|
||||
Description: req.Description,
|
||||
Prompt: req.Prompt,
|
||||
Greeting: req.Greeting,
|
||||
Language: req.Language,
|
||||
}
|
||||
|
||||
if err := h.repo.Create(c.Request.Context(), scenario); err != nil {
|
||||
logger.Log.Errorw("创建用户情景失败", "user_id", userID, "error", err)
|
||||
if err.Error() == "duplicate key value violates unique constraint" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "情景名称已存在"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建失败"})
|
||||
return
|
||||
}
|
||||
|
||||
logger.Log.Infow("创建用户情景成功", "user_id", userID, "scenario_id", scenario.ID)
|
||||
c.JSON(http.StatusCreated, scenario)
|
||||
}
|
||||
|
||||
// Get 获取单个情景详情。
|
||||
// GET /api/scenarios/:id
|
||||
func (h *UserScenarioHandler) Get(c *gin.Context) {
|
||||
userID, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||
return
|
||||
}
|
||||
|
||||
scenarioID := c.Param("id")
|
||||
scenario, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||
if err != nil {
|
||||
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, scenario)
|
||||
}
|
||||
|
||||
// Update 更新用户情景。
|
||||
// PATCH /api/scenarios/:id
|
||||
func (h *UserScenarioHandler) Update(c *gin.Context) {
|
||||
userID, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||
return
|
||||
}
|
||||
|
||||
scenarioID := c.Param("id")
|
||||
|
||||
// 查询并校验所有权
|
||||
scenario, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||
if err != nil {
|
||||
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||
return
|
||||
}
|
||||
|
||||
var req models.UpdateUserScenarioRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// 更新字段
|
||||
if req.Name != nil {
|
||||
scenario.Name = *req.Name
|
||||
}
|
||||
if req.Icon != nil {
|
||||
scenario.Icon = *req.Icon
|
||||
}
|
||||
if req.Description != nil {
|
||||
scenario.Description = *req.Description
|
||||
}
|
||||
if req.Prompt != nil {
|
||||
scenario.Prompt = *req.Prompt
|
||||
}
|
||||
if req.Greeting != nil {
|
||||
scenario.Greeting = *req.Greeting
|
||||
}
|
||||
if req.Language != nil {
|
||||
scenario.Language = *req.Language
|
||||
}
|
||||
|
||||
if err := h.repo.Update(c.Request.Context(), scenario); err != nil {
|
||||
logger.Log.Errorw("更新用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "更新失败"})
|
||||
return
|
||||
}
|
||||
|
||||
logger.Log.Infow("更新用户情景成功", "user_id", userID, "scenario_id", scenarioID)
|
||||
c.JSON(http.StatusOK, scenario)
|
||||
}
|
||||
|
||||
// Delete 删除用户情景。
|
||||
// DELETE /api/scenarios/:id
|
||||
func (h *UserScenarioHandler) Delete(c *gin.Context) {
|
||||
userID, exists := c.Get("user_id")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "未登录"})
|
||||
return
|
||||
}
|
||||
|
||||
scenarioID := c.Param("id")
|
||||
|
||||
// 查询并校验所有权
|
||||
_, err := h.repo.FindByIDAndUserID(c.Request.Context(), scenarioID, userID.(string))
|
||||
if err != nil {
|
||||
logger.Log.Errorw("查询用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "情景不存在或无权限"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.Delete(c.Request.Context(), scenarioID); err != nil {
|
||||
logger.Log.Errorw("删除用户情景失败", "user_id", userID, "scenario_id", scenarioID, "error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "删除失败"})
|
||||
return
|
||||
}
|
||||
|
||||
logger.Log.Infow("删除用户情景成功", "user_id", userID, "scenario_id", scenarioID)
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
134
backend/internal/auth/jwt.go
Normal file
134
backend/internal/auth/jwt.go
Normal 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[:])
|
||||
}
|
||||
165
backend/internal/auth/jwt_test.go
Normal file
165
backend/internal/auth/jwt_test.go
Normal 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)))
|
||||
}
|
||||
70
backend/internal/auth/middleware.go
Normal file
70
backend/internal/auth/middleware.go
Normal 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()
|
||||
}
|
||||
}
|
||||
19
backend/internal/auth/password.go
Normal file
19
backend/internal/auth/password.go
Normal 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))
|
||||
}
|
||||
224
backend/internal/auth/service.go
Normal file
224
backend/internal/auth/service.go
Normal 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 token(rotation)
|
||||
_ = 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)
|
||||
}
|
||||
235
backend/internal/auth/service_test.go
Normal file
235
backend/internal/auth/service_test.go
Normal 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 失败已间接验证
|
||||
}
|
||||
272
backend/internal/config/config.go
Normal file
272
backend/internal/config/config.go
Normal file
@@ -0,0 +1,272 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/joho/godotenv"
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
// Config 应用配置。
|
||||
type Config struct {
|
||||
App AppConfig `mapstructure:"app"`
|
||||
Server ServerConfig `mapstructure:"server"`
|
||||
Session SessionConfig `mapstructure:"session"`
|
||||
Redis RedisConfig `mapstructure:"redis"`
|
||||
AI AIConfig `mapstructure:"ai"`
|
||||
Storage StorageConfig `mapstructure:"storage"`
|
||||
Log LogConfig `mapstructure:"log"`
|
||||
Auth AuthConfig `mapstructure:"auth"`
|
||||
RateLimit RateLimitConfig `mapstructure:"ratelimit"`
|
||||
}
|
||||
|
||||
// SessionConfig 会话管理配置。
|
||||
type SessionConfig struct {
|
||||
TTL int `mapstructure:"ttl"` // 会话过期时间(分钟)
|
||||
MaxHistory int `mapstructure:"max_history"` // 对话历史上限(条)
|
||||
}
|
||||
|
||||
type AppConfig struct {
|
||||
Env string `mapstructure:"env"`
|
||||
Version string `mapstructure:"version"`
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
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 地址。
|
||||
func (s ServerConfig) Addr() string {
|
||||
return fmt.Sprintf("%s:%d", s.Host, s.Port)
|
||||
}
|
||||
|
||||
type RedisConfig struct {
|
||||
Addr string `mapstructure:"addr"`
|
||||
Password string `mapstructure:"password"`
|
||||
DB int `mapstructure:"db"`
|
||||
}
|
||||
|
||||
type AIConfig struct {
|
||||
STT STTConfig `mapstructure:"stt"`
|
||||
LLM LLMConfig `mapstructure:"llm"`
|
||||
TTS TTSConfig `mapstructure:"tts"`
|
||||
}
|
||||
|
||||
type STTConfig struct {
|
||||
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"`
|
||||
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"`
|
||||
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"`
|
||||
}
|
||||
|
||||
// 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 过期时间(分钟),默认 10080(7天)
|
||||
}
|
||||
|
||||
// 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(filepath.Join(workDir, "config")) // 配置文件在 config/ 目录下
|
||||
v.AddConfigPath(workDir) // 兼容旧路径
|
||||
|
||||
// 2. 设置默认值(与 config.yaml 保持一致,仅作为兜底)
|
||||
setDefaults(v)
|
||||
|
||||
// 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() // 文件不存在也不报错
|
||||
}
|
||||
|
||||
// 5. 显式绑定敏感信息环境变量(不用 AutomaticEnv,避免隐式映射)
|
||||
bindEnvVars(v)
|
||||
|
||||
var cfg Config
|
||||
if err := v.Unmarshal(&cfg); err != nil {
|
||||
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")
|
||||
}
|
||||
172
backend/internal/eino/adapter.go
Normal file
172
backend/internal/eino/adapter.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// EinoOrchestrator 实现 orchestrator.Orchestrator 接口。
|
||||
// 将 Eino Graph 包装为现有接口,WS Handler 几乎不用改。
|
||||
type EinoOrchestrator struct {
|
||||
graph *PipelineGraph
|
||||
sessionMgr session.Manager
|
||||
model string
|
||||
callbacks compose.Option // 运行时 Callback option
|
||||
}
|
||||
|
||||
// NewEinoOrchestrator 创建 Eino 编排器适配器。
|
||||
func NewEinoOrchestrator(graph *PipelineGraph, sessionMgr session.Manager, model string) *EinoOrchestrator {
|
||||
return &EinoOrchestrator{
|
||||
graph: graph,
|
||||
sessionMgr: sessionMgr,
|
||||
model: model,
|
||||
callbacks: compose.WithCallbacks(BuildCallbackHandler()),
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessQuery 实现 orchestrator.Orchestrator 接口。
|
||||
func (e *EinoOrchestrator) ProcessQuery(
|
||||
ctx context.Context,
|
||||
sessionID string,
|
||||
req models.WsQuery,
|
||||
sender orchestrator.Sender,
|
||||
) error {
|
||||
log := trace.FromContext(ctx)
|
||||
startTime := time.Now()
|
||||
|
||||
// 1. 设置活跃请求
|
||||
if err := e.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil {
|
||||
return err
|
||||
}
|
||||
defer e.sessionMgr.ClearActiveRequest(ctx, sessionID)
|
||||
|
||||
// 2. 获取会话配置
|
||||
sess, err := e.sessionMgr.Get(ctx, sessionID)
|
||||
if err != nil {
|
||||
log.Errorw("get session failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
Code: "SESSION_NOT_FOUND",
|
||||
Message: "会话不存在",
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. 解码音频和图片
|
||||
var audioData []byte
|
||||
if req.Text == "" && req.Audio != "" {
|
||||
audioData, err = base64.StdEncoding.DecodeString(req.Audio)
|
||||
if err != nil {
|
||||
log.Errorw("audio decode failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "音频数据解码失败",
|
||||
})
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
var imageData []byte
|
||||
if req.Image != "" {
|
||||
imageData, err = base64.StdEncoding.DecodeString(req.Image)
|
||||
if err != nil {
|
||||
log.Errorw("image decode failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "图片数据解码失败",
|
||||
})
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 4. 构建 Graph 输入
|
||||
input := buildPipelineInput(req, sessionID, sess, audioData, imageData)
|
||||
|
||||
// 5. 注入 context 值(供 Callback 和 Lambda 节点使用)
|
||||
ctx = WithSender(ctx, sender)
|
||||
ctx = WithRequestID(ctx, req.RequestID)
|
||||
ctx = trace.WithSessionID(ctx, sessionID)
|
||||
ctx = WithStartTime(ctx, startTime)
|
||||
|
||||
// 创建 State 并从 input 复制元数据
|
||||
state := genLocalState(ctx)
|
||||
state.SessionID = input.SessionID
|
||||
state.RequestID = input.RequestID
|
||||
state.ImageData = input.ImageData
|
||||
state.Scenario = input.Scenario
|
||||
state.Language = input.Language
|
||||
state.DetailLevel = sess.Config.DetailLevel
|
||||
state.TTSEnabled = input.TTSEnabled
|
||||
state.UserID = input.UserID
|
||||
ctx = WithPipelineState(ctx, state)
|
||||
|
||||
// 6. 调用 Graph(Stream 模式 + 运行时 Callback)
|
||||
streamReader, err := e.graph.Runnable.Stream(ctx, input, e.callbacks)
|
||||
if err != nil {
|
||||
log.Errorw("graph stream start failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
Code: "INTERNAL_ERROR",
|
||||
Message: "编排器启动失败",
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// 7. 消费 StreamReader(触发整条链路执行,side effects 推送消息到客户端)
|
||||
var output PipelineOutput
|
||||
for {
|
||||
o, err := streamReader.Recv()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
log.Errorw("graph stream consume error", "error", err)
|
||||
break
|
||||
}
|
||||
output = o
|
||||
}
|
||||
|
||||
// 8. 追加用户消息到历史(使用 STT 结果,兼容文本输入和语音输入)
|
||||
userText := output.TranscribedText
|
||||
if userText == "" {
|
||||
userText = req.Text // fallback 到原始文本输入
|
||||
}
|
||||
if userText != "" {
|
||||
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
||||
Role: "user",
|
||||
Content: userText,
|
||||
}); err != nil {
|
||||
log.Errorw("append user message failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 9. 追加助手消息到历史
|
||||
if output.FullResponse != "" {
|
||||
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
|
||||
Role: "assistant",
|
||||
Content: output.FullResponse,
|
||||
}); err != nil {
|
||||
log.Errorw("append assistant message failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
latency := time.Since(startTime).Milliseconds()
|
||||
log.Infow("eino pipeline completed", "latency_ms", latency)
|
||||
|
||||
return nil
|
||||
}
|
||||
131
backend/internal/eino/callback.go
Normal file
131
backend/internal/eino/callback.go
Normal file
@@ -0,0 +1,131 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
|
||||
"github.com/cloudwego/eino/callbacks"
|
||||
"github.com/cloudwego/eino/components/model"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
callbacksHelper "github.com/cloudwego/eino/utils/callbacks"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// context key 类型,避免与其他包冲突。
|
||||
type ctxKeySender struct{}
|
||||
type ctxKeyState struct{}
|
||||
|
||||
// WithSender 将 Sender 注入 context。
|
||||
func WithSender(ctx context.Context, sender orchestrator.Sender) context.Context {
|
||||
return context.WithValue(ctx, ctxKeySender{}, sender)
|
||||
}
|
||||
|
||||
// WithRequestID 将 requestID 注入 context(使用 trace 包)。
|
||||
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||
return trace.WithRequestID(ctx, requestID)
|
||||
}
|
||||
|
||||
// WithPipelineState 将 PipelineState 注入 context。
|
||||
func WithPipelineState(ctx context.Context, state *PipelineState) context.Context {
|
||||
return context.WithValue(ctx, ctxKeyState{}, state)
|
||||
}
|
||||
|
||||
// senderFromCtx 从 context 获取 Sender。
|
||||
func senderFromCtx(ctx context.Context) orchestrator.Sender {
|
||||
s, _ := ctx.Value(ctxKeySender{}).(orchestrator.Sender)
|
||||
return s
|
||||
}
|
||||
|
||||
// requestIDFromCtx 从 context 获取 requestID(使用 trace 包)。
|
||||
func requestIDFromCtx(ctx context.Context) string {
|
||||
return trace.GetRequestID(ctx)
|
||||
}
|
||||
|
||||
// stateFromCtx 从 context 获取 PipelineState。
|
||||
func stateFromCtx(ctx context.Context) *PipelineState {
|
||||
s, _ := ctx.Value(ctxKeyState{}).(*PipelineState)
|
||||
return s
|
||||
}
|
||||
|
||||
// BuildCallbackHandler 构建 Eino Callback Handler。
|
||||
//
|
||||
// 核心职责:ChatModel 节点通过 OnEndWithStreamOutput 逐 token 推送 llm_chunk 到客户端,
|
||||
// 同时累积完整文本到 PipelineState。
|
||||
//
|
||||
// 其他节点的消息推送(stt_result、tts_audio、llm_done)由各 Lambda 内部直接调用 Sender。
|
||||
func BuildCallbackHandler() callbacks.Handler {
|
||||
return callbacksHelper.NewHandlerHelper().
|
||||
ChatModel(&callbacksHelper.ModelCallbackHandler{
|
||||
OnEndWithStreamOutput: func(ctx context.Context, info *callbacks.RunInfo, output *schema.StreamReader[*model.CallbackOutput]) context.Context {
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
state := stateFromCtx(ctx)
|
||||
|
||||
if sender == nil || requestID == "" {
|
||||
log.Warnw("ModelCallback: missing sender or request_id in context",
|
||||
"node", info.Name)
|
||||
return ctx
|
||||
}
|
||||
|
||||
// 异步消费流,避免阻塞框架的下游处理。
|
||||
// 框架对流做了内部拷贝,此 goroutine 读取独立副本。
|
||||
go func() {
|
||||
defer output.Close()
|
||||
|
||||
for {
|
||||
chunk, err := output.Recv()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
return
|
||||
}
|
||||
log.Errorw("ModelCallback: stream recv error",
|
||||
"node", info.Name, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
if chunk == nil || chunk.Message == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
delta := chunk.Message.Content
|
||||
if delta == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// 推送 llm_chunk 到客户端
|
||||
if err := sender.SendLLMChunk(models.WsLLMChunk{
|
||||
Type: "llm_chunk",
|
||||
RequestID: requestID,
|
||||
Delta: delta,
|
||||
Role: "assistant",
|
||||
}); err != nil {
|
||||
log.Errorw("ModelCallback: send llm_chunk failed", "error", err)
|
||||
}
|
||||
|
||||
// 累积完整文本到 State
|
||||
if state != nil {
|
||||
state.AppendText(delta)
|
||||
}
|
||||
|
||||
// 记录 token 用量(流的最后一帧携带)
|
||||
if chunk.TokenUsage != nil && state != nil {
|
||||
state.mu.Lock()
|
||||
state.TokenUsage = &TokenUsage{
|
||||
Prompt: chunk.TokenUsage.PromptTokens,
|
||||
Completion: chunk.TokenUsage.CompletionTokens,
|
||||
Total: chunk.TokenUsage.TotalTokens,
|
||||
}
|
||||
state.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
return ctx
|
||||
},
|
||||
}).
|
||||
Handler()
|
||||
}
|
||||
119
backend/internal/eino/graph.go
Normal file
119
backend/internal/eino/graph.go
Normal file
@@ -0,0 +1,119 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
openaiImpl "github.com/cloudwego/eino-ext/components/model/openai"
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/stt"
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
nodeSTT = "stt"
|
||||
nodeHistory = "history"
|
||||
nodeLLM = "llm"
|
||||
nodeMessageToString = "msg2str"
|
||||
nodeSplitter = "splitter"
|
||||
nodeTTS = "tts"
|
||||
nodeDone = "done"
|
||||
)
|
||||
|
||||
// PipelineGraph 封装编译后的 Eino Graph。
|
||||
type PipelineGraph struct {
|
||||
Runnable compose.Runnable[PipelineInput, PipelineOutput]
|
||||
}
|
||||
|
||||
// NewPipelineGraph 构建 CamTalk AI 编排 Graph。
|
||||
//
|
||||
// 拓扑:START → STT → History → ChatModel → Splitter → TTS → Done → END
|
||||
//
|
||||
// Graph 使用 Stream 模式调用,ChatModel 实现真正的 token 级流式输出。
|
||||
// LLM token 通过 Callback 的 OnEndWithStreamOutput 实时推送到客户端。
|
||||
func NewPipelineGraph(
|
||||
ctx context.Context,
|
||||
cfg *config.Config,
|
||||
sttService stt.Service,
|
||||
ttsService tts.Service,
|
||||
sessionMgr session.Manager,
|
||||
scenarioRepo store.UserScenarioRepository,
|
||||
) (*PipelineGraph, error) {
|
||||
log := logger.Log
|
||||
|
||||
// 1. 创建 eino-ext ChatModel(对接 DashScope OpenAI 兼容接口)
|
||||
chatModel, err := openaiImpl.NewChatModel(ctx, &openaiImpl.ChatModelConfig{
|
||||
APIKey: cfg.AI.LLM.APIKey,
|
||||
Model: cfg.AI.LLM.Model,
|
||||
BaseURL: cfg.AI.LLM.Endpoint,
|
||||
Timeout: time.Duration(cfg.AI.LLM.Timeout) * time.Second,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
log.Infow("Eino ChatModel 初始化成功",
|
||||
"model", cfg.AI.LLM.Model,
|
||||
"endpoint", cfg.AI.LLM.Endpoint)
|
||||
|
||||
// 2. 构建 Graph(值类型,非指针)
|
||||
g := compose.NewGraph[PipelineInput, PipelineOutput](
|
||||
compose.WithGenLocalState(genLocalState),
|
||||
)
|
||||
|
||||
// 3. 添加节点
|
||||
maxHistory := cfg.Session.MaxHistory
|
||||
|
||||
_ = g.AddLambdaNode(nodeSTT, NewSTTLambda(sttService))
|
||||
_ = g.AddLambdaNode(nodeHistory, NewHistoryLambda(sessionMgr.GetHistory, scenarioRepo, maxHistory))
|
||||
_ = g.AddChatModelNode(nodeLLM, chatModel)
|
||||
_ = g.AddLambdaNode(nodeMessageToString, NewMessageToStringLambda())
|
||||
_ = g.AddLambdaNode(nodeSplitter, NewSplitterLambda())
|
||||
_ = g.AddLambdaNode(nodeTTS, NewTTSLambda(
|
||||
ttsService,
|
||||
cfg.AI.TTS.Voice,
|
||||
cfg.AI.TTS.Speed,
|
||||
cfg.AI.TTS.OutputFormat,
|
||||
cfg.AI.TTS.SampleRate,
|
||||
))
|
||||
_ = g.AddLambdaNode(nodeDone, NewDoneLambda(cfg.AI.LLM.Model))
|
||||
|
||||
// 4. 连接边
|
||||
_ = g.AddEdge(compose.START, nodeSTT)
|
||||
_ = g.AddEdge(nodeSTT, nodeHistory)
|
||||
_ = g.AddEdge(nodeHistory, nodeLLM)
|
||||
_ = g.AddEdge(nodeLLM, nodeMessageToString)
|
||||
_ = g.AddEdge(nodeMessageToString, nodeSplitter)
|
||||
_ = g.AddEdge(nodeSplitter, nodeTTS)
|
||||
_ = g.AddEdge(nodeTTS, nodeDone)
|
||||
_ = g.AddEdge(nodeDone, compose.END)
|
||||
|
||||
// 5. 编译(回调在运行时通过 Stream option 传入)
|
||||
runnable, err := g.Compile(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Infow("Eino Graph 编译成功", "nodes", 7)
|
||||
return &PipelineGraph{Runnable: runnable}, nil
|
||||
}
|
||||
|
||||
// buildPipelineInput 从 WebSocket 请求和会话配置构建 Graph 输入。
|
||||
func buildPipelineInput(req models.WsQuery, sessionID string, sess *models.Session, audioData, imageData []byte) PipelineInput {
|
||||
return PipelineInput{
|
||||
AudioData: audioData,
|
||||
ImageData: imageData,
|
||||
Text: req.Text,
|
||||
SessionID: sessionID,
|
||||
RequestID: req.RequestID,
|
||||
Language: sess.Config.Language,
|
||||
Scenario: sess.Config.Scenario,
|
||||
TTSEnabled: sess.Config.TTSEnabled,
|
||||
UserID: sess.UserID,
|
||||
}
|
||||
}
|
||||
236
backend/internal/eino/graph_test.go
Normal file
236
backend/internal/eino/graph_test.go
Normal file
@@ -0,0 +1,236 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/mock"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/stt"
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// --- Mock STT Service ---
|
||||
|
||||
type mockSTTService struct {
|
||||
mock.Mock
|
||||
}
|
||||
|
||||
func (m *mockSTTService) Recognize(ctx context.Context, audio []byte, opts stt.Options) (string, error) {
|
||||
args := m.Called(ctx, audio, opts)
|
||||
return args.String(0), args.Error(1)
|
||||
}
|
||||
|
||||
// --- Mock TTS Service ---
|
||||
|
||||
type mockTTSService struct {
|
||||
mock.Mock
|
||||
}
|
||||
|
||||
func (m *mockTTSService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts tts.Options) (<-chan tts.Chunk, error) {
|
||||
args := m.Called(ctx, textStream, opts)
|
||||
return args.Get(0).(<-chan tts.Chunk), args.Error(1)
|
||||
}
|
||||
|
||||
// --- Mock Sender ---
|
||||
|
||||
type mockSender struct {
|
||||
mock.Mock
|
||||
STTResults []models.WsSTTResult
|
||||
LLMChunks []models.WsLLMChunk
|
||||
LLMDones []models.WsLLMDone
|
||||
TTSAudios []models.WsTTSAudio
|
||||
Errors []models.WsError
|
||||
}
|
||||
|
||||
func (m *mockSender) SendSTTResult(result models.WsSTTResult) error {
|
||||
m.STTResults = append(m.STTResults, result)
|
||||
return m.Called(result).Error(0)
|
||||
}
|
||||
|
||||
func (m *mockSender) SendLLMChunk(chunk models.WsLLMChunk) error {
|
||||
m.LLMChunks = append(m.LLMChunks, chunk)
|
||||
return m.Called(chunk).Error(0)
|
||||
}
|
||||
|
||||
func (m *mockSender) SendLLMDone(done models.WsLLMDone) error {
|
||||
m.LLMDones = append(m.LLMDones, done)
|
||||
return m.Called(done).Error(0)
|
||||
}
|
||||
|
||||
func (m *mockSender) SendTTSAudio(audio models.WsTTSAudio) error {
|
||||
m.TTSAudios = append(m.TTSAudios, audio)
|
||||
return m.Called(audio).Error(0)
|
||||
}
|
||||
|
||||
func (m *mockSender) SendError(err models.WsError) error {
|
||||
m.Errors = append(m.Errors, err)
|
||||
return m.Called(err).Error(0)
|
||||
}
|
||||
|
||||
// --- Tests ---
|
||||
|
||||
func TestDetectImageMimeType(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
expected string
|
||||
}{
|
||||
{"JPEG", []byte{0xFF, 0xD8, 0xFF, 0xE0}, "image/jpeg"},
|
||||
{"PNG", []byte{0x89, 0x50, 0x4E, 0x47}, "image/png"},
|
||||
{"GIF", []byte{0x47, 0x49, 0x46, 0x38}, "image/gif"},
|
||||
{"WebP", []byte{0x52, 0x49, 0x46, 0x46}, "image/webp"},
|
||||
{"Unknown", []byte{0x00, 0x00, 0x00}, "image/jpeg"},
|
||||
{"Short", []byte{0xFF}, "image/jpeg"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := detectImageMimeType(tt.data)
|
||||
assert.Equal(t, tt.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildPipelineInput(t *testing.T) {
|
||||
req := models.WsQuery{
|
||||
Text: "你好",
|
||||
RequestID: "req-1",
|
||||
}
|
||||
sess := &models.Session{
|
||||
Config: models.SessionConfig{
|
||||
Language: "zh-CN",
|
||||
Scenario: "free_chat",
|
||||
TTSEnabled: true,
|
||||
},
|
||||
}
|
||||
|
||||
input := buildPipelineInput(req, "sess-1", sess, nil, nil)
|
||||
require.Equal(t, "你好", input.Text)
|
||||
require.Equal(t, "sess-1", input.SessionID)
|
||||
require.Equal(t, "req-1", input.RequestID)
|
||||
require.Equal(t, "zh-CN", input.Language)
|
||||
require.Equal(t, "free_chat", input.Scenario)
|
||||
require.True(t, input.TTSEnabled)
|
||||
}
|
||||
|
||||
func TestBuildPipelineInput_WithAudioData(t *testing.T) {
|
||||
req := models.WsQuery{
|
||||
Audio: "base64audio",
|
||||
RequestID: "req-2",
|
||||
}
|
||||
sess := &models.Session{
|
||||
Config: models.SessionConfig{
|
||||
Language: "en",
|
||||
Scenario: "free_chat",
|
||||
TTSEnabled: false,
|
||||
},
|
||||
}
|
||||
|
||||
audioData := []byte("fake-audio-bytes")
|
||||
imageData := []byte("fake-image-bytes")
|
||||
|
||||
input := buildPipelineInput(req, "sess-2", sess, audioData, imageData)
|
||||
require.Equal(t, audioData, input.AudioData)
|
||||
require.Equal(t, imageData, input.ImageData)
|
||||
require.False(t, input.TTSEnabled)
|
||||
require.Equal(t, "en", input.Language)
|
||||
}
|
||||
|
||||
func TestPipelineState_AppendAndGet(t *testing.T) {
|
||||
state := genLocalState(context.Background())
|
||||
|
||||
state.AppendText("Hello ")
|
||||
state.AppendText("World")
|
||||
|
||||
require.Equal(t, "Hello World", state.GetFullResponse())
|
||||
}
|
||||
|
||||
func TestPipelineState_ConcurrentAccess(t *testing.T) {
|
||||
state := genLocalState(context.Background())
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for i := 0; i < 100; i++ {
|
||||
state.AppendText("a")
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
|
||||
for i := 0; i < 100; i++ {
|
||||
_ = state.GetFullResponse()
|
||||
}
|
||||
|
||||
<-done
|
||||
require.Equal(t, 100, len(state.GetFullResponse()))
|
||||
}
|
||||
|
||||
func TestContextInjection(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
sender := &mockSender{}
|
||||
ctx = WithSender(ctx, sender)
|
||||
ctx = WithRequestID(ctx, "req-123")
|
||||
ctx = trace.WithSessionID(ctx, "sess-456")
|
||||
ctx = WithStartTime(ctx, time.Now())
|
||||
ctx = WithPipelineState(ctx, genLocalState(ctx))
|
||||
|
||||
require.NotNil(t, senderFromCtx(ctx))
|
||||
require.Equal(t, "req-123", requestIDFromCtx(ctx))
|
||||
require.NotNil(t, stateFromCtx(ctx))
|
||||
}
|
||||
|
||||
func TestLatencyFromCtx(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// No start time set
|
||||
require.Equal(t, int64(0), latencyFromCtx(ctx))
|
||||
|
||||
// With start time
|
||||
start := time.Now().Add(-100 * time.Millisecond)
|
||||
ctx = WithStartTime(ctx, start)
|
||||
latency := latencyFromCtx(ctx)
|
||||
require.Greater(t, latency, int64(0))
|
||||
require.Less(t, latency, int64(1000)) // should be < 1 second
|
||||
}
|
||||
|
||||
func TestEinoOrchestrator_ImplementsInterface(t *testing.T) {
|
||||
// Compile-time check that EinoOrchestrator implements orchestrator.Orchestrator
|
||||
var _ orchestrator.Orchestrator = (*EinoOrchestrator)(nil)
|
||||
}
|
||||
|
||||
func TestNewSTTLambda_ReturnsNonNil(t *testing.T) {
|
||||
mockSTT := &mockSTTService{}
|
||||
lambda := NewSTTLambda(mockSTT)
|
||||
require.NotNil(t, lambda)
|
||||
}
|
||||
|
||||
func TestNewHistoryLambda_ReturnsNonNil(t *testing.T) {
|
||||
fetcher := func(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||
return nil, nil
|
||||
}
|
||||
lambda := NewHistoryLambda(fetcher, nil, 10)
|
||||
require.NotNil(t, lambda)
|
||||
}
|
||||
|
||||
func TestNewSplitterLambda_ReturnsNonNil(t *testing.T) {
|
||||
lambda := NewSplitterLambda()
|
||||
require.NotNil(t, lambda)
|
||||
}
|
||||
|
||||
func TestNewTTSLambda_ReturnsNonNil(t *testing.T) {
|
||||
mockTTS := &mockTTSService{}
|
||||
lambda := NewTTSLambda(mockTTS, "alloy", 1.0, "mp3", 24000)
|
||||
require.NotNil(t, lambda)
|
||||
}
|
||||
|
||||
func TestNewDoneLambda_ReturnsNonNil(t *testing.T) {
|
||||
lambda := NewDoneLambda("test-model")
|
||||
require.NotNil(t, lambda)
|
||||
}
|
||||
86
backend/internal/eino/nodes_done.go
Normal file
86
backend/internal/eino/nodes_done.go
Normal file
@@ -0,0 +1,86 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// ctxKeyStartTime 请求开始时间的 context key。
|
||||
type ctxKeyStartTime struct{}
|
||||
|
||||
// WithStartTime 将请求开始时间注入 context。
|
||||
func WithStartTime(ctx context.Context, t time.Time) context.Context {
|
||||
return context.WithValue(ctx, ctxKeyStartTime{}, t)
|
||||
}
|
||||
|
||||
// latencyFromCtx 从 context 获取开始时间并计算延迟(毫秒)。
|
||||
func latencyFromCtx(ctx context.Context) int64 {
|
||||
if startTime, ok := ctx.Value(ctxKeyStartTime{}).(time.Time); ok {
|
||||
return time.Since(startTime).Milliseconds()
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// NewDoneLambda 创建 Done Lambda 节点。
|
||||
// 输入: struct{}(TTS 完成信号)→ 输出: *PipelineOutput
|
||||
//
|
||||
// 从 PipelineState 读取完整回复和 token 用量,发送 llm_done 到客户端。
|
||||
// 历史消息追加由适配器负责(避免重复写入)。
|
||||
func NewDoneLambda(defaultModel string) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, _ struct{}) (PipelineOutput, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
state := stateFromCtx(ctx)
|
||||
|
||||
if state == nil {
|
||||
return PipelineOutput{}, nil
|
||||
}
|
||||
|
||||
state.mu.Lock()
|
||||
fullResponse := state.FullResponse.String()
|
||||
transcribedText := state.TranscribedText
|
||||
tokenUsage := state.TokenUsage
|
||||
requestID := state.RequestID
|
||||
modelName := defaultModel
|
||||
state.mu.Unlock()
|
||||
|
||||
// 发送 llm_done
|
||||
if sender != nil && requestID != "" {
|
||||
done := models.WsLLMDone{
|
||||
Type: "llm_done",
|
||||
RequestID: requestID,
|
||||
FullText: fullResponse,
|
||||
Model: modelName,
|
||||
LatencyMs: latencyFromCtx(ctx),
|
||||
}
|
||||
if tokenUsage != nil {
|
||||
done.TokensUsed = struct {
|
||||
Prompt int `json:"prompt"`
|
||||
Completion int `json:"completion"`
|
||||
Total int `json:"total"`
|
||||
}{
|
||||
Prompt: tokenUsage.Prompt,
|
||||
Completion: tokenUsage.Completion,
|
||||
Total: tokenUsage.Total,
|
||||
}
|
||||
}
|
||||
if err := sender.SendLLMDone(done); err != nil {
|
||||
log.Errorw("send llm_done failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
log.Infow("query processing completed", "response_length", len(fullResponse))
|
||||
|
||||
return PipelineOutput{
|
||||
TranscribedText: transcribedText,
|
||||
FullResponse: fullResponse,
|
||||
Model: modelName,
|
||||
TokenUsage: tokenUsage,
|
||||
}, nil
|
||||
})
|
||||
}
|
||||
151
backend/internal/eino/nodes_history.go
Normal file
151
backend/internal/eino/nodes_history.go
Normal file
@@ -0,0 +1,151 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/llm"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// NewHistoryLambda 创建历史组装 Lambda 节点。
|
||||
// 输入: *STTOutput → 输出: []*schema.Message
|
||||
//
|
||||
// 从 PipelineState 读取请求元数据(SessionID、Scenario、ImageData 等),
|
||||
// 构建系统提示词,组装历史消息和当前用户输入(含多模态图片)。
|
||||
func NewHistoryLambda(
|
||||
historyFetcher func(ctx context.Context, sessionID string, limit int) ([]models.Message, error),
|
||||
scenarioRepo store.UserScenarioRepository,
|
||||
maxHistory int,
|
||||
) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, sttOut STTOutput) ([]*schema.Message, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
// 从 State 读取请求元数据
|
||||
state := stateFromCtx(ctx)
|
||||
if state == nil {
|
||||
return []*schema.Message{}, nil
|
||||
}
|
||||
|
||||
state.mu.Lock()
|
||||
sessionID := state.SessionID
|
||||
requestID := state.RequestID
|
||||
imageData := state.ImageData
|
||||
scenario := state.Scenario
|
||||
detailLevel := state.DetailLevel
|
||||
language := sttOut.Language
|
||||
userID := state.UserID
|
||||
state.mu.Unlock()
|
||||
|
||||
// 加载用户自建情景(如果有 userID 和 scenarioRepo)
|
||||
var customScenarios map[string]string
|
||||
var customGreetings map[string]string
|
||||
if userID != "" && scenarioRepo != nil {
|
||||
scenarios, err := scenarioRepo.FindByUserID(ctx, userID)
|
||||
if err != nil {
|
||||
log.Warnw("load user scenarios failed", "user_id", userID, "error", err)
|
||||
} else if len(scenarios) > 0 {
|
||||
customScenarios = make(map[string]string, len(scenarios))
|
||||
customGreetings = make(map[string]string, len(scenarios))
|
||||
for _, s := range scenarios {
|
||||
customScenarios[s.ID] = s.Prompt
|
||||
if s.Greeting != "" {
|
||||
customGreetings[s.ID] = s.Greeting
|
||||
}
|
||||
}
|
||||
log.Debugw("loaded user scenarios", "user_id", userID, "count", len(scenarios))
|
||||
}
|
||||
}
|
||||
|
||||
// 构建系统提示词(支持用户自建情景)
|
||||
scenarioPrompt := llm.GetScenarioPrompt(scenario, language, customScenarios)
|
||||
systemPrompt := llm.BuildSystemPrompt(language, detailLevel, scenarioPrompt)
|
||||
|
||||
// 构建 system message(仅文本,多模态内容只能放在 user 角色)
|
||||
systemMsg := &schema.Message{
|
||||
Role: schema.System,
|
||||
Content: systemPrompt,
|
||||
}
|
||||
|
||||
messages := []*schema.Message{systemMsg}
|
||||
|
||||
// 获取并追加历史消息
|
||||
if historyFetcher != nil && sessionID != "" {
|
||||
history, err := historyFetcher(ctx, sessionID, maxHistory)
|
||||
if err != nil {
|
||||
log.Warnw("fetch history failed, continuing", "error", err, "request_id", requestID)
|
||||
} else {
|
||||
for _, msg := range history {
|
||||
messages = append(messages, &schema.Message{
|
||||
Role: schema.RoleType(msg.Role),
|
||||
Content: msg.Content,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 追加当前用户输入(含图片,多模态内容只能放在 user 角色)
|
||||
// 注意:不能同时设置 Content 和 UserInputMultiContent,需要统一放到 MultiContent 中
|
||||
if len(imageData) > 0 {
|
||||
base64Str := base64.StdEncoding.EncodeToString(imageData)
|
||||
mimeType := detectImageMimeType(imageData)
|
||||
parts := []schema.MessageInputPart{
|
||||
{
|
||||
Type: schema.ChatMessagePartTypeText,
|
||||
Text: sttOut.Text,
|
||||
},
|
||||
{
|
||||
Type: schema.ChatMessagePartTypeImageURL,
|
||||
Image: &schema.MessageInputImage{
|
||||
MessagePartCommon: schema.MessagePartCommon{
|
||||
Base64Data: &base64Str,
|
||||
MIMEType: mimeType,
|
||||
},
|
||||
Detail: schema.ImageURLDetailAuto,
|
||||
},
|
||||
},
|
||||
}
|
||||
messages = append(messages, &schema.Message{
|
||||
Role: schema.User,
|
||||
UserInputMultiContent: parts,
|
||||
})
|
||||
} else {
|
||||
messages = append(messages, &schema.Message{
|
||||
Role: schema.User,
|
||||
Content: sttOut.Text,
|
||||
})
|
||||
}
|
||||
|
||||
log.Debugw("history assembled",
|
||||
"message_count", len(messages),
|
||||
"has_image", len(imageData) > 0,
|
||||
"scenario", scenario)
|
||||
|
||||
return messages, nil
|
||||
})
|
||||
}
|
||||
|
||||
// detectImageMimeType 简单检测图片 MIME 类型。
|
||||
func detectImageMimeType(data []byte) string {
|
||||
if len(data) < 4 {
|
||||
return "image/jpeg"
|
||||
}
|
||||
if data[0] == 0xFF && data[1] == 0xD8 && data[2] == 0xFF {
|
||||
return "image/jpeg"
|
||||
}
|
||||
if data[0] == 0x89 && data[1] == 0x50 && data[2] == 0x4E && data[3] == 0x47 {
|
||||
return "image/png"
|
||||
}
|
||||
if data[0] == 0x47 && data[1] == 0x49 && data[2] == 0x46 {
|
||||
return "image/gif"
|
||||
}
|
||||
if data[0] == 0x52 && data[1] == 0x49 && data[2] == 0x46 && data[3] == 0x46 {
|
||||
return "image/webp"
|
||||
}
|
||||
return "image/jpeg"
|
||||
}
|
||||
102
backend/internal/eino/nodes_splitter.go
Normal file
102
backend/internal/eino/nodes_splitter.go
Normal file
@@ -0,0 +1,102 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
)
|
||||
|
||||
// sentenceDelimiters 句子分隔符集合。
|
||||
var sentenceDelimiters = map[rune]bool{
|
||||
'。': true,
|
||||
'!': true,
|
||||
'?': true,
|
||||
'\n': true,
|
||||
'.': true,
|
||||
'!': true,
|
||||
'?': true,
|
||||
}
|
||||
|
||||
// NewMessageToStringLambda 创建 Message → String 转换 Lambda 节点。
|
||||
// 输入: *schema.Message → 输出: string
|
||||
//
|
||||
// 提取 Message.Content 文本,供 Splitter 节点消费。
|
||||
func NewMessageToStringLambda() *compose.Lambda {
|
||||
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[*schema.Message]) (*schema.StreamReader[string], error) {
|
||||
sr, sw := schema.Pipe[string](8)
|
||||
|
||||
go func() {
|
||||
defer sw.Close()
|
||||
defer input.Close()
|
||||
|
||||
for {
|
||||
msg, err := input.Recv()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
return
|
||||
}
|
||||
sw.Send("", err)
|
||||
return
|
||||
}
|
||||
if msg != nil && msg.Content != "" {
|
||||
sw.Send(msg.Content, nil)
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
return sr, nil
|
||||
})
|
||||
}
|
||||
|
||||
// NewSplitterLambda 创建句子分割 Transform Lambda 节点。
|
||||
// 输入: StreamReader[string](LLM token 流)→ 输出: StreamReader[string](完整句子流)
|
||||
//
|
||||
// 逐字符累积,按句子分隔符切分。每切出一个完整句子就输出一次,
|
||||
// 供下游 TTS 节点实时合成。
|
||||
func NewSplitterLambda() *compose.Lambda {
|
||||
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[string], error) {
|
||||
sr, sw := schema.Pipe[string](8)
|
||||
|
||||
go func() {
|
||||
defer sw.Close()
|
||||
defer input.Close()
|
||||
|
||||
var buffer strings.Builder
|
||||
|
||||
for {
|
||||
chunk, err := input.Recv()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
// 流结束,flush 剩余缓冲
|
||||
if buffer.Len() > 0 {
|
||||
text := strings.TrimSpace(buffer.String())
|
||||
if text != "" {
|
||||
sw.Send(text, nil)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
sw.Send("", err)
|
||||
return
|
||||
}
|
||||
|
||||
// 逐字符累积,按句子分隔符切分
|
||||
for _, r := range chunk {
|
||||
buffer.WriteRune(r)
|
||||
if sentenceDelimiters[r] {
|
||||
text := strings.TrimSpace(buffer.String())
|
||||
if text != "" {
|
||||
sw.Send(text, nil)
|
||||
}
|
||||
buffer.Reset()
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
return sr, nil
|
||||
})
|
||||
}
|
||||
135
backend/internal/eino/nodes_stt.go
Normal file
135
backend/internal/eino/nodes_stt.go
Normal file
@@ -0,0 +1,135 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/stt"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
)
|
||||
|
||||
// NewSTTLambda 创建 STT Lambda 节点。
|
||||
// 输入: PipelineInput → 输出: STTOutput
|
||||
//
|
||||
// 文本输入模式:跳过 STT,直接返回用户输入文本。
|
||||
// 语音模式:调用 sttService.Recognize() 进行语音识别。
|
||||
// 识别结果通过 Sender 发送 stt_result 到客户端。
|
||||
func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, input PipelineInput) (STTOutput, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
|
||||
// 将输入元数据写入 State,供下游节点(History、Done)读取
|
||||
if state := stateFromCtx(ctx); state != nil {
|
||||
state.mu.Lock()
|
||||
state.SessionID = input.SessionID
|
||||
state.RequestID = input.RequestID
|
||||
state.ImageData = input.ImageData
|
||||
state.Scenario = input.Scenario
|
||||
state.DetailLevel = "low"
|
||||
state.Language = input.Language
|
||||
state.TTSEnabled = input.TTSEnabled
|
||||
state.mu.Unlock()
|
||||
}
|
||||
|
||||
// 文本输入模式:跳过 STT
|
||||
if input.Text != "" {
|
||||
log.Debugw("text input mode, skipping stt",
|
||||
"text_len", len(input.Text),
|
||||
"text_preview", util.Truncate(input.Text, 50))
|
||||
|
||||
// 发送 stt_result 保持前端消息流一致性
|
||||
if sender != nil {
|
||||
if err := sender.SendSTTResult(models.WsSTTResult{
|
||||
Type: "stt_result",
|
||||
RequestID: requestID,
|
||||
Text: input.Text,
|
||||
IsFinal: true,
|
||||
}); err != nil {
|
||||
log.Errorw("send stt_result failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 写入 State
|
||||
if state := stateFromCtx(ctx); state != nil {
|
||||
state.mu.Lock()
|
||||
state.TranscribedText = input.Text
|
||||
state.mu.Unlock()
|
||||
}
|
||||
|
||||
return STTOutput{
|
||||
Text: input.Text,
|
||||
Language: input.Language,
|
||||
IsSkipped: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// 语音模式:解码音频
|
||||
if len(input.AudioData) == 0 {
|
||||
return STTOutput{}, fmt.Errorf("stt: no audio data provided")
|
||||
}
|
||||
|
||||
log.Debugw("stt recognition started", "audio_bytes", len(input.AudioData))
|
||||
|
||||
// 调用 STT 服务
|
||||
text, err := sttService.Recognize(ctx, input.AudioData, stt.Options{
|
||||
Encoding: "pcm_s16le",
|
||||
SampleRate: 16000,
|
||||
Language: input.Language,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorw("stt recognition failed", "error", err)
|
||||
if sender != nil {
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: requestID,
|
||||
Code: "STT_ERROR",
|
||||
Message: "语音识别失败: " + err.Error(),
|
||||
})
|
||||
}
|
||||
return STTOutput{}, fmt.Errorf("stt: recognize: %w", err)
|
||||
}
|
||||
|
||||
// STT 返回空文本
|
||||
if strings.TrimSpace(text) == "" {
|
||||
log.Infow("stt returned empty text")
|
||||
text = "(未识别到语音)"
|
||||
}
|
||||
|
||||
log.Debugw("stt recognition completed",
|
||||
"text_len", len(text),
|
||||
"text_preview", util.Truncate(text, 50))
|
||||
|
||||
// 发送 stt_result
|
||||
if sender != nil {
|
||||
if err := sender.SendSTTResult(models.WsSTTResult{
|
||||
Type: "stt_result",
|
||||
RequestID: requestID,
|
||||
Text: text,
|
||||
IsFinal: true,
|
||||
}); err != nil {
|
||||
log.Errorw("send stt_result failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 写入 State
|
||||
if state := stateFromCtx(ctx); state != nil {
|
||||
state.mu.Lock()
|
||||
state.TranscribedText = text
|
||||
state.mu.Unlock()
|
||||
}
|
||||
|
||||
return STTOutput{
|
||||
Text: text,
|
||||
Language: input.Language,
|
||||
IsSkipped: false,
|
||||
}, nil
|
||||
})
|
||||
}
|
||||
|
||||
116
backend/internal/eino/nodes_tts.go
Normal file
116
backend/internal/eino/nodes_tts.go
Normal file
@@ -0,0 +1,116 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/cloudwego/eino/schema"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// NewTTSLambda 创建 TTS Transform Lambda 节点。
|
||||
// 输入: StreamReader[string](句子流)→ 输出: StreamReader[struct{}](结果流)
|
||||
//
|
||||
// 流式消费每个句子,调用 ttsService.SynthesizeStream() 合成,
|
||||
// 逐 chunk 推送 tts_audio 到客户端。TTS 失败静默跳过。
|
||||
func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, ttsOutputFmt string, ttsSampleRate int) *compose.Lambda {
|
||||
return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[struct{}], error) {
|
||||
sr, sw := schema.Pipe[struct{}](8)
|
||||
|
||||
go func() {
|
||||
defer sw.Close()
|
||||
defer input.Close()
|
||||
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
|
||||
if sender == nil || requestID == "" {
|
||||
// 消费并丢弃流
|
||||
for {
|
||||
_, err := input.Recv()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 收集句子,按批次合成 TTS
|
||||
var sentences []string
|
||||
for {
|
||||
sentence, err := input.Recv()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
log.Errorw("TTS: stream recv error", "error", err)
|
||||
break
|
||||
}
|
||||
if sentence != "" {
|
||||
sentences = append(sentences, sentence)
|
||||
}
|
||||
}
|
||||
|
||||
if len(sentences) == 0 {
|
||||
sw.Send(struct{}{}, nil)
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("开始 TTS 合成", "sentence_count", len(sentences))
|
||||
|
||||
// 将句子数组转为 channel
|
||||
sentenceCh := make(chan string, len(sentences))
|
||||
for _, s := range sentences {
|
||||
sentenceCh <- s
|
||||
}
|
||||
close(sentenceCh)
|
||||
|
||||
// 调用 TTS 服务
|
||||
ttsStream, err := ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{
|
||||
Voice: ttsVoice,
|
||||
Speed: ttsSpeed,
|
||||
OutputFmt: ttsOutputFmt,
|
||||
SampleRate: ttsSampleRate,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorw("TTS 合成启动失败(已跳过)", "error", err)
|
||||
sw.Send(struct{}{}, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// 消费 TTS 音频流,推送到客户端
|
||||
for chunk := range ttsStream {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
log.Debugw("tts stream interrupted")
|
||||
sw.Send(struct{}{}, ctx.Err())
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
audioBase64 := base64.StdEncoding.EncodeToString(chunk.Audio)
|
||||
|
||||
if err := sender.SendTTSAudio(models.WsTTSAudio{
|
||||
Type: "tts_audio",
|
||||
RequestID: requestID,
|
||||
Audio: audioBase64,
|
||||
MimeType: "audio/mp3",
|
||||
IsLast: chunk.IsLast,
|
||||
Final: chunk.Final,
|
||||
}); err != nil {
|
||||
log.Errorw("发送 tts_audio 失败", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
log.Infow("TTS 合成完成")
|
||||
sw.Send(struct{}{}, nil)
|
||||
}()
|
||||
|
||||
return sr, nil
|
||||
})
|
||||
}
|
||||
46
backend/internal/eino/state.go
Normal file
46
backend/internal/eino/state.go
Normal file
@@ -0,0 +1,46 @@
|
||||
package eino
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// PipelineState Graph 全局状态,用于跨节点收集数据。
|
||||
// 通过 compose.WithGenLocalState 注册,各节点通过 compose.ProcessState 读写。
|
||||
type PipelineState struct {
|
||||
mu sync.Mutex
|
||||
FullResponse strings.Builder // LLM 完整回复(由 Callback 累积)
|
||||
TranscribedText string // STT 识别文本
|
||||
Model string // 实际使用的模型名
|
||||
TokenUsage *TokenUsage // token 用量
|
||||
|
||||
// 从 PipelineInput 复制的元数据,供下游节点(History、Done)读取
|
||||
SessionID string
|
||||
RequestID string
|
||||
ImageData []byte
|
||||
Scenario string
|
||||
DetailLevel string
|
||||
Language string
|
||||
TTSEnabled bool
|
||||
UserID string // 新增:用户 ID,用于加载自建情景
|
||||
}
|
||||
|
||||
// genLocalState 创建每请求的 PipelineState 实例。
|
||||
func genLocalState(ctx context.Context) *PipelineState {
|
||||
return &PipelineState{}
|
||||
}
|
||||
|
||||
// AppendText 追加文本到 FullResponse(线程安全)。
|
||||
func (s *PipelineState) AppendText(text string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.FullResponse.WriteString(text)
|
||||
}
|
||||
|
||||
// GetFullResponse 获取完整回复文本(线程安全)。
|
||||
func (s *PipelineState) GetFullResponse() string {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.FullResponse.String()
|
||||
}
|
||||
38
backend/internal/eino/types.go
Normal file
38
backend/internal/eino/types.go
Normal file
@@ -0,0 +1,38 @@
|
||||
// Package eino 基于 CloudWeGo Eino 框架的 AI 编排层。
|
||||
// 使用 Eino Graph 替代手写 goroutine 管道,实现声明式 STT → LLM → TTS 编排。
|
||||
package eino
|
||||
|
||||
// PipelineInput Graph 统一输入。
|
||||
type PipelineInput struct {
|
||||
AudioData []byte // base64 解码后的音频(可选)
|
||||
ImageData []byte // base64 解码后的图像(可选)
|
||||
Text string // 直接文本输入(可选,跳过 STT)
|
||||
SessionID string
|
||||
RequestID string
|
||||
Language string // zh / en
|
||||
Scenario string // free_chat, interviewer, etc.
|
||||
TTSEnabled bool
|
||||
UserID string // 用户 ID,用于加载自建情景
|
||||
}
|
||||
|
||||
// PipelineOutput Graph 统一输出。
|
||||
type PipelineOutput struct {
|
||||
TranscribedText string // STT 结果
|
||||
FullResponse string // LLM 完整回复
|
||||
Model string // 实际使用的模型名
|
||||
TokenUsage *TokenUsage // token 用量
|
||||
}
|
||||
|
||||
// STTOutput STT 节点输出。
|
||||
type STTOutput struct {
|
||||
Text string
|
||||
Language string
|
||||
IsSkipped bool // 文本输入模式跳过了 STT
|
||||
}
|
||||
|
||||
// TokenUsage token 用量统计。
|
||||
type TokenUsage struct {
|
||||
Prompt int
|
||||
Completion int
|
||||
Total int
|
||||
}
|
||||
39
backend/internal/errors/codes.go
Normal file
39
backend/internal/errors/codes.go
Normal file
@@ -0,0 +1,39 @@
|
||||
package errors
|
||||
|
||||
import "github.com/hhs/camtalk/internal/models"
|
||||
|
||||
// 错误码常量,与 docs/03-接口文档.md 保持一致。
|
||||
const (
|
||||
CodeInvalidMessage = "INVALID_MESSAGE"
|
||||
CodeSessionNotFound = "SESSION_NOT_FOUND"
|
||||
CodeRateLimited = "RATE_LIMITED"
|
||||
CodeImageTooLarge = "IMAGE_TOO_LARGE"
|
||||
CodeAudioTooShort = "AUDIO_TOO_SHORT"
|
||||
CodeLLMTimeout = "LLM_TIMEOUT"
|
||||
CodeLLMError = "LLM_ERROR"
|
||||
CodeSTTError = "STT_ERROR"
|
||||
CodeTTSError = "TTS_ERROR"
|
||||
CodeInternalError = "INTERNAL_ERROR"
|
||||
|
||||
// 认证相关错误码
|
||||
CodeUsernameTaken = "USERNAME_TAKEN"
|
||||
CodeInvalidCredentials = "INVALID_CREDENTIALS"
|
||||
CodeInvalidToken = "INVALID_TOKEN"
|
||||
CodeInvalidInput = "INVALID_INPUT"
|
||||
)
|
||||
|
||||
// Sender 定义发送 WS 错误消息的接口,便于测试 mock。
|
||||
type Sender interface {
|
||||
SendError(code, requestID, message string)
|
||||
}
|
||||
|
||||
// SendWSError 向客户端发送 error 消息。
|
||||
// sender 是一个具有 sendJSON 方法的对象,这里用接口抽象。
|
||||
func SendWSError(sender interface{ SendJSON(v any) error }, code, requestID string, err error) {
|
||||
_ = sender.SendJSON(models.WsError{
|
||||
Type: "error",
|
||||
Code: code,
|
||||
RequestID: requestID,
|
||||
Message: err.Error(),
|
||||
})
|
||||
}
|
||||
57
backend/internal/logger/logger.go
Normal file
57
backend/internal/logger/logger.go
Normal file
@@ -0,0 +1,57 @@
|
||||
package logger
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
)
|
||||
|
||||
// Log 是全局 SugaredLogger,由 Init 初始化。
|
||||
var Log *zap.SugaredLogger
|
||||
|
||||
// Init 初始化全局日志器。
|
||||
// level: "debug", "info", "warn", "error"
|
||||
// format: "json" 或 "console"
|
||||
func Init(level, format string) {
|
||||
var lvl zapcore.Level
|
||||
switch level {
|
||||
case "debug":
|
||||
lvl = zapcore.DebugLevel
|
||||
case "warn":
|
||||
lvl = zapcore.WarnLevel
|
||||
case "error":
|
||||
lvl = zapcore.ErrorLevel
|
||||
default:
|
||||
lvl = zapcore.InfoLevel
|
||||
}
|
||||
|
||||
encoderCfg := zap.NewProductionEncoderConfig()
|
||||
encoderCfg.TimeKey = "ts"
|
||||
encoderCfg.EncodeTime = zapcore.ISO8601TimeEncoder
|
||||
|
||||
var core zapcore.Core
|
||||
if format == "console" {
|
||||
core = zapcore.NewCore(
|
||||
zapcore.NewConsoleEncoder(encoderCfg),
|
||||
zapcore.AddSync(os.Stdout),
|
||||
lvl,
|
||||
)
|
||||
} else {
|
||||
core = zapcore.NewCore(
|
||||
zapcore.NewJSONEncoder(encoderCfg),
|
||||
zapcore.AddSync(os.Stdout),
|
||||
lvl,
|
||||
)
|
||||
}
|
||||
|
||||
logger := zap.New(core, zap.AddCaller(), zap.AddStacktrace(zapcore.ErrorLevel))
|
||||
Log = logger.Sugar()
|
||||
}
|
||||
|
||||
// Sync 刷新缓冲区,退出前调用。
|
||||
func Sync() {
|
||||
if Log != nil {
|
||||
_ = Log.Sync()
|
||||
}
|
||||
}
|
||||
@@ -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,48 @@ 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 会话配置增量更新(指针字段表示"未传则不更新")。
|
||||
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。
|
||||
func (p SessionConfigPatch) Apply(cfg *SessionConfig) {
|
||||
if p.TTSEnabled != nil {
|
||||
cfg.TTSEnabled = *p.TTSEnabled
|
||||
}
|
||||
if p.DetailLevel != nil {
|
||||
cfg.DetailLevel = *p.DetailLevel
|
||||
}
|
||||
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 对话消息。
|
||||
@@ -35,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"
|
||||
}
|
||||
|
||||
@@ -45,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"`
|
||||
}
|
||||
|
||||
@@ -91,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 消息。
|
||||
|
||||
43
backend/internal/models/user_scenario.go
Normal file
43
backend/internal/models/user_scenario.go
Normal file
@@ -0,0 +1,43 @@
|
||||
package models
|
||||
|
||||
import "time"
|
||||
|
||||
// UserScenario 用户自建情景。
|
||||
type UserScenario struct {
|
||||
ID string `json:"id"`
|
||||
UserID string `json:"user_id"`
|
||||
Name string `json:"name"`
|
||||
Icon string `json:"icon"`
|
||||
Description string `json:"description"`
|
||||
Prompt string `json:"prompt"`
|
||||
Greeting string `json:"greeting,omitempty"`
|
||||
Language string `json:"language"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// CreateUserScenarioRequest 创建用户情景请求。
|
||||
type CreateUserScenarioRequest struct {
|
||||
Name string `json:"name" binding:"required,min=2,max=50"`
|
||||
Icon string `json:"icon,omitempty"`
|
||||
Description string `json:"description,omitempty" binding:"omitempty,max=100"`
|
||||
Prompt string `json:"prompt" binding:"required,min=10,max=2000"`
|
||||
Greeting string `json:"greeting,omitempty" binding:"omitempty,max=500"`
|
||||
Language string `json:"language,omitempty"`
|
||||
}
|
||||
|
||||
// UpdateUserScenarioRequest 更新用户情景请求。
|
||||
type UpdateUserScenarioRequest struct {
|
||||
Name *string `json:"name,omitempty" binding:"omitempty,min=2,max=50"`
|
||||
Icon *string `json:"icon,omitempty"`
|
||||
Description *string `json:"description,omitempty" binding:"omitempty,max=100"`
|
||||
Prompt *string `json:"prompt,omitempty" binding:"omitempty,min=10,max=2000"`
|
||||
Greeting *string `json:"greeting,omitempty" binding:"omitempty,max=500"`
|
||||
Language *string `json:"language,omitempty"`
|
||||
}
|
||||
|
||||
// UserScenarioListResponse 用户情景列表响应。
|
||||
type UserScenarioListResponse struct {
|
||||
Scenarios []*UserScenario `json:"scenarios"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
31
backend/internal/orchestrator/orchestrator.go
Normal file
31
backend/internal/orchestrator/orchestrator.go
Normal file
@@ -0,0 +1,31 @@
|
||||
// Package orchestrator 实现 STT → LLM → TTS 流式并行管道。
|
||||
package orchestrator
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
// Orchestrator AI 编排器接口。
|
||||
// 接收查询并执行完整的 STT → LLM → TTS 管道。
|
||||
type Orchestrator interface {
|
||||
// ProcessQuery 处理一次用户查询。
|
||||
// ctx 用于整体超时和中断控制。
|
||||
// sessionID 用于会话管理和历史获取。
|
||||
// req 包含图像和音频数据。
|
||||
// sender 用于向客户端推送消息。
|
||||
ProcessQuery(
|
||||
ctx context.Context,
|
||||
sessionID string,
|
||||
req models.WsQuery,
|
||||
sender Sender,
|
||||
) error
|
||||
}
|
||||
|
||||
// QueryRequest 查询请求(内部使用)。
|
||||
type QueryRequest struct {
|
||||
Image []byte // JPEG 图片(已从 Base64 解码)
|
||||
Audio []byte // 音频数据(已从 Base64 解码)
|
||||
Language string // 语言,如 "zh-CN"
|
||||
}
|
||||
22
backend/internal/orchestrator/sender.go
Normal file
22
backend/internal/orchestrator/sender.go
Normal file
@@ -0,0 +1,22 @@
|
||||
package orchestrator
|
||||
|
||||
import "github.com/hhs/camtalk/internal/models"
|
||||
|
||||
// Sender 抽象 WebSocket 消息推送能力。
|
||||
// 便于测试时 mock,避免依赖真实 WebSocket 连接。
|
||||
type Sender interface {
|
||||
// SendSTTResult 发送语音识别结果。
|
||||
SendSTTResult(result models.WsSTTResult) error
|
||||
|
||||
// SendLLMChunk 发送 LLM 流式文本增量。
|
||||
SendLLMChunk(chunk models.WsLLMChunk) error
|
||||
|
||||
// SendLLMDone 发送 LLM 流结束信号。
|
||||
SendLLMDone(done models.WsLLMDone) error
|
||||
|
||||
// SendTTSAudio 发送 TTS 音频数据。
|
||||
SendTTSAudio(audio models.WsTTSAudio) error
|
||||
|
||||
// SendError 发送错误消息。
|
||||
SendError(err models.WsError) error
|
||||
}
|
||||
172
backend/internal/ratelimit/bucket.go
Normal file
172
backend/internal/ratelimit/bucket.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
)
|
||||
|
||||
// TokenBucket 内存令牌桶,适用于单实例部署。
|
||||
type TokenBucket struct {
|
||||
capacity int // 桶容量
|
||||
rate float64 // 每秒填充令牌数
|
||||
tokens float64 // 当前令牌数
|
||||
lastRefill time.Time // 上次填充时间
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// newTokenBucket 创建令牌桶。
|
||||
func newTokenBucket(capacity int, rate float64) *TokenBucket {
|
||||
return &TokenBucket{
|
||||
capacity: capacity,
|
||||
rate: rate,
|
||||
tokens: float64(capacity), // 初始满桶
|
||||
lastRefill: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// allow 尝试消耗一个令牌。
|
||||
func (b *TokenBucket) allow() (bool, time.Duration) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
elapsed := now.Sub(b.lastRefill).Seconds()
|
||||
|
||||
// 补充令牌
|
||||
newTokens := elapsed * b.rate
|
||||
b.tokens = min(float64(b.capacity), b.tokens+newTokens)
|
||||
b.lastRefill = now
|
||||
|
||||
// 尝试消耗一个令牌
|
||||
if b.tokens >= 1 {
|
||||
b.tokens -= 1
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// 计算需要等待的时间
|
||||
if b.rate == 0 {
|
||||
// rate=0 时永远无法补充令牌
|
||||
return false, 24 * time.Hour // 返回一个很大的值
|
||||
}
|
||||
retryAfter := time.Duration((1-b.tokens)/b.rate*1000) * time.Millisecond
|
||||
return false, retryAfter
|
||||
}
|
||||
|
||||
// MemoryLimiter 管理多个用户的令牌桶。
|
||||
type MemoryLimiter struct {
|
||||
buckets map[string]*TokenBucket
|
||||
config config.RateLimitConfig
|
||||
mu sync.RWMutex
|
||||
stopOnce sync.Once
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
// NewMemoryLimiter 创建内存限流器。
|
||||
func NewMemoryLimiter(cfg config.RateLimitConfig) *MemoryLimiter {
|
||||
limiter := &MemoryLimiter{
|
||||
buckets: make(map[string]*TokenBucket),
|
||||
config: cfg,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
// 启动后台清理 goroutine
|
||||
go limiter.cleanup()
|
||||
|
||||
return limiter
|
||||
}
|
||||
|
||||
// Allow 实现 Limiter 接口。
|
||||
func (l *MemoryLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
bucket := l.getOrCreateBucket(key)
|
||||
return bucket.allow()
|
||||
}
|
||||
|
||||
// Stop 实现 Limiter 接口。
|
||||
func (l *MemoryLimiter) Stop() {
|
||||
l.stopOnce.Do(func() {
|
||||
close(l.done)
|
||||
})
|
||||
}
|
||||
|
||||
// getOrCreateBucket 获取或创建令牌桶。
|
||||
func (l *MemoryLimiter) getOrCreateBucket(key string) *TokenBucket {
|
||||
// 先尝试读锁
|
||||
l.mu.RLock()
|
||||
bucket, exists := l.buckets[key]
|
||||
l.mu.RUnlock()
|
||||
|
||||
if exists {
|
||||
return bucket
|
||||
}
|
||||
|
||||
// 需要创建新桶,升级为写锁
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
// 双重检查(可能其他 goroutine 已创建)
|
||||
bucket, exists = l.buckets[key]
|
||||
if exists {
|
||||
return bucket
|
||||
}
|
||||
|
||||
// 根据 key 确定配置(简化版:假设 key 格式为 "userID:action")
|
||||
cfg := l.getBucketConfig(key)
|
||||
bucket = newTokenBucket(cfg.Capacity, cfg.Rate)
|
||||
l.buckets[key] = bucket
|
||||
|
||||
return bucket
|
||||
}
|
||||
|
||||
// getBucketConfig 根据 key 获取桶配置。
|
||||
func (l *MemoryLimiter) getBucketConfig(key string) config.BucketConfig {
|
||||
// 简化实现:从 key 后缀判断动作类型
|
||||
// 实际使用时调用方会传递正确的 key
|
||||
// 默认使用 query 配置
|
||||
return l.config.Query
|
||||
}
|
||||
|
||||
// cleanup 定期清理不活跃的桶。
|
||||
func (l *MemoryLimiter) cleanup() {
|
||||
ticker := time.NewTicker(10 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
l.removeInactiveBuckets()
|
||||
case <-l.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// removeInactiveBuckets 移除超过 10 分钟无活动的桶。
|
||||
func (l *MemoryLimiter) removeInactiveBuckets() {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, bucket := range l.buckets {
|
||||
bucket.mu.Lock()
|
||||
inactive := now.Sub(bucket.lastRefill) > 10*time.Minute
|
||||
bucket.mu.Unlock()
|
||||
|
||||
if inactive {
|
||||
delete(l.buckets, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// min 返回两个 float64 中的较小值。
|
||||
func min(a, b float64) float64 {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// 编译期接口检查
|
||||
var _ Limiter = (*MemoryLimiter)(nil)
|
||||
203
backend/internal/ratelimit/bucket_test.go
Normal file
203
backend/internal/ratelimit/bucket_test.go
Normal file
@@ -0,0 +1,203 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTokenBucket_Allow_FirstRequest(t *testing.T) {
|
||||
bucket := newTokenBucket(5, 0.2)
|
||||
|
||||
allowed, retryAfter := bucket.allow()
|
||||
|
||||
assert.True(t, allowed)
|
||||
assert.Equal(t, time.Duration(0), retryAfter)
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_ConsumeUntilEmpty(t *testing.T) {
|
||||
bucket := newTokenBucket(3, 0.2)
|
||||
|
||||
// 连续消耗 3 个令牌
|
||||
for i := 0; i < 3; i++ {
|
||||
allowed, _ := bucket.allow()
|
||||
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||
}
|
||||
|
||||
// 第 4 个请求应被拒绝
|
||||
allowed, retryAfter := bucket.allow()
|
||||
assert.False(t, allowed)
|
||||
assert.Greater(t, retryAfter, time.Duration(0))
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_RetryAfterCorrect(t *testing.T) {
|
||||
bucket := newTokenBucket(1, 1.0) // 每秒 1 个令牌
|
||||
|
||||
// 消耗唯一的令牌
|
||||
allowed, _ := bucket.allow()
|
||||
require.True(t, allowed)
|
||||
|
||||
// 立即再次请求应被拒绝
|
||||
allowed, retryAfter := bucket.allow()
|
||||
assert.False(t, allowed)
|
||||
// retryAfter 应约为 1 秒(允许一定误差)
|
||||
assert.InDelta(t, 1000, retryAfter.Milliseconds(), 100)
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_RefillAfterWait(t *testing.T) {
|
||||
bucket := newTokenBucket(2, 10.0) // 每秒 10 个令牌(每 100ms 一个)
|
||||
|
||||
// 消耗 2 个令牌
|
||||
bucket.allow()
|
||||
bucket.allow()
|
||||
|
||||
// 等待 150ms,应补充至少 1 个令牌
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
|
||||
allowed, _ := bucket.allow()
|
||||
assert.True(t, allowed)
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_CapacityLimit(t *testing.T) {
|
||||
bucket := newTokenBucket(3, 1.0)
|
||||
|
||||
// 等待足够长时间让桶"溢出"
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// 但最多只能消耗 capacity 个令牌
|
||||
for i := 0; i < 3; i++ {
|
||||
allowed, _ := bucket.allow()
|
||||
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||
}
|
||||
|
||||
// 第 4 个应被拒绝
|
||||
allowed, _ := bucket.allow()
|
||||
assert.False(t, allowed)
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_ConcurrentSafe(t *testing.T) {
|
||||
bucket := newTokenBucket(100, 10.0)
|
||||
var wg sync.WaitGroup
|
||||
successCount := 0
|
||||
var mu sync.Mutex
|
||||
|
||||
// 100 个并发请求
|
||||
for i := 0; i < 100; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
allowed, _ := bucket.allow()
|
||||
if allowed {
|
||||
mu.Lock()
|
||||
successCount++
|
||||
mu.Unlock()
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
|
||||
// 应该正好 100 个成功(桶容量为 100)
|
||||
assert.Equal(t, 100, successCount)
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_ZeroCapacity(t *testing.T) {
|
||||
bucket := newTokenBucket(0, 1.0)
|
||||
|
||||
allowed, retryAfter := bucket.allow()
|
||||
assert.False(t, allowed)
|
||||
assert.Greater(t, retryAfter, time.Duration(0))
|
||||
}
|
||||
|
||||
func TestTokenBucket_Allow_ZeroRate(t *testing.T) {
|
||||
bucket := newTokenBucket(1, 0.0)
|
||||
|
||||
// 第一个通过
|
||||
allowed, _ := bucket.allow()
|
||||
assert.True(t, allowed)
|
||||
|
||||
// 第二个被拒绝,且 retryAfter 应为无限大(实际上会很大)
|
||||
allowed, retryAfter := bucket.allow()
|
||||
assert.False(t, allowed)
|
||||
// rate=0 时,retryAfter 理论上无限大,实际会是一个很大的值
|
||||
assert.Greater(t, retryAfter, 1*time.Hour)
|
||||
}
|
||||
|
||||
func TestMemoryLimiter_Allow_DifferentKeys(t *testing.T) {
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
|
||||
}
|
||||
limiter := NewMemoryLimiter(cfg)
|
||||
defer limiter.Stop()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// user1 消耗 2 个令牌
|
||||
allowed, _ := limiter.Allow(ctx, "user1:query")
|
||||
assert.True(t, allowed)
|
||||
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||
assert.True(t, allowed)
|
||||
|
||||
// user1 第 3 个被拒绝
|
||||
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||
assert.False(t, allowed)
|
||||
|
||||
// user2 应该不受影响
|
||||
allowed, _ = limiter.Allow(ctx, "user2:query")
|
||||
assert.True(t, allowed)
|
||||
}
|
||||
|
||||
func TestMemoryLimiter_Cleanup(t *testing.T) {
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||
}
|
||||
limiter := NewMemoryLimiter(cfg)
|
||||
defer limiter.Stop()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建一个桶
|
||||
limiter.Allow(ctx, "user1:query")
|
||||
|
||||
// 验证桶已创建
|
||||
limiter.mu.RLock()
|
||||
initialCount := len(limiter.buckets)
|
||||
limiter.mu.RUnlock()
|
||||
assert.Equal(t, 1, initialCount)
|
||||
|
||||
// 手动触发清理(模拟 10 分钟后)
|
||||
limiter.mu.Lock()
|
||||
for _, bucket := range limiter.buckets {
|
||||
bucket.mu.Lock()
|
||||
bucket.lastRefill = time.Now().Add(-11 * time.Minute)
|
||||
bucket.mu.Unlock()
|
||||
}
|
||||
limiter.mu.Unlock()
|
||||
|
||||
limiter.removeInactiveBuckets()
|
||||
|
||||
// 验证桶已被清理
|
||||
limiter.mu.RLock()
|
||||
finalCount := len(limiter.buckets)
|
||||
limiter.mu.RUnlock()
|
||||
assert.Equal(t, 0, finalCount)
|
||||
}
|
||||
|
||||
func TestMemoryLimiter_Stop(t *testing.T) {
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||
}
|
||||
limiter := NewMemoryLimiter(cfg)
|
||||
|
||||
// 多次调用 Stop 不应 panic
|
||||
limiter.Stop()
|
||||
limiter.Stop()
|
||||
}
|
||||
17
backend/internal/ratelimit/limiter.go
Normal file
17
backend/internal/ratelimit/limiter.go
Normal file
@@ -0,0 +1,17 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Limiter 速率限制器接口。
|
||||
type Limiter interface {
|
||||
// Allow 判断 key 是否允许执行一次操作。
|
||||
// key 通常为 "userID:action" 格式。
|
||||
// 返回 (allowed, retryAfter)。retryAfter 表示需要等待的时间。
|
||||
Allow(ctx context.Context, key string) (bool, time.Duration)
|
||||
|
||||
// Stop 停止限流器,清理资源(如后台 goroutine)。
|
||||
Stop()
|
||||
}
|
||||
51
backend/internal/ratelimit/middleware.go
Normal file
51
backend/internal/ratelimit/middleware.go
Normal file
@@ -0,0 +1,51 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// Middleware 返回 Gin 中间件,按 key 维度限流。
|
||||
// keyFunc 从请求中提取限流 key(如 IP、用户 ID)。
|
||||
func Middleware(limiter Limiter, keyFunc func(*gin.Context) string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if limiter == nil {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
key := keyFunc(c)
|
||||
if key == "" {
|
||||
// key 为空时跳过限流
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
allowed, retryAfter := limiter.Allow(c.Request.Context(), key)
|
||||
|
||||
if !allowed {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
log.Warnw("rate limited",
|
||||
"client_ip", c.ClientIP(),
|
||||
"path", c.Request.URL.Path,
|
||||
"limit_key", key,
|
||||
"retry_after_sec", int(retryAfter.Seconds()+0.5))
|
||||
|
||||
// 设置 Retry-After header(秒)
|
||||
c.Header("Retry-After", fmt.Sprintf("%d", int(retryAfter.Seconds()+0.5)))
|
||||
|
||||
c.JSON(http.StatusTooManyRequests, gin.H{
|
||||
"code": "RATE_LIMITED",
|
||||
"message": fmt.Sprintf("too many requests, retry after %s", retryAfter.Round(1)),
|
||||
})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
196
backend/internal/ratelimit/middleware_test.go
Normal file
196
backend/internal/ratelimit/middleware_test.go
Normal file
@@ -0,0 +1,196 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// mockLimiter 用于测试的 mock 限流器。
|
||||
type mockLimiter struct {
|
||||
allowFunc func(ctx context.Context, key string) (bool, time.Duration)
|
||||
}
|
||||
|
||||
func (m *mockLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
if m.allowFunc != nil {
|
||||
return m.allowFunc(ctx, key)
|
||||
}
|
||||
return true, 0
|
||||
}
|
||||
|
||||
func (m *mockLimiter) Stop() {}
|
||||
|
||||
// 编译期接口检查
|
||||
var _ Limiter = (*mockLimiter)(nil)
|
||||
|
||||
func TestMiddleware_Allow(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
limiter := &mockLimiter{
|
||||
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||
return true, 0
|
||||
},
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||
return "user1:test"
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "ok", resp["status"])
|
||||
}
|
||||
|
||||
func TestMiddleware_Deny(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
limiter := &mockLimiter{
|
||||
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||
return false, 5 * time.Second
|
||||
},
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||
return "user1:test"
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
// 验证返回 429
|
||||
assert.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||
|
||||
// 验证 Retry-After header
|
||||
assert.Equal(t, "5", w.Header().Get("Retry-After"))
|
||||
|
||||
// 验证响应体
|
||||
var resp map[string]interface{}
|
||||
err := json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "RATE_LIMITED", resp["code"])
|
||||
assert.Contains(t, resp["message"], "retry after")
|
||||
}
|
||||
|
||||
func TestMiddleware_NilLimiter(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(nil, func(c *gin.Context) string {
|
||||
return "user1:test"
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
// nil limiter 应该放行
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func TestMiddleware_EmptyKey(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
limiter := &mockLimiter{
|
||||
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||
// 不应该被调用
|
||||
t.Error("Allow should not be called with empty key")
|
||||
return false, 0
|
||||
},
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||
return "" // 返回空 key
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
// 空 key 应该放行
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func TestMiddleware_KeyFunc(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
var capturedKey string
|
||||
limiter := &mockLimiter{
|
||||
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||
capturedKey = key
|
||||
return true, 0
|
||||
},
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||
// 从 query 参数提取 user_id
|
||||
userID := c.Query("user_id")
|
||||
return userID + ":test"
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test?user_id=user123", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
assert.Equal(t, "user123:test", capturedKey)
|
||||
}
|
||||
|
||||
func TestMiddleware_RetryAfterRounding(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
limiter := &mockLimiter{
|
||||
allowFunc: func(ctx context.Context, key string) (bool, time.Duration) {
|
||||
return false, 2500 * time.Millisecond // 2.5 秒
|
||||
},
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(Middleware(limiter, func(c *gin.Context) string {
|
||||
return "user1:test"
|
||||
}))
|
||||
router.GET("/test", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||
// 2.5 秒向上取整为 3 秒
|
||||
assert.Equal(t, "3", w.Header().Get("Retry-After"))
|
||||
}
|
||||
132
backend/internal/ratelimit/redis_bucket.go
Normal file
132
backend/internal/ratelimit/redis_bucket.go
Normal file
@@ -0,0 +1,132 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// luaScript 是 Redis 令牌桶算法的 Lua 脚本。
|
||||
// 保证原子性:读取-计算-回写在一个事务中完成。
|
||||
const luaScript = `
|
||||
-- KEYS[1] = 限流 key
|
||||
-- ARGV[1] = capacity(桶容量)
|
||||
-- ARGV[2] = rate(每秒填充数)
|
||||
-- ARGV[3] = now(当前时间戳,秒,浮点)
|
||||
-- ARGV[4] = ttl(key 过期时间,秒)
|
||||
|
||||
local key = KEYS[1]
|
||||
local capacity = tonumber(ARGV[1])
|
||||
local rate = tonumber(ARGV[2])
|
||||
local now = tonumber(ARGV[3])
|
||||
local ttl = tonumber(ARGV[4])
|
||||
|
||||
local data = redis.call('HMGET', key, 'tokens', 'last_refill')
|
||||
local tokens = tonumber(data[1]) or capacity
|
||||
local last_refill = tonumber(data[2]) or now
|
||||
|
||||
-- 计算新令牌
|
||||
local elapsed = math.max(0, now - last_refill)
|
||||
tokens = math.min(capacity, tokens + elapsed * rate)
|
||||
|
||||
local allowed = 0
|
||||
local retry_after = 0
|
||||
|
||||
if tokens >= 1 then
|
||||
tokens = tokens - 1
|
||||
allowed = 1
|
||||
else
|
||||
if rate == 0 then
|
||||
retry_after = 86400 -- 24小时
|
||||
else
|
||||
retry_after = (1 - tokens) / rate
|
||||
end
|
||||
end
|
||||
|
||||
-- 回写状态
|
||||
redis.call('HMSET', key, 'tokens', tokens, 'last_refill', now)
|
||||
redis.call('EXPIRE', key, ttl)
|
||||
|
||||
return {allowed, tostring(retry_after)}
|
||||
`
|
||||
|
||||
// RedisLimiter Redis 令牌桶限流器。
|
||||
type RedisLimiter struct {
|
||||
client *redis.Client
|
||||
config config.RateLimitConfig
|
||||
script *redis.Script
|
||||
}
|
||||
|
||||
// NewRedisLimiter 创建 Redis 限流器。
|
||||
func NewRedisLimiter(client *redis.Client, cfg config.RateLimitConfig) *RedisLimiter {
|
||||
return &RedisLimiter{
|
||||
client: client,
|
||||
config: cfg,
|
||||
script: redis.NewScript(luaScript),
|
||||
}
|
||||
}
|
||||
|
||||
// Allow 实现 Limiter 接口。
|
||||
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
log := trace.FromContext(ctx)
|
||||
cfg := l.getBucketConfig(key)
|
||||
|
||||
now := float64(time.Now().UnixNano()) / 1e9 // 秒,浮点
|
||||
ttl := 600 // key 过期时间 10 分钟
|
||||
|
||||
result, err := l.script.Run(ctx, l.client, []string{key},
|
||||
cfg.Capacity, cfg.Rate, now, ttl).Result()
|
||||
|
||||
if err != nil {
|
||||
log.Errorw("rate limit check failed", "key", key, "error", err)
|
||||
// Redis 错误时降级:允许请求(fail-open 策略)
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// 解析返回值
|
||||
vals, ok := result.([]interface{})
|
||||
if !ok || len(vals) != 2 {
|
||||
return true, 0
|
||||
}
|
||||
|
||||
allowed, _ := vals[0].(int64)
|
||||
retryAfterStr, _ := vals[1].(string)
|
||||
retryAfterSec, _ := strconv.ParseFloat(retryAfterStr, 64)
|
||||
|
||||
if allowed == 1 {
|
||||
return true, 0
|
||||
}
|
||||
|
||||
retryAfter := time.Duration(retryAfterSec*1000) * time.Millisecond
|
||||
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
|
||||
return false, retryAfter
|
||||
}
|
||||
|
||||
// Stop 实现 Limiter 接口(Redis 不需要清理资源)。
|
||||
func (l *RedisLimiter) Stop() {
|
||||
// Redis 客户端由外部管理,这里不需要操作
|
||||
}
|
||||
|
||||
// getBucketConfig 根据 key 获取桶配置。
|
||||
func (l *RedisLimiter) getBucketConfig(key string) config.BucketConfig {
|
||||
// 简化实现:默认使用 query 配置
|
||||
return l.config.Query
|
||||
}
|
||||
|
||||
// KeyPrefix 返回限流 key 的前缀。
|
||||
func KeyPrefix() string {
|
||||
return "ratelimit:"
|
||||
}
|
||||
|
||||
// FormatKey 格式化限流 key。
|
||||
func FormatKey(userID, action string) string {
|
||||
return fmt.Sprintf("%s%s:%s", KeyPrefix(), userID, action)
|
||||
}
|
||||
|
||||
// 编译期接口检查
|
||||
var _ Limiter = (*RedisLimiter)(nil)
|
||||
228
backend/internal/ratelimit/redis_bucket_test.go
Normal file
228
backend/internal/ratelimit/redis_bucket_test.go
Normal file
@@ -0,0 +1,228 @@
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// setupMiniRedis 创建一个内存 Redis 实例用于测试。
|
||||
func setupMiniRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) {
|
||||
mr, err := miniredis.Run()
|
||||
require.NoError(t, err)
|
||||
|
||||
client := redis.NewClient(&redis.Options{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
|
||||
t.Cleanup(func() {
|
||||
client.Close()
|
||||
mr.Close()
|
||||
})
|
||||
|
||||
return mr, client
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_FirstRequest(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 5, Rate: 0.2},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
|
||||
|
||||
assert.True(t, allowed)
|
||||
assert.Equal(t, time.Duration(0), retryAfter)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_ConsumeUntilEmpty(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 3, Rate: 0.2},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
key := "user1:query"
|
||||
|
||||
// 连续消耗 3 个令牌
|
||||
for i := 0; i < 3; i++ {
|
||||
allowed, _ := limiter.Allow(ctx, key)
|
||||
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||
}
|
||||
|
||||
// 第 4 个请求应被拒绝
|
||||
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||
assert.False(t, allowed)
|
||||
assert.Greater(t, retryAfter, time.Duration(0))
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_DifferentKeys(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 2, Rate: 1.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// user1 消耗 2 个令牌
|
||||
allowed, _ := limiter.Allow(ctx, "user1:query")
|
||||
assert.True(t, allowed)
|
||||
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||
assert.True(t, allowed)
|
||||
|
||||
// user1 第 3 个被拒绝
|
||||
allowed, _ = limiter.Allow(ctx, "user1:query")
|
||||
assert.False(t, allowed)
|
||||
|
||||
// user2 应该不受影响
|
||||
allowed, _ = limiter.Allow(ctx, "user2:query")
|
||||
assert.True(t, allowed)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_RefillAfterWait(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 2, Rate: 10.0}, // 每秒 10 个令牌
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
key := "user1:query"
|
||||
|
||||
// 消耗 2 个令牌
|
||||
limiter.Allow(ctx, key)
|
||||
limiter.Allow(ctx, key)
|
||||
|
||||
// 真实等待 150ms(Lua 脚本使用系统时间)
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
|
||||
// 应该补充了至少 1 个令牌
|
||||
allowed, _ := limiter.Allow(ctx, key)
|
||||
assert.True(t, allowed)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_CapacityLimit(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 3, Rate: 1.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
key := "user1:query"
|
||||
|
||||
// 真实等待让桶"溢出"
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// 但最多只能消耗 capacity 个令牌
|
||||
for i := 0; i < 3; i++ {
|
||||
allowed, _ := limiter.Allow(ctx, key)
|
||||
assert.True(t, allowed, "request %d should be allowed", i+1)
|
||||
}
|
||||
|
||||
// 第 4 个应被拒绝
|
||||
allowed, _ := limiter.Allow(ctx, key)
|
||||
assert.False(t, allowed)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_ZeroRate(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 1, Rate: 0.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
key := "user1:query"
|
||||
|
||||
// 第一个通过
|
||||
allowed, _ := limiter.Allow(ctx, key)
|
||||
assert.True(t, allowed)
|
||||
|
||||
// 第二个被拒绝,retryAfter 应该很大
|
||||
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||
assert.False(t, allowed)
|
||||
assert.Greater(t, retryAfter, 1*time.Hour)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_KeyTTL(t *testing.T) {
|
||||
mr, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 5, Rate: 1.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
key := "user1:query"
|
||||
|
||||
// 第一次请求
|
||||
limiter.Allow(ctx, key)
|
||||
|
||||
// 验证 key 已设置 TTL
|
||||
ttl := mr.TTL(key)
|
||||
assert.Greater(t, ttl, time.Duration(0))
|
||||
assert.LessOrEqual(t, ttl, 600*time.Second)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Allow_FailOpen(t *testing.T) {
|
||||
mr, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 关闭 Redis 模拟故障
|
||||
mr.Close()
|
||||
|
||||
// 应该 fail-open(允许请求)
|
||||
allowed, retryAfter := limiter.Allow(ctx, "user1:query")
|
||||
assert.True(t, allowed)
|
||||
assert.Equal(t, time.Duration(0), retryAfter)
|
||||
}
|
||||
|
||||
func TestRedisLimiter_Stop(t *testing.T) {
|
||||
_, client := setupMiniRedis(t)
|
||||
|
||||
cfg := config.RateLimitConfig{
|
||||
Enabled: true,
|
||||
Query: config.BucketConfig{Capacity: 1, Rate: 1.0},
|
||||
}
|
||||
limiter := NewRedisLimiter(client, cfg)
|
||||
|
||||
// Stop 应该不会 panic(即使多次调用)
|
||||
limiter.Stop()
|
||||
limiter.Stop()
|
||||
}
|
||||
|
||||
func TestFormatKey(t *testing.T) {
|
||||
key := FormatKey("user123", "query")
|
||||
assert.Equal(t, "ratelimit:user123:query", key)
|
||||
}
|
||||
66
backend/internal/session/manager.go
Normal file
66
backend/internal/session/manager.go
Normal file
@@ -0,0 +1,66 @@
|
||||
// Package session 提供会话生命周期管理能力。
|
||||
package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
// 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。userID 为空表示匿名会话。
|
||||
Create(ctx context.Context, userID string, config models.SessionConfig) (string, error)
|
||||
|
||||
// Get 获取会话(含 config)。不存在返回 ErrSessionNotFound。
|
||||
Get(ctx context.Context, sessionID string) (*models.Session, error)
|
||||
|
||||
// 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)
|
||||
|
||||
// AppendMessage 追加一条对话消息,同时刷新 TTL。
|
||||
AppendMessage(ctx context.Context, sessionID string, msg models.Message) error
|
||||
|
||||
// SetActiveRequest 标记当前正在处理的请求 ID(interrupt 用)。
|
||||
SetActiveRequest(ctx context.Context, sessionID string, requestID string) error
|
||||
|
||||
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||
GetActiveRequestID(ctx context.Context, sessionID string) (string, error)
|
||||
|
||||
// ClearActiveRequest 清除活跃请求标记(请求完成或中断后)。
|
||||
ClearActiveRequest(ctx context.Context, sessionID string) error
|
||||
|
||||
// Touch 刷新 TTL(心跳时调用)。
|
||||
Touch(ctx context.Context, sessionID string) error
|
||||
|
||||
// Destroy 显式销毁会话(REST API DELETE 或连接断开清理)。
|
||||
Destroy(ctx context.Context, sessionID string) error
|
||||
|
||||
// ActiveCount 返回当前活跃会话数(健康检查用)。
|
||||
ActiveCount() int
|
||||
}
|
||||
607
backend/internal/session/memory.go
Normal file
607
backend/internal/session/memory.go
Normal file
@@ -0,0 +1,607 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTTL = 30 * time.Minute
|
||||
defaultHistorySize = 20
|
||||
)
|
||||
|
||||
// sessionEntry 内部会话条目。
|
||||
type sessionEntry struct {
|
||||
session models.Session
|
||||
history []models.Message
|
||||
activeReqID string
|
||||
lastActive time.Time
|
||||
}
|
||||
|
||||
// MemoryManager 基于内存的 SessionManager 实现。
|
||||
// 适用于 MVP 和无 Redis 的开发环境。
|
||||
type MemoryManager struct {
|
||||
mu sync.RWMutex
|
||||
sessions map[string]*sessionEntry
|
||||
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)。
|
||||
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。
|
||||
func NewMemoryManager(ttl time.Duration, maxHistory int, opts ...Option) *MemoryManager {
|
||||
if ttl <= 0 {
|
||||
ttl = defaultTTL
|
||||
}
|
||||
if maxHistory <= 0 {
|
||||
maxHistory = defaultHistorySize
|
||||
}
|
||||
|
||||
m := &MemoryManager{
|
||||
sessions: make(map[string]*sessionEntry),
|
||||
ttl: ttl,
|
||||
maxHistory: maxHistory,
|
||||
stopCleaner: make(chan struct{}),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(m)
|
||||
}
|
||||
|
||||
// 启动后台清理 goroutine,每分钟清除过期会话。
|
||||
go m.cleanLoop()
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// cleanLoop 后台定期清理过期会话。
|
||||
func (m *MemoryManager) cleanLoop() {
|
||||
ticker := time.NewTicker(1 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
m.cleanExpired()
|
||||
case <-m.stopCleaner:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// cleanExpired 清除所有过期会话。
|
||||
func (m *MemoryManager) cleanExpired() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for id, entry := range m.sessions {
|
||||
if now.Sub(entry.lastActive) > m.ttl {
|
||||
delete(m.sessions, id)
|
||||
logger.Log.Debugw("session expired (cleaner)", "session", id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Stop 停止后台清理 goroutine。应用退出前调用。
|
||||
func (m *MemoryManager) Stop() {
|
||||
close(m.stopCleaner)
|
||||
}
|
||||
|
||||
// isExpired 检查会话是否过期(调用方需持锁或在已知 entry 存在时调用)。
|
||||
func (m *MemoryManager) isExpired(entry *sessionEntry) bool {
|
||||
return time.Since(entry.lastActive) > m.ttl
|
||||
}
|
||||
|
||||
// Create 创建新会话。userID 为空表示匿名会话。
|
||||
func (m *MemoryManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
m.mu.Lock()
|
||||
|
||||
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()
|
||||
|
||||
// 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 获取会话。内存中不存在时,尝试从 PG 加载(透明恢复)。
|
||||
func (m *MemoryManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||
m.mu.RLock()
|
||||
entry, ok := m.sessions[sessionID]
|
||||
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
|
||||
}
|
||||
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
// UpdateConfig 更新会话配置。
|
||||
func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
||||
m.mu.Lock()
|
||||
|
||||
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()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
if limit <= 0 || limit > len(entry.history) {
|
||||
limit = len(entry.history)
|
||||
}
|
||||
|
||||
// 返回最近 limit 条的副本
|
||||
result := make([]models.Message, limit)
|
||||
copy(result, entry.history[len(entry.history)-limit:])
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// AppendMessage 追加一条对话消息,同时刷新 TTL。
|
||||
// 若配置了 MessageRepository,消息会异步写入 PostgreSQL(Write-Through)。
|
||||
func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg models.Message) error {
|
||||
m.mu.Lock()
|
||||
|
||||
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:]
|
||||
}
|
||||
|
||||
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()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
entry.activeReqID = requestID
|
||||
entry.lastActive = time.Now()
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||
func (m *MemoryManager) GetActiveRequestID(_ context.Context, sessionID string) (string, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return "", ErrSessionNotFound
|
||||
}
|
||||
|
||||
return entry.activeReqID, nil
|
||||
}
|
||||
|
||||
// ClearActiveRequest 清除活跃请求标记。
|
||||
func (m *MemoryManager) ClearActiveRequest(_ context.Context, sessionID string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
entry.activeReqID = ""
|
||||
entry.lastActive = time.Now()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Touch 刷新 TTL。
|
||||
func (m *MemoryManager) Touch(_ context.Context, sessionID string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
entry.lastActive = time.Now()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Destroy 显式销毁会话。
|
||||
func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
|
||||
m.mu.Lock()
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// ActiveCount 返回当前活跃会话数。
|
||||
func (m *MemoryManager) ActiveCount() int {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
now := time.Now()
|
||||
count := 0
|
||||
for _, entry := range m.sessions {
|
||||
if now.Sub(entry.lastActive) <= m.ttl {
|
||||
count++
|
||||
}
|
||||
}
|
||||
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,
|
||||
}
|
||||
}
|
||||
481
backend/internal/session/memory_test.go
Normal file
481
backend/internal/session/memory_test.go
Normal file
@@ -0,0 +1,481 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
func init() {
|
||||
logger.Init("debug", "console")
|
||||
}
|
||||
|
||||
func TestCreateAndGet(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
config := models.DefaultConfig()
|
||||
id, err := m.Create(ctx, "", config)
|
||||
if err != nil {
|
||||
t.Fatalf("Create: %v", err)
|
||||
}
|
||||
if id == "" {
|
||||
t.Fatal("Create returned empty ID")
|
||||
}
|
||||
|
||||
sess, err := m.Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
if sess.ID != id {
|
||||
t.Errorf("ID = %q, want %q", sess.ID, id)
|
||||
}
|
||||
if sess.Config.Language != "zh-CN" {
|
||||
t.Errorf("Language = %q, want %q", sess.Config.Language, "zh-CN")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetNotFound(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := m.Get(ctx, "nonexistent")
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("Get nonexistent: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpire(t *testing.T) {
|
||||
// 使用极短 TTL 测试过期
|
||||
m := NewMemoryManager(50*time.Millisecond, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 未过期时应能获取
|
||||
_, err := m.Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("Get before expire: %v", err)
|
||||
}
|
||||
|
||||
// 等待过期
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
|
||||
_, err = m.Get(ctx, id)
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("Get after expire: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDestroy(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
if err := m.Destroy(ctx, id); err != nil {
|
||||
t.Fatalf("Destroy: %v", err)
|
||||
}
|
||||
|
||||
_, err := m.Get(ctx, id)
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("Get after Destroy: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDestroyNotFound(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
err := m.Destroy(ctx, "nonexistent")
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("Destroy nonexistent: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendMessageAndGetHistory(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
msgs := []models.Message{
|
||||
{Role: "user", Content: "你好"},
|
||||
{Role: "assistant", Content: "你好!有什么可以帮你的吗?"},
|
||||
{Role: "user", Content: "这是什么?"},
|
||||
{Role: "assistant", Content: "这是一朵花。"},
|
||||
}
|
||||
|
||||
for _, msg := range msgs {
|
||||
if err := m.AppendMessage(ctx, id, msg); err != nil {
|
||||
t.Fatalf("AppendMessage: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
history, err := m.GetHistory(ctx, id, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("GetHistory: %v", err)
|
||||
}
|
||||
if len(history) != 4 {
|
||||
t.Fatalf("GetHistory len = %d, want 4", len(history))
|
||||
}
|
||||
if history[0].Content != "你好" {
|
||||
t.Errorf("history[0] = %q, want %q", history[0].Content, "你好")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetHistoryLimit(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
|
||||
}
|
||||
|
||||
history, err := m.GetHistory(ctx, id, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("GetHistory: %v", err)
|
||||
}
|
||||
if len(history) != 3 {
|
||||
t.Fatalf("GetHistory limit=3: len = %d, want 3", len(history))
|
||||
}
|
||||
}
|
||||
|
||||
func TestHistoryLimit(t *testing.T) {
|
||||
const maxHistory = 5
|
||||
m := NewMemoryManager(30*time.Minute, maxHistory)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 插入超过上限的消息
|
||||
for i := 0; i < 10; i++ {
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
|
||||
}
|
||||
|
||||
history, err := m.GetHistory(ctx, id, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("GetHistory: %v", err)
|
||||
}
|
||||
if len(history) != maxHistory {
|
||||
t.Fatalf("GetHistory after overflow: len = %d, want %d", len(history), maxHistory)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateConfig(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
ttsEnabled := false
|
||||
detailLevel := "high"
|
||||
patch := models.SessionConfigPatch{
|
||||
TTSEnabled: &ttsEnabled,
|
||||
DetailLevel: &detailLevel,
|
||||
}
|
||||
|
||||
if err := m.UpdateConfig(ctx, id, patch); err != nil {
|
||||
t.Fatalf("UpdateConfig: %v", err)
|
||||
}
|
||||
|
||||
sess, _ := m.Get(ctx, id)
|
||||
if sess.Config.TTSEnabled != false {
|
||||
t.Errorf("TTSEnabled = %v, want false", sess.Config.TTSEnabled)
|
||||
}
|
||||
if sess.Config.DetailLevel != "high" {
|
||||
t.Errorf("DetailLevel = %q, want %q", sess.Config.DetailLevel, "high")
|
||||
}
|
||||
// Language 未传,应保持原值
|
||||
if sess.Config.Language != "zh-CN" {
|
||||
t.Errorf("Language = %q, want %q", sess.Config.Language, "zh-CN")
|
||||
}
|
||||
}
|
||||
|
||||
func TestActiveRequest(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 初始应为空
|
||||
reqID, err := m.GetActiveRequestID(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("GetActiveRequestID: %v", err)
|
||||
}
|
||||
if reqID != "" {
|
||||
t.Errorf("initial active request = %q, want empty", reqID)
|
||||
}
|
||||
|
||||
// 设置
|
||||
if err := m.SetActiveRequest(ctx, id, "req-123"); err != nil {
|
||||
t.Fatalf("SetActiveRequest: %v", err)
|
||||
}
|
||||
reqID, _ = m.GetActiveRequestID(ctx, id)
|
||||
if reqID != "req-123" {
|
||||
t.Errorf("active request = %q, want %q", reqID, "req-123")
|
||||
}
|
||||
|
||||
// 清除
|
||||
if err := m.ClearActiveRequest(ctx, id); err != nil {
|
||||
t.Fatalf("ClearActiveRequest: %v", err)
|
||||
}
|
||||
reqID, _ = m.GetActiveRequestID(ctx, id)
|
||||
if reqID != "" {
|
||||
t.Errorf("active request after clear = %q, want empty", reqID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTouchRefreshesTTL(t *testing.T) {
|
||||
m := NewMemoryManager(100*time.Millisecond, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 50ms 后 Touch,应重置 TTL
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
if err := m.Touch(ctx, id); err != nil {
|
||||
t.Fatalf("Touch: %v", err)
|
||||
}
|
||||
|
||||
// 再等 70ms(距创建 120ms,但距 Touch 只有 70ms),不应过期
|
||||
time.Sleep(70 * time.Millisecond)
|
||||
_, err := m.Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Errorf("Get after Touch: %v, want nil (should not expire yet)", err)
|
||||
}
|
||||
|
||||
// 再等 50ms(距 Touch 120ms),应过期
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
_, err = m.Get(ctx, id)
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("Get after TTL: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActiveCount(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
if m.ActiveCount() != 0 {
|
||||
t.Errorf("initial ActiveCount = %d, want 0", m.ActiveCount())
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
}
|
||||
487
backend/internal/session/redis.go
Normal file
487
backend/internal/session/redis.go
Normal file
@@ -0,0 +1,487 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"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
|
||||
maxHistory int
|
||||
}
|
||||
|
||||
// NewRedisManager 创建 Redis 版 SessionManager。
|
||||
func NewRedisManager(rdb *redis.Client, ttl time.Duration, maxHistory int) *RedisManager {
|
||||
if ttl <= 0 {
|
||||
ttl = defaultTTL
|
||||
}
|
||||
if maxHistory <= 0 {
|
||||
maxHistory = defaultHistorySize
|
||||
}
|
||||
return &RedisManager{rdb: rdb, ttl: ttl, maxHistory: maxHistory}
|
||||
}
|
||||
|
||||
// Ping 检查 Redis 连接是否正常。
|
||||
func (m *RedisManager) Ping(ctx context.Context) error {
|
||||
return m.rdb.Ping(ctx).Err()
|
||||
}
|
||||
|
||||
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
|
||||
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
|
||||
func userSessKey(id string) string { return fmt.Sprintf("user:%s:sessions", id) }
|
||||
|
||||
// Create 创建新会话。userID 为空表示匿名会话。
|
||||
func (m *RedisManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
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
|
||||
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),
|
||||
"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)
|
||||
}
|
||||
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session created", "session_id", id, "user_id", userID)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// placeholderHistoryMark 占位符,避免 Redis 对空 key 的特殊行为。
|
||||
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 {
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
sess := &models.Session{
|
||||
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
|
||||
}
|
||||
|
||||
// UpdateConfig 更新会话配置。
|
||||
func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) 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)
|
||||
fields := map[string]interface{}{
|
||||
"last_active": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
if patch.TTSEnabled != nil {
|
||||
fields["config.tts_enabled"] = strconv.FormatBool(*patch.TTSEnabled)
|
||||
}
|
||||
if patch.DetailLevel != nil {
|
||||
fields["config.detail_level"] = *patch.DetailLevel
|
||||
}
|
||||
if patch.Language != nil {
|
||||
fields["config.language"] = *patch.Language
|
||||
}
|
||||
|
||||
if err := m.rdb.HSet(ctx, metaKey(sessionID), fields).Err(); err != nil {
|
||||
return fmt.Errorf("redis update config: %w", err)
|
||||
}
|
||||
|
||||
// 刷新 TTL
|
||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
|
||||
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) {
|
||||
// 检查会话是否存在
|
||||
exists, err := m.rdb.Exists(ctx, metaKey(sessionID)).Result()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("redis check session: %w", err)
|
||||
}
|
||||
if exists == 0 {
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
if limit <= 0 {
|
||||
limit = m.maxHistory
|
||||
}
|
||||
|
||||
// LRANGE 0 {limit-1},最新在前(LPUSH),需要反转为时间顺序
|
||||
raws, err := m.rdb.LRange(ctx, histKey(sessionID), 0, int64(limit)).Result()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("redis get history: %w", err)
|
||||
}
|
||||
|
||||
var msgs []models.Message
|
||||
for _, raw := range raws {
|
||||
if raw == placeholderHistoryMark {
|
||||
continue
|
||||
}
|
||||
var msg models.Message
|
||||
if err := json.Unmarshal([]byte(raw), &msg); err != nil {
|
||||
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)
|
||||
}
|
||||
|
||||
// 反转为时间顺序(LPUSH 最新在前 → 需要最旧在前)
|
||||
for i, j := 0, len(msgs)-1; i < j; i, j = i+1, j-1 {
|
||||
msgs[i], msgs[j] = msgs[j], msgs[i]
|
||||
}
|
||||
|
||||
return msgs, nil
|
||||
}
|
||||
|
||||
// AppendMessage 追加一条对话消息,同时刷新 TTL。
|
||||
func (m *RedisManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) 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
|
||||
}
|
||||
|
||||
data, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal message: %w", err)
|
||||
}
|
||||
|
||||
pipe := m.rdb.Pipeline()
|
||||
// LPUSH 新消息到左头(最新在前)
|
||||
pipe.LPush(ctx, histKey(sessionID), string(data))
|
||||
// LTRIM 保留最近 maxHistory 条(+1 是因为有占位符)
|
||||
pipe.LTrim(ctx, histKey(sessionID), 0, int64(m.maxHistory))
|
||||
// 刷新 TTL
|
||||
pipe.Expire(ctx, histKey(sessionID), m.ttl)
|
||||
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetActiveRequest 标记当前正在处理的请求 ID。
|
||||
func (m *RedisManager) SetActiveRequest(ctx context.Context, sessionID string, requestID 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
|
||||
}
|
||||
|
||||
pipe := m.rdb.Pipeline()
|
||||
pipe.HSet(ctx, metaKey(sessionID), "active_request_id", requestID)
|
||||
pipe.HSet(ctx, metaKey(sessionID), "last_active", time.Now().UTC().Format(time.RFC3339))
|
||||
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
return fmt.Errorf("redis set active request: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||
func (m *RedisManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
|
||||
val, err := m.rdb.HGet(ctx, metaKey(sessionID), "active_request_id").Result()
|
||||
if err == redis.Nil {
|
||||
return "", ErrSessionNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("redis get active request: %w", err)
|
||||
}
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// ClearActiveRequest 清除活跃请求标记。
|
||||
func (m *RedisManager) ClearActiveRequest(ctx context.Context, sessionID 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
|
||||
}
|
||||
|
||||
pipe := m.rdb.Pipeline()
|
||||
pipe.HSet(ctx, metaKey(sessionID), "active_request_id", "")
|
||||
pipe.HSet(ctx, metaKey(sessionID), "last_active", time.Now().UTC().Format(time.RFC3339))
|
||||
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
return fmt.Errorf("redis clear active request: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Touch 刷新 TTL。
|
||||
func (m *RedisManager) Touch(ctx context.Context, sessionID 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
|
||||
}
|
||||
|
||||
pipe := m.rdb.Pipeline()
|
||||
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
pipe.Expire(ctx, histKey(sessionID), m.ttl)
|
||||
pipe.HSet(ctx, metaKey(sessionID), "last_active", time.Now().UTC().Format(time.RFC3339))
|
||||
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
return fmt.Errorf("redis touch: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
if deleted == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
// 清理用户会话索引
|
||||
if userID != "" {
|
||||
m.rdb.SRem(ctx, userSessKey(userID), sessionID)
|
||||
}
|
||||
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session destroyed", "session_id", sessionID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ActiveCount 返回当前活跃会话数。
|
||||
// Redis 实现通过 SCAN 遍历 meta key,适用于中等规模。
|
||||
// 大规模部署建议维护独立的活跃会话集合。
|
||||
func (m *RedisManager) ActiveCount() int {
|
||||
ctx := context.Background()
|
||||
count := 0
|
||||
var cursor uint64
|
||||
for {
|
||||
keys, nextCursor, err := m.rdb.Scan(ctx, cursor, "session:*:meta", 100).Result()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
for _, key := range keys {
|
||||
exists, _ := m.rdb.Exists(ctx, key).Result()
|
||||
if exists > 0 {
|
||||
count++
|
||||
}
|
||||
}
|
||||
cursor = nextCursor
|
||||
if cursor == 0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// uuidNew 生成 UUID,便于测试时 mock。
|
||||
var uuidNew = func() string {
|
||||
return uuid.New().String()
|
||||
}
|
||||
361
backend/internal/session/tiered.go
Normal file
361
backend/internal/session/tiered.go
Normal file
@@ -0,0 +1,361 @@
|
||||
package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
)
|
||||
|
||||
// TieredManager 三级存储 SessionManager 实现。
|
||||
//
|
||||
// L1(内存)→ L2(Redis)→ L3(PostgreSQL)
|
||||
//
|
||||
// 读:L1 miss → L2 miss → L3,回填到 L1+L2
|
||||
// 写:L1 → L2(同步)→ L3(异步)
|
||||
// 降级:Redis 不可用时,回退到 L1+L3 模式
|
||||
type TieredManager struct {
|
||||
l1 *MemoryManager // L1: 内存缓存
|
||||
l2 *RedisManager // L2: Redis(可选)
|
||||
sessRepo store.SessionRepository // L3: PostgreSQL 会话持久化(可选)
|
||||
msgRepo store.MessageRepository // L3: PostgreSQL 消息持久化(可选)
|
||||
|
||||
redisOK atomic.Bool // Redis 健康状态
|
||||
stopCh chan struct{} // 停止信号
|
||||
}
|
||||
|
||||
// TieredOption TieredManager 的函数式选项。
|
||||
type TieredOption func(*TieredManager)
|
||||
|
||||
// WithTieredSessionRepository 注入 L3 会话持久化仓库。
|
||||
func WithTieredSessionRepository(repo store.SessionRepository) TieredOption {
|
||||
return func(m *TieredManager) {
|
||||
m.sessRepo = repo
|
||||
}
|
||||
}
|
||||
|
||||
// WithTieredMessageRepository 注入 L3 消息持久化仓库。
|
||||
func WithTieredMessageRepository(repo store.MessageRepository) TieredOption {
|
||||
return func(m *TieredManager) {
|
||||
m.msgRepo = repo
|
||||
}
|
||||
}
|
||||
|
||||
// NewTieredManager 创建三级存储 SessionManager。
|
||||
// l2 为 nil 时降级为 L1+L3 模式。
|
||||
func NewTieredManager(
|
||||
ttl time.Duration,
|
||||
maxHistory int,
|
||||
l2 *RedisManager,
|
||||
opts ...TieredOption,
|
||||
) *TieredManager {
|
||||
m := &TieredManager{
|
||||
l2: l2,
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(m)
|
||||
}
|
||||
|
||||
// 初始化 L1(内存),注入 L3 仓库实现 Write-Through
|
||||
var l1Opts []Option
|
||||
if m.sessRepo != nil {
|
||||
l1Opts = append(l1Opts, WithSessionRepository(m.sessRepo))
|
||||
}
|
||||
if m.msgRepo != nil {
|
||||
l1Opts = append(l1Opts, WithMessageRepository(m.msgRepo))
|
||||
}
|
||||
m.l1 = NewMemoryManager(ttl, maxHistory, l1Opts...)
|
||||
|
||||
// 初始化 Redis 健康状态
|
||||
if l2 != nil {
|
||||
m.redisOK.Store(true)
|
||||
go m.healthCheck()
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// healthCheck 定期检查 Redis 健康状态。
|
||||
func (m *TieredManager) healthCheck() {
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
err := m.l2.Ping(ctx)
|
||||
cancel()
|
||||
|
||||
wasOK := m.redisOK.Load()
|
||||
isOK := err == nil
|
||||
m.redisOK.Store(isOK)
|
||||
|
||||
if wasOK && !isOK {
|
||||
logger.Log.Warn("Redis connection lost, degrading to L1+L3 mode")
|
||||
} else if !wasOK && isOK {
|
||||
logger.Log.Info("Redis connection restored, resuming L1+L2+L3 mode")
|
||||
}
|
||||
case <-m.stopCh:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// isRedisOK 检查 Redis 是否可用。
|
||||
func (m *TieredManager) isRedisOK() bool {
|
||||
return m.l2 != nil && m.redisOK.Load()
|
||||
}
|
||||
|
||||
// Stop 停止 TieredManager(清理后台 goroutine)。
|
||||
func (m *TieredManager) Stop() {
|
||||
close(m.stopCh)
|
||||
m.l1.Stop()
|
||||
}
|
||||
|
||||
// Create 创建新会话。
|
||||
// 写入:L1 → L2(同步)→ L3(异步)
|
||||
func (m *TieredManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
// L1: 内存
|
||||
id, err := m.l1.Create(ctx, userID, config)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// L2: Redis(同步),使用 L1 生成的 ID 保证一致性
|
||||
if m.isRedisOK() {
|
||||
if _, err := m.l2.CreateWithID(ctx, id, userID, config); err != nil {
|
||||
logger.Log.Warnw("Redis Create failed, continuing without L2",
|
||||
"session", id, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// Get 获取会话。
|
||||
// 读取:L1 → L2(回填 L1)→ L3(回填 L1+L2)
|
||||
func (m *TieredManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||
// L1: 内存
|
||||
sess, err := m.l1.Get(ctx, sessionID)
|
||||
if err == nil {
|
||||
return sess, nil
|
||||
}
|
||||
if err != ErrSessionNotFound {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// L2: Redis
|
||||
if m.isRedisOK() {
|
||||
sess, err = m.l2.Get(ctx, sessionID)
|
||||
if err == nil {
|
||||
// 回填 L1
|
||||
history, _ := m.l2.GetHistory(ctx, sessionID, 0)
|
||||
m.l1.LoadSession(sess, history)
|
||||
return sess, nil
|
||||
}
|
||||
if err != ErrSessionNotFound {
|
||||
logger.Log.Warnw("Redis Get failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// L3: PostgreSQL(由 L1 的 Cache-Aside 处理)
|
||||
// L1.Get 已经实现了从 PostgreSQL 恢复的逻辑
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
// UpdateConfig 更新会话配置。
|
||||
// 写入:L1 → L2(同步)→ L3(异步)
|
||||
func (m *TieredManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
||||
// L1: 内存
|
||||
if err := m.l1.UpdateConfig(ctx, sessionID, patch); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.UpdateConfig(ctx, sessionID, patch); err != nil {
|
||||
logger.Log.Warnw("Redis UpdateConfig failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateTitle 更新会话标题。
|
||||
// 写入:L1 → L2(同步)→ L3(异步)
|
||||
func (m *TieredManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
||||
// L1: 内存
|
||||
if err := m.l1.UpdateTitle(ctx, sessionID, title); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.UpdateTitle(ctx, sessionID, title); err != nil {
|
||||
logger.Log.Warnw("Redis UpdateTitle failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// L3: PostgreSQL(异步,由 L1 的 Write-Through 处理)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListByUser 查询用户的会话列表。
|
||||
// 读取:L1 + L2 + L3 合并去重
|
||||
func (m *TieredManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
||||
// 优先使用 L1(已集成 L3 回退逻辑)
|
||||
return m.l1.ListByUser(ctx, userID, page, size)
|
||||
}
|
||||
|
||||
// GetHistory 获取对话历史。
|
||||
// 读取:L1 → L2(回填 L1)→ L3(回填 L1+L2)
|
||||
func (m *TieredManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||
// L1: 内存
|
||||
msgs, err := m.l1.GetHistory(ctx, sessionID, limit)
|
||||
if err == nil && len(msgs) > 0 {
|
||||
return msgs, nil
|
||||
}
|
||||
|
||||
// L2: Redis
|
||||
if m.isRedisOK() {
|
||||
msgs, err = m.l2.GetHistory(ctx, sessionID, limit)
|
||||
if err == nil && len(msgs) > 0 {
|
||||
// 回填 L1(通过 Get 触发)
|
||||
m.l1.Get(ctx, sessionID)
|
||||
return msgs, nil
|
||||
}
|
||||
}
|
||||
|
||||
// L3: PostgreSQL(由 L1 的 Cache-Aside 处理)
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
|
||||
// AppendMessage 追加消息。
|
||||
// 写入:L1 → L2(同步)→ L3(异步)
|
||||
func (m *TieredManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
|
||||
// L1: 内存(Write-Through 到 L3)
|
||||
if err := m.l1.AppendMessage(ctx, sessionID, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.AppendMessage(ctx, sessionID, msg); err != nil {
|
||||
logger.Log.Warnw("Redis AppendMessage failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetActiveRequest 设置当前活跃请求。
|
||||
func (m *TieredManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
|
||||
// L1: 内存
|
||||
if err := m.l1.SetActiveRequest(ctx, sessionID, requestID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.SetActiveRequest(ctx, sessionID, requestID); err != nil {
|
||||
logger.Log.Warnw("Redis SetActiveRequest failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||
func (m *TieredManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
|
||||
// L1: 内存
|
||||
id, err := m.l1.GetActiveRequestID(ctx, sessionID)
|
||||
if err == nil && id != "" {
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// L2: Redis
|
||||
if m.isRedisOK() {
|
||||
id, err = m.l2.GetActiveRequestID(ctx, sessionID)
|
||||
if err == nil && id != "" {
|
||||
return id, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ClearActiveRequest 清除当前活跃请求。
|
||||
func (m *TieredManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
|
||||
// L1: 内存
|
||||
if err := m.l1.ClearActiveRequest(ctx, sessionID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.ClearActiveRequest(ctx, sessionID); err != nil {
|
||||
logger.Log.Warnw("Redis ClearActiveRequest failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Touch 刷新会话活跃时间。
|
||||
func (m *TieredManager) Touch(ctx context.Context, sessionID string) error {
|
||||
// L1: 内存
|
||||
if err := m.l1.Touch(ctx, sessionID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.Touch(ctx, sessionID); err != nil {
|
||||
logger.Log.Warnw("Redis Touch failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Destroy 销毁会话。
|
||||
// 写入:L1 → L2 → L3
|
||||
func (m *TieredManager) Destroy(ctx context.Context, sessionID string) error {
|
||||
// L1: 内存(Write-Through 到 L3)
|
||||
if err := m.l1.Destroy(ctx, sessionID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// L2: Redis(同步)
|
||||
if m.isRedisOK() {
|
||||
if err := m.l2.Destroy(ctx, sessionID); err != nil {
|
||||
logger.Log.Warnw("Redis Destroy failed",
|
||||
"session", sessionID, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ActiveCount 返回活跃会话数量。
|
||||
func (m *TieredManager) ActiveCount() int {
|
||||
return m.l1.ActiveCount()
|
||||
}
|
||||
172
backend/internal/store/cached_user.go
Normal file
172
backend/internal/store/cached_user.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// Redis key 前缀。
|
||||
const (
|
||||
refreshTokenPrefix = "auth:refresh:" // auth:refresh:{token_hash} → user_id
|
||||
userRefreshPrefix = "auth:user_refresh:" // auth:user_refresh:{user_id} → Set of token_hash
|
||||
)
|
||||
|
||||
// CachedUserRepository 装饰器,为 UserRepository 的 refresh token 操作增加 Redis 缓存。
|
||||
// 读路径:Redis miss → DB → 回填 Redis。
|
||||
// 写路径:同步双写 Redis + DB。
|
||||
// 删路径:同步双删 Redis + DB。
|
||||
// Redis 操作失败时降级到纯 DB,不阻断主流程。
|
||||
type CachedUserRepository struct {
|
||||
inner UserRepository
|
||||
rdb *redis.Client
|
||||
backfillTTL time.Duration // DB 回填 Redis 时使用的默认 TTL
|
||||
}
|
||||
|
||||
// NewCachedUserRepository 创建带 Redis 缓存的 UserRepository 装饰器。
|
||||
// backfillTTL: 从 DB 回填 Redis 时使用的 TTL(因 DB 接口不返回 expiresAt)。
|
||||
func NewCachedUserRepository(inner UserRepository, rdb *redis.Client, backfillTTL time.Duration) *CachedUserRepository {
|
||||
if backfillTTL <= 0 {
|
||||
backfillTTL = 24 * time.Hour
|
||||
}
|
||||
return &CachedUserRepository{
|
||||
inner: inner,
|
||||
rdb: rdb,
|
||||
backfillTTL: backfillTTL,
|
||||
}
|
||||
}
|
||||
|
||||
// refreshTokenKey 生成 refresh token 的 Redis key。
|
||||
func refreshTokenKey(tokenHash string) string {
|
||||
return refreshTokenPrefix + tokenHash
|
||||
}
|
||||
|
||||
// userRefreshKey 生成用户 refresh token 集合的 Redis key。
|
||||
func userRefreshKey(userID string) string {
|
||||
return userRefreshPrefix + userID
|
||||
}
|
||||
|
||||
// --- 委托方法(不做缓存) ---
|
||||
|
||||
func (r *CachedUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
|
||||
return r.inner.Create(ctx, username, passwordHash)
|
||||
}
|
||||
|
||||
func (r *CachedUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
|
||||
return r.inner.FindByUsername(ctx, username)
|
||||
}
|
||||
|
||||
func (r *CachedUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
|
||||
return r.inner.FindByID(ctx, id)
|
||||
}
|
||||
|
||||
// --- 缓存方法 ---
|
||||
|
||||
// SaveRefreshToken Write-Through:先写 DB,再写 Redis。
|
||||
func (r *CachedUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||
// 先写 DB
|
||||
if err := r.inner.SaveRefreshToken(ctx, userID, tokenHash, expiresAt); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 写 Redis(SET + SADD),设置 TTL 为 token 剩余有效期
|
||||
ttl := time.Until(expiresAt)
|
||||
if ttl <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
key := refreshTokenKey(tokenHash)
|
||||
pipe := r.rdb.Pipeline()
|
||||
pipe.Set(ctx, key, userID, ttl)
|
||||
pipe.SAdd(ctx, userRefreshKey(userID), tokenHash)
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache write failed for refresh token", "error", err)
|
||||
// 降级:DB 已写入成功,Redis 失败不影响正确性
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindRefreshToken Read-Through:先查 Redis,miss 时查 DB 并回填。
|
||||
func (r *CachedUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
key := refreshTokenKey(tokenHash)
|
||||
|
||||
// 查 Redis
|
||||
userID, err := r.rdb.Get(ctx, key).Result()
|
||||
if err == nil {
|
||||
return userID, nil
|
||||
}
|
||||
// redis.Nil 表示 key 不存在,其他错误记录日志后降级到 DB
|
||||
if err != redis.Nil {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache read failed for refresh token", "error", err)
|
||||
}
|
||||
|
||||
// 降级到 DB
|
||||
userID, err = r.inner.FindRefreshToken(ctx, tokenHash)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 回填 Redis(SET + SADD),TTL 使用保守默认值
|
||||
go func() {
|
||||
bgCtx := context.Background()
|
||||
pipe := r.rdb.Pipeline()
|
||||
pipe.Set(bgCtx, key, userID, r.backfillTTL)
|
||||
pipe.SAdd(bgCtx, userRefreshKey(userID), tokenHash)
|
||||
_, _ = pipe.Exec(bgCtx)
|
||||
}()
|
||||
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
// DeleteRefreshToken 双删:先删 DB,再删 Redis。
|
||||
func (r *CachedUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
// 先从 Redis 获取 user_id(用于从集合中移除)
|
||||
userID, _ := r.rdb.Get(ctx, refreshTokenKey(tokenHash)).Result()
|
||||
|
||||
// 删 DB
|
||||
if err := r.inner.DeleteRefreshToken(ctx, tokenHash); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 删 Redis
|
||||
key := refreshTokenKey(tokenHash)
|
||||
pipe := r.rdb.Pipeline()
|
||||
pipe.Del(ctx, key)
|
||||
if userID != "" {
|
||||
pipe.SRem(ctx, userRefreshKey(userID), tokenHash)
|
||||
}
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache delete failed for refresh token", "error", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteUserRefreshTokens 批量清理:先从 Redis 获取集合,逐个删缓存,再删 DB。
|
||||
func (r *CachedUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
|
||||
userKey := userRefreshKey(userID)
|
||||
|
||||
// 从 Redis 获取该用户所有 token hash
|
||||
hashes, _ := r.rdb.SMembers(ctx, userKey).Result()
|
||||
|
||||
// 批量删除 Redis 缓存
|
||||
if len(hashes) > 0 {
|
||||
keys := make([]string, 0, len(hashes)+1)
|
||||
for _, h := range hashes {
|
||||
keys = append(keys, refreshTokenKey(h))
|
||||
}
|
||||
keys = append(keys, userKey)
|
||||
if err := r.rdb.Del(ctx, keys...).Err(); err != nil {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache batch delete failed for user refresh tokens", "error", err, "user_id", userID)
|
||||
}
|
||||
}
|
||||
|
||||
// 删 DB(无论 Redis 是否成功都执行)
|
||||
return r.inner.DeleteUserRefreshTokens(ctx, userID)
|
||||
}
|
||||
17
backend/internal/store/db.go
Normal file
17
backend/internal/store/db.go
Normal 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)
|
||||
}
|
||||
50
backend/internal/store/message.go
Normal file
50
backend/internal/store/message.go
Normal 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"`
|
||||
}
|
||||
198
backend/internal/store/message_pg.go
Normal file
198
backend/internal/store/message_pg.go
Normal 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
|
||||
}
|
||||
80
backend/internal/store/migrate.go
Normal file
80
backend/internal/store/migrate.go
Normal 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
|
||||
}
|
||||
47
backend/internal/store/session.go
Normal file
47
backend/internal/store/session.go
Normal 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
|
||||
}
|
||||
189
backend/internal/store/session_pg.go
Normal file
189
backend/internal/store/session_pg.go
Normal 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
|
||||
}
|
||||
46
backend/internal/store/user.go
Normal file
46
backend/internal/store/user.go
Normal 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
|
||||
}
|
||||
120
backend/internal/store/user_mem.go
Normal file
120
backend/internal/store/user_mem.go
Normal 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 ©, 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 ©, 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
|
||||
}
|
||||
147
backend/internal/store/user_pg.go
Normal file
147
backend/internal/store/user_pg.go
Normal 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
|
||||
}
|
||||
276
backend/internal/store/user_scenario_repository.go
Normal file
276
backend/internal/store/user_scenario_repository.go
Normal file
@@ -0,0 +1,276 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// UserScenarioRepository 用户自建情景仓储接口。
|
||||
type UserScenarioRepository interface {
|
||||
Create(ctx context.Context, scenario *models.UserScenario) error
|
||||
FindByID(ctx context.Context, id string) (*models.UserScenario, error)
|
||||
FindByIDAndUserID(ctx context.Context, id, userID string) (*models.UserScenario, error)
|
||||
FindByUserID(ctx context.Context, userID string) ([]*models.UserScenario, error)
|
||||
Update(ctx context.Context, scenario *models.UserScenario) error
|
||||
Delete(ctx context.Context, id string) error
|
||||
CountByUserID(ctx context.Context, userID string) (int, error)
|
||||
}
|
||||
|
||||
// PostgresUserScenarioRepo PostgreSQL 实现。
|
||||
type PostgresUserScenarioRepo struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewPostgresUserScenarioRepo 创建 PostgreSQL 用户情景仓储。
|
||||
func NewPostgresUserScenarioRepo(pool *pgxpool.Pool) UserScenarioRepository {
|
||||
return &PostgresUserScenarioRepo{pool: pool}
|
||||
}
|
||||
|
||||
// Create 创建用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Create(ctx context.Context, scenario *models.UserScenario) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
INSERT INTO user_scenarios (id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at)
|
||||
VALUES ($1, $2, $3, $4, NULLIF($5, ''), $6, NULLIF($7, ''), $8, $9, $10)
|
||||
RETURNING id, created_at, updated_at
|
||||
`
|
||||
|
||||
now := time.Now()
|
||||
scenario.CreatedAt = now
|
||||
scenario.UpdatedAt = now
|
||||
|
||||
if scenario.ID == "" {
|
||||
scenario.ID = uuid.New().String()
|
||||
}
|
||||
if scenario.Icon == "" {
|
||||
scenario.Icon = "✨"
|
||||
}
|
||||
if scenario.Language == "" {
|
||||
scenario.Language = "zh-CN"
|
||||
}
|
||||
|
||||
err := r.pool.QueryRow(ctx, query,
|
||||
scenario.ID,
|
||||
scenario.UserID,
|
||||
scenario.Name,
|
||||
scenario.Icon,
|
||||
scenario.Description,
|
||||
scenario.Prompt,
|
||||
scenario.Greeting,
|
||||
scenario.Language,
|
||||
scenario.CreatedAt,
|
||||
scenario.UpdatedAt,
|
||||
).Scan(&scenario.ID, &scenario.CreatedAt, &scenario.UpdatedAt)
|
||||
|
||||
if err != nil {
|
||||
log.Errorw("create user scenario failed", "user_id", scenario.UserID, "name", scenario.Name, "error", err)
|
||||
return fmt.Errorf("create user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario created", "scenario_id", scenario.ID, "user_id", scenario.UserID, "name", scenario.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindByID 根据 ID 查找情景。
|
||||
func (r *PostgresUserScenarioRepo) FindByID(ctx context.Context, id string) (*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
WHERE id = $1
|
||||
`
|
||||
|
||||
var scenario models.UserScenario
|
||||
err := r.pool.QueryRow(ctx, query, id).Scan(
|
||||
&scenario.ID,
|
||||
&scenario.UserID,
|
||||
&scenario.Name,
|
||||
&scenario.Icon,
|
||||
&scenario.Description,
|
||||
&scenario.Prompt,
|
||||
&scenario.Greeting,
|
||||
&scenario.Language,
|
||||
&scenario.CreatedAt,
|
||||
&scenario.UpdatedAt,
|
||||
)
|
||||
|
||||
if err == pgx.ErrNoRows {
|
||||
return nil, fmt.Errorf("user scenario not found: %s", id)
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user scenario failed", "scenario_id", id, "error", err)
|
||||
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario found", "scenario_id", id)
|
||||
return &scenario, nil
|
||||
}
|
||||
|
||||
// FindByIDAndUserID 根据 ID 和用户 ID 查找情景(权限校验)。
|
||||
func (r *PostgresUserScenarioRepo) FindByIDAndUserID(ctx context.Context, id, userID string) (*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
WHERE id = $1 AND user_id = $2
|
||||
`
|
||||
|
||||
var scenario models.UserScenario
|
||||
err := r.pool.QueryRow(ctx, query, id, userID).Scan(
|
||||
&scenario.ID,
|
||||
&scenario.UserID,
|
||||
&scenario.Name,
|
||||
&scenario.Icon,
|
||||
&scenario.Description,
|
||||
&scenario.Prompt,
|
||||
&scenario.Greeting,
|
||||
&scenario.Language,
|
||||
&scenario.CreatedAt,
|
||||
&scenario.UpdatedAt,
|
||||
)
|
||||
|
||||
if err == pgx.ErrNoRows {
|
||||
return nil, fmt.Errorf("user scenario not found or no permission")
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user scenario by id and user failed", "scenario_id", id, "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario found by id and user", "scenario_id", id, "user_id", userID)
|
||||
return &scenario, nil
|
||||
}
|
||||
|
||||
// FindByUserID 查找用户的所有情景。
|
||||
func (r *PostgresUserScenarioRepo) FindByUserID(ctx context.Context, userID string) ([]*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
WHERE user_id = $1
|
||||
ORDER BY created_at DESC
|
||||
`
|
||||
|
||||
rows, err := r.pool.Query(ctx, query, userID)
|
||||
if err != nil {
|
||||
log.Errorw("find user scenarios failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("find user scenarios: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var scenarios []*models.UserScenario
|
||||
for rows.Next() {
|
||||
var s models.UserScenario
|
||||
err := rows.Scan(
|
||||
&s.ID,
|
||||
&s.UserID,
|
||||
&s.Name,
|
||||
&s.Icon,
|
||||
&s.Description,
|
||||
&s.Prompt,
|
||||
&s.Greeting,
|
||||
&s.Language,
|
||||
&s.CreatedAt,
|
||||
&s.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("scan user scenario row failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("scan user scenario: %w", err)
|
||||
}
|
||||
scenarios = append(scenarios, &s)
|
||||
}
|
||||
|
||||
if err = rows.Err(); err != nil {
|
||||
log.Errorw("iterate user scenarios failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("iterate user scenarios: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenarios found", "user_id", userID, "count", len(scenarios))
|
||||
return scenarios, nil
|
||||
}
|
||||
|
||||
// Update 更新用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Update(ctx context.Context, scenario *models.UserScenario) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
UPDATE user_scenarios
|
||||
SET name = $1, icon = $2, description = $3, prompt = $4, greeting = $5, language = $6, updated_at = $7
|
||||
WHERE id = $8 AND user_id = $9
|
||||
RETURNING updated_at
|
||||
`
|
||||
|
||||
scenario.UpdatedAt = time.Now()
|
||||
|
||||
err := r.pool.QueryRow(ctx, query,
|
||||
scenario.Name,
|
||||
scenario.Icon,
|
||||
scenario.Description,
|
||||
scenario.Prompt,
|
||||
scenario.Greeting,
|
||||
scenario.Language,
|
||||
scenario.UpdatedAt,
|
||||
scenario.ID,
|
||||
scenario.UserID,
|
||||
).Scan(&scenario.UpdatedAt)
|
||||
|
||||
if err == pgx.ErrNoRows {
|
||||
return fmt.Errorf("user scenario not found or no permission")
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("update user scenario failed", "scenario_id", scenario.ID, "user_id", scenario.UserID, "error", err)
|
||||
return fmt.Errorf("update user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario updated", "scenario_id", scenario.ID, "user_id", scenario.UserID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 删除用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Delete(ctx context.Context, id string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `DELETE FROM user_scenarios WHERE id = $1`
|
||||
|
||||
result, err := r.pool.Exec(ctx, query, id)
|
||||
if err != nil {
|
||||
log.Errorw("delete user scenario failed", "scenario_id", id, "error", err)
|
||||
return fmt.Errorf("delete user scenario: %w", err)
|
||||
}
|
||||
|
||||
if result.RowsAffected() == 0 {
|
||||
return fmt.Errorf("user scenario not found")
|
||||
}
|
||||
|
||||
log.Debugw("user scenario deleted", "scenario_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
// CountByUserID 统计用户的情景数量。
|
||||
func (r *PostgresUserScenarioRepo) CountByUserID(ctx context.Context, userID string) (int, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `SELECT COUNT(*) FROM user_scenarios WHERE user_id = $1`
|
||||
|
||||
var count int
|
||||
err := r.pool.QueryRow(ctx, query, userID).Scan(&count)
|
||||
if err != nil {
|
||||
log.Errorw("count user scenarios failed", "user_id", userID, "error", err)
|
||||
return 0, fmt.Errorf("count user scenarios: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenarios counted", "user_id", userID, "count", count)
|
||||
return count, nil
|
||||
}
|
||||
185
backend/internal/store/user_test.go
Normal file
185
backend/internal/store/user_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
46
backend/internal/trace/context.go
Normal file
46
backend/internal/trace/context.go
Normal file
@@ -0,0 +1,46 @@
|
||||
package trace
|
||||
|
||||
import "context"
|
||||
|
||||
type traceIDKey struct{}
|
||||
type requestIDKey struct{}
|
||||
type sessionIDKey struct{}
|
||||
|
||||
// WithTraceID 将 trace ID 注入 context(连接级/会话级标识)
|
||||
func WithTraceID(ctx context.Context, traceID string) context.Context {
|
||||
return context.WithValue(ctx, traceIDKey{}, traceID)
|
||||
}
|
||||
|
||||
// GetTraceID 从 context 提取 trace ID
|
||||
func GetTraceID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// WithRequestID 将 request ID 注入 context(单次请求/查询标识)
|
||||
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||
return context.WithValue(ctx, requestIDKey{}, requestID)
|
||||
}
|
||||
|
||||
// GetRequestID 从 context 提取 request ID
|
||||
func GetRequestID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(requestIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// WithSessionID 将 session ID 注入 context(会话存储标识)
|
||||
func WithSessionID(ctx context.Context, sessionID string) context.Context {
|
||||
return context.WithValue(ctx, sessionIDKey{}, sessionID)
|
||||
}
|
||||
|
||||
// GetSessionID 从 context 提取 session ID
|
||||
func GetSessionID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(sessionIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
42
backend/internal/trace/eino_test.go
Normal file
42
backend/internal/trace/eino_test.go
Normal file
@@ -0,0 +1,42 @@
|
||||
package trace_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
func TestEinoContextPropagation(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
testTraceID := "01J5TEST123456789"
|
||||
ctx = trace.WithTraceID(ctx, testTraceID)
|
||||
|
||||
var capturedTraceID string
|
||||
|
||||
g := compose.NewGraph[string, string]()
|
||||
g.AddLambdaNode("test_node", compose.InvokableLambda(
|
||||
func(ctx context.Context, input string) (string, error) {
|
||||
capturedTraceID = trace.GetTraceID(ctx)
|
||||
return "ok", nil
|
||||
},
|
||||
))
|
||||
g.AddEdge(compose.START, "test_node")
|
||||
g.AddEdge("test_node", compose.END)
|
||||
|
||||
runnable, err := g.Compile(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("compile failed: %v", err)
|
||||
}
|
||||
|
||||
_, err = runnable.Invoke(ctx, "test_input")
|
||||
if err != nil {
|
||||
t.Fatalf("invoke failed: %v", err)
|
||||
}
|
||||
|
||||
if capturedTraceID != testTraceID {
|
||||
t.Errorf("trace_id lost in Eino propagation: got %q, want %q",
|
||||
capturedTraceID, testTraceID)
|
||||
}
|
||||
}
|
||||
63
backend/internal/trace/gin_logger.go
Normal file
63
backend/internal/trace/gin_logger.go
Normal file
@@ -0,0 +1,63 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GinLogger 记录每个 HTTP 请求的 method/path/status/latency
|
||||
func GinLogger() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
start := time.Now()
|
||||
path := c.Request.URL.Path
|
||||
query := c.Request.URL.RawQuery
|
||||
|
||||
c.Next()
|
||||
|
||||
latency := time.Since(start).Milliseconds()
|
||||
status := c.Writer.Status()
|
||||
log := FromContext(c.Request.Context())
|
||||
|
||||
fields := []interface{}{
|
||||
"method", c.Request.Method,
|
||||
"path", path,
|
||||
"status", status,
|
||||
"latency_ms", latency,
|
||||
"client_ip", c.ClientIP(),
|
||||
}
|
||||
if query != "" {
|
||||
fields = append(fields, "query", query)
|
||||
}
|
||||
if errStr := c.Errors.String(); errStr != "" {
|
||||
fields = append(fields, "errors", errStr)
|
||||
}
|
||||
|
||||
switch {
|
||||
case status >= 500:
|
||||
log.Errorw("request completed", fields...)
|
||||
case status >= 400:
|
||||
log.Warnw("request completed", fields...)
|
||||
default:
|
||||
log.Infow("request completed", fields...)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GinRecovery 自定义 panic 恢复中间件,使用 zap 记录
|
||||
func GinRecovery() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log := FromContext(c.Request.Context())
|
||||
log.Errorw("panic recovered",
|
||||
"error", err,
|
||||
"path", c.Request.URL.Path,
|
||||
"method", c.Request.Method,
|
||||
"client_ip", c.ClientIP())
|
||||
c.AbortWithStatus(500)
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
22
backend/internal/trace/id.go
Normal file
22
backend/internal/trace/id.go
Normal file
@@ -0,0 +1,22 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
cryptorand "crypto/rand"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/oklog/ulid/v2"
|
||||
)
|
||||
|
||||
var entropyPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return ulid.Monotonic(cryptorand.Reader, 0)
|
||||
},
|
||||
}
|
||||
|
||||
// GenerateTraceID 生成并发安全的 ULID trace ID
|
||||
func GenerateTraceID() string {
|
||||
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
|
||||
defer entropyPool.Put(entropy)
|
||||
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
|
||||
}
|
||||
25
backend/internal/trace/logger.go
Normal file
25
backend/internal/trace/logger.go
Normal file
@@ -0,0 +1,25 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
|
||||
func FromContext(ctx context.Context) *zap.SugaredLogger {
|
||||
log := logger.Log
|
||||
|
||||
if traceID := GetTraceID(ctx); traceID != "" {
|
||||
log = log.With("trace_id", traceID)
|
||||
}
|
||||
if requestID := GetRequestID(ctx); requestID != "" {
|
||||
log = log.With("request_id", requestID)
|
||||
}
|
||||
if sessionID := GetSessionID(ctx); sessionID != "" {
|
||||
log = log.With("session_id", sessionID)
|
||||
}
|
||||
|
||||
return log
|
||||
}
|
||||
17
backend/internal/trace/middleware.go
Normal file
17
backend/internal/trace/middleware.go
Normal file
@@ -0,0 +1,17 @@
|
||||
package trace
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
|
||||
func TraceMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
traceID := GenerateTraceID()
|
||||
ctx := WithTraceID(c.Request.Context(), traceID)
|
||||
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
|
||||
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
9
backend/internal/util/string.go
Normal file
9
backend/internal/util/string.go
Normal file
@@ -0,0 +1,9 @@
|
||||
package util
|
||||
|
||||
// Truncate 截断字符串到指定长度,超出部分用 "..." 替换
|
||||
func Truncate(s string, maxLen int) string {
|
||||
if len(s) <= maxLen {
|
||||
return s
|
||||
}
|
||||
return s[:maxLen] + "..."
|
||||
}
|
||||
@@ -1,55 +1,190 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"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/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
|
||||
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
|
||||
}
|
||||
|
||||
func (c *Client) sendJSON(v any) error {
|
||||
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
||||
func (c *Client) SendJSON(v any) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.conn.WriteJSON(v)
|
||||
}
|
||||
|
||||
// WSClient 实现 orchestrator.Sender 接口,将消息推送到 WebSocket 连接。
|
||||
type WSClient struct {
|
||||
client *Client
|
||||
requestID string
|
||||
}
|
||||
|
||||
// SendSTTResult 发送语音识别结果。
|
||||
func (w *WSClient) SendSTTResult(result models.WsSTTResult) error {
|
||||
result.RequestID = w.requestID
|
||||
return w.client.SendJSON(result)
|
||||
}
|
||||
|
||||
// SendLLMChunk 发送 LLM 流式文本增量。
|
||||
func (w *WSClient) SendLLMChunk(chunk models.WsLLMChunk) error {
|
||||
chunk.RequestID = w.requestID
|
||||
return w.client.SendJSON(chunk)
|
||||
}
|
||||
|
||||
// SendLLMDone 发送 LLM 流结束信号。
|
||||
func (w *WSClient) SendLLMDone(done models.WsLLMDone) error {
|
||||
done.RequestID = w.requestID
|
||||
return w.client.SendJSON(done)
|
||||
}
|
||||
|
||||
// SendTTSAudio 发送 TTS 音频数据。
|
||||
func (w *WSClient) SendTTSAudio(audio models.WsTTSAudio) error {
|
||||
audio.RequestID = w.requestID
|
||||
return w.client.SendJSON(audio)
|
||||
}
|
||||
|
||||
// SendError 发送错误消息。
|
||||
func (w *WSClient) SendError(err models.WsError) error {
|
||||
err.RequestID = w.requestID
|
||||
return w.client.SendJSON(err)
|
||||
}
|
||||
|
||||
// ServeWS 处理 WebSocket 升级请求。
|
||||
func ServeWS(c *gin.Context) {
|
||||
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, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr, limiter, scenarioRepo)
|
||||
}
|
||||
}
|
||||
|
||||
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
|
||||
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, 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 {
|
||||
log.Printf("websocket upgrade failed: %v", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Errorw("websocket upgrade failed", "error", err)
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
sessionID := uuid.New().String()
|
||||
client := &Client{conn: conn, sessionID: sessionID}
|
||||
// 创建或复用会话
|
||||
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{
|
||||
conn: conn,
|
||||
sessionID: sessionID,
|
||||
sessionMgr: sessionMgr,
|
||||
orchestrator: orch,
|
||||
cancelFuncs: make(map[string]context.CancelFunc),
|
||||
}
|
||||
|
||||
// 发送 connected 消息
|
||||
_ = client.sendJSON(models.WsConnected{
|
||||
_ = client.SendJSON(models.WsConnected{
|
||||
Type: "connected",
|
||||
SessionID: sessionID,
|
||||
ServerVersion: "0.1.0",
|
||||
ServerVersion: version,
|
||||
})
|
||||
log.Printf("client connected: session=%s", sessionID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("client connected", "user_id", userID, "username", username)
|
||||
|
||||
// 心跳检测
|
||||
lastPong := time.Now()
|
||||
@@ -61,13 +196,14 @@ func ServeWS(c *gin.Context) {
|
||||
// 启动心跳检查 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 {
|
||||
log.Printf("heartbeat timeout: session=%s", sessionID)
|
||||
if time.Since(lastPong) > heartbeatTimeout {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("heartbeat timeout")
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
@@ -82,7 +218,8 @@ func ServeWS(c *gin.Context) {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
|
||||
log.Printf("ws read error: %v", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("ws read error", "error", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
@@ -92,51 +229,168 @@ func ServeWS(c *gin.Context) {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if err := json.Unmarshal(message, &envelope); err != nil {
|
||||
_ = client.sendJSON(models.WsError{
|
||||
Type: "error",
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "invalid JSON",
|
||||
})
|
||||
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
|
||||
continue
|
||||
}
|
||||
|
||||
switch envelope.Type {
|
||||
case "ping":
|
||||
_ = client.sendJSON(models.WsPong{Type: "pong"})
|
||||
lastPong = time.Now() // 刷新心跳计时器
|
||||
_ = client.SendJSON(models.WsPong{Type: "pong"})
|
||||
|
||||
case "query":
|
||||
var msg models.WsQuery
|
||||
if err := json.Unmarshal(message, &msg); err != nil {
|
||||
_ = client.sendJSON(models.WsError{
|
||||
Type: "error",
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "invalid query message",
|
||||
RequestID: msg.RequestID,
|
||||
})
|
||||
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
|
||||
continue
|
||||
}
|
||||
log.Printf("query received: session=%s request=%s", sessionID, msg.RequestID)
|
||||
// TODO: 调用 AI 编排流程(STT → LLM → TTS)
|
||||
|
||||
// 注入 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 {
|
||||
log.Warnw("touch session failed", "error", err)
|
||||
}
|
||||
|
||||
// 标记活跃请求
|
||||
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
||||
log.Warnw("set active request failed", "error", err)
|
||||
}
|
||||
|
||||
// 创建可取消的 context
|
||||
processCtx, cancel := context.WithCancel(queryCtx)
|
||||
client.mu.Lock()
|
||||
client.cancelFuncs[msg.RequestID] = cancel
|
||||
client.mu.Unlock()
|
||||
|
||||
// 创建 sender
|
||||
sender := &WSClient{client: client, requestID: msg.RequestID}
|
||||
|
||||
// 启动 orchestrator 处理 goroutine
|
||||
go func() {
|
||||
defer func() {
|
||||
// 清理 cancel func
|
||||
client.mu.Lock()
|
||||
delete(client.cancelFuncs, msg.RequestID)
|
||||
client.mu.Unlock()
|
||||
cancel()
|
||||
// 清除活跃请求
|
||||
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||
}()
|
||||
|
||||
if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil {
|
||||
log := trace.FromContext(processCtx)
|
||||
log.Errorw("process query failed", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
case "config":
|
||||
var msg models.WsConfig
|
||||
if err := json.Unmarshal(message, &msg); err != nil {
|
||||
_ = client.sendJSON(models.WsError{
|
||||
Type: "error",
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "invalid config message",
|
||||
})
|
||||
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
|
||||
continue
|
||||
}
|
||||
log.Printf("config update: session=%s", sessionID)
|
||||
// TODO: 更新会话配置
|
||||
|
||||
patch := models.SessionConfigPatch{
|
||||
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
|
||||
}
|
||||
|
||||
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":
|
||||
log.Printf("interrupt received: session=%s", sessionID)
|
||||
// TODO: 中断当前 AI 响应
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("interrupt received")
|
||||
|
||||
// 获取活跃请求 ID 并取消
|
||||
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
||||
if reqID != "" {
|
||||
client.mu.Lock()
|
||||
if cancel, ok := client.cancelFuncs[reqID]; ok {
|
||||
cancel()
|
||||
delete(client.cancelFuncs, reqID)
|
||||
}
|
||||
client.mu.Unlock()
|
||||
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||
}
|
||||
|
||||
default:
|
||||
_ = client.sendJSON(models.WsError{
|
||||
_ = client.SendJSON(models.WsError{
|
||||
Type: "error",
|
||||
Code: "INVALID_MESSAGE",
|
||||
Message: "unknown message type: " + envelope.Type,
|
||||
@@ -145,5 +399,18 @@ func ServeWS(c *gin.Context) {
|
||||
}
|
||||
|
||||
close(done)
|
||||
log.Printf("client disconnected: session=%s", sessionID)
|
||||
|
||||
// 取消所有活跃请求
|
||||
client.mu.Lock()
|
||||
for reqID, cancel := range client.cancelFuncs {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("canceling active request on disconnect", "request", reqID)
|
||||
cancel()
|
||||
}
|
||||
client.cancelFuncs = make(map[string]context.CancelFunc)
|
||||
client.mu.Unlock()
|
||||
|
||||
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
||||
log = trace.FromContext(ctx)
|
||||
log.Infow("client disconnected")
|
||||
}
|
||||
|
||||
728
backend/internal/ws/handler_test.go
Normal file
728
backend/internal/ws/handler_test.go
Normal file
@@ -0,0 +1,728 @@
|
||||
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"
|
||||
|
||||
"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"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
)
|
||||
|
||||
func init() {
|
||||
logger.Init("debug", "console")
|
||||
gin.SetMode(gin.TestMode)
|
||||
}
|
||||
|
||||
// --- Mock Orchestrator ---
|
||||
|
||||
// MockOrchestrator 实现 orchestrator.Orchestrator 接口,
|
||||
// 模拟完整的 STT → LLM → TTS 管道,通过 sender 推送消息。
|
||||
type MockOrchestrator struct {
|
||||
// STTResult 模拟的语音识别结果
|
||||
STTResult string
|
||||
// LLMDeltas 模拟的 LLM 流式输出
|
||||
LLMDeltas []string
|
||||
// TTSAudios 模拟的 TTS 音频数据(每项一个 base64 编码的 MP3 片段)
|
||||
TTSAudios []string
|
||||
// Err 如果非 nil,ProcessQuery 直接返回此错误
|
||||
Err error
|
||||
// Delay 每个消息之间的延迟(用于 interrupt 测试)
|
||||
Delay time.Duration
|
||||
}
|
||||
|
||||
func (m *MockOrchestrator) ProcessQuery(
|
||||
ctx context.Context,
|
||||
sessionID string,
|
||||
req models.WsQuery,
|
||||
sender orchestrator.Sender,
|
||||
) error {
|
||||
if m.Err != nil {
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
Code: "INTERNAL_ERROR",
|
||||
Message: m.Err.Error(),
|
||||
})
|
||||
return m.Err
|
||||
}
|
||||
|
||||
// Step 1: 发送 STT 结果
|
||||
if m.STTResult != "" {
|
||||
_ = sender.SendSTTResult(models.WsSTTResult{
|
||||
Type: "stt_result",
|
||||
RequestID: req.RequestID,
|
||||
Text: m.STTResult,
|
||||
IsFinal: true,
|
||||
})
|
||||
}
|
||||
if m.Delay > 0 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil // 中断视为正常完成
|
||||
case <-time.After(m.Delay):
|
||||
}
|
||||
}
|
||||
|
||||
// Step 2: 发送 LLM chunks
|
||||
var fullText strings.Builder
|
||||
for _, delta := range m.LLMDeltas {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil // 中断视为正常完成
|
||||
default:
|
||||
}
|
||||
fullText.WriteString(delta)
|
||||
_ = sender.SendLLMChunk(models.WsLLMChunk{
|
||||
Type: "llm_chunk",
|
||||
RequestID: req.RequestID,
|
||||
Delta: delta,
|
||||
Role: "assistant",
|
||||
})
|
||||
if m.Delay > 0 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil // 中断视为正常完成
|
||||
case <-time.After(m.Delay):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Step 3: 发送 TTS 音频
|
||||
for i, audio := range m.TTSAudios {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil // 中断视为正常完成
|
||||
default:
|
||||
}
|
||||
isLast := i == len(m.TTSAudios)-1
|
||||
_ = sender.SendTTSAudio(models.WsTTSAudio{
|
||||
Type: "tts_audio",
|
||||
RequestID: req.RequestID,
|
||||
Audio: audio,
|
||||
MimeType: "audio/mp3",
|
||||
IsLast: isLast,
|
||||
})
|
||||
}
|
||||
|
||||
// Step 4: 发送 llm_done
|
||||
_ = sender.SendLLMDone(models.WsLLMDone{
|
||||
Type: "llm_done",
|
||||
RequestID: req.RequestID,
|
||||
FullText: fullText.String(),
|
||||
Model: "gpt-4o",
|
||||
LatencyMs: 100,
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// --- 测试辅助函数 ---
|
||||
|
||||
// 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()
|
||||
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)
|
||||
|
||||
// 生成有效 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
|
||||
}
|
||||
|
||||
// connectWS 建立 WebSocket 连接并返回 conn。
|
||||
func connectWS(t *testing.T, wsURL string) *websocket.Conn {
|
||||
t.Helper()
|
||||
|
||||
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
||||
require.NoError(t, err, "WebSocket 连接失败")
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
return conn
|
||||
}
|
||||
|
||||
// readJSON 从 WebSocket 读取一条 JSON 消息。
|
||||
func readJSON(t *testing.T, conn *websocket.Conn) map[string]any {
|
||||
t.Helper()
|
||||
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
var msg map[string]any
|
||||
err := conn.ReadJSON(&msg)
|
||||
require.NoError(t, err, "读取 WebSocket 消息失败")
|
||||
return msg
|
||||
}
|
||||
|
||||
// --- 测试用例 ---
|
||||
|
||||
// TestWS_Connected 验证连接建立后收到 connected 消息。
|
||||
func TestWS_Connected(t *testing.T) {
|
||||
srv, wsURL := setupTestServer(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
|
||||
msg := readJSON(t, conn)
|
||||
assert.Equal(t, "connected", msg["type"])
|
||||
assert.NotEmpty(t, msg["session_id"])
|
||||
assert.Equal(t, "test", msg["server_version"])
|
||||
}
|
||||
|
||||
// TestWS_PingPong 验证 ping/pong 心跳。
|
||||
func TestWS_PingPong(t *testing.T) {
|
||||
srv, wsURL := setupTestServer(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
|
||||
// 读取 connected 消息
|
||||
_ = readJSON(t, conn)
|
||||
|
||||
// 发送 ping
|
||||
err := conn.WriteJSON(map[string]string{"type": "ping"})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 读取 pong
|
||||
msg := readJSON(t, conn)
|
||||
assert.Equal(t, "pong", msg["type"])
|
||||
}
|
||||
|
||||
// TestWS_QueryFullFlow 验证完整的 query → stt_result → llm_chunk → tts_audio → llm_done 流程。
|
||||
func TestWS_QueryFullFlow(t *testing.T) {
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("fake-audio-data"))
|
||||
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"))},
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
|
||||
// 1. 读取 connected
|
||||
connected := readJSON(t, conn)
|
||||
assert.Equal(t, "connected", connected["type"])
|
||||
sessionID := connected["session_id"].(string)
|
||||
assert.NotEmpty(t, sessionID)
|
||||
|
||||
// 2. 发送 query
|
||||
queryMsg := models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-test-001",
|
||||
Image: imageB64,
|
||||
Audio: audioB64,
|
||||
MimeType: "audio/pcm",
|
||||
}
|
||||
err := conn.WriteJSON(queryMsg)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 3. 读取 stt_result
|
||||
sttResult := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", sttResult["type"])
|
||||
assert.Equal(t, "req-test-001", sttResult["request_id"])
|
||||
assert.Equal(t, "你好,世界", sttResult["text"])
|
||||
assert.Equal(t, true, sttResult["is_final"])
|
||||
|
||||
// 4. 读取 llm_chunk 消息
|
||||
chunk1 := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk1["type"])
|
||||
assert.Equal(t, "req-test-001", chunk1["request_id"])
|
||||
assert.Equal(t, "你好", chunk1["delta"])
|
||||
assert.Equal(t, "assistant", chunk1["role"])
|
||||
|
||||
chunk2 := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk2["type"])
|
||||
assert.Equal(t, ",世界!", chunk2["delta"])
|
||||
|
||||
// 5. 读取 tts_audio 消息
|
||||
tts1 := readJSON(t, conn)
|
||||
assert.Equal(t, "tts_audio", tts1["type"])
|
||||
assert.Equal(t, "req-test-001", tts1["request_id"])
|
||||
assert.NotEmpty(t, tts1["audio"])
|
||||
assert.Equal(t, "audio/mp3", tts1["mime_type"])
|
||||
assert.Equal(t, false, tts1["is_last"])
|
||||
|
||||
tts2 := readJSON(t, conn)
|
||||
assert.Equal(t, "tts_audio", tts2["type"])
|
||||
assert.Equal(t, true, tts2["is_last"])
|
||||
|
||||
// 6. 读取 llm_done
|
||||
llmDone := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_done", llmDone["type"])
|
||||
assert.Equal(t, "req-test-001", llmDone["request_id"])
|
||||
assert.Equal(t, "你好,世界!", llmDone["full_text"])
|
||||
assert.Equal(t, "gpt-4o", llmDone["model"])
|
||||
assert.NotNil(t, llmDone["latency_ms"])
|
||||
}
|
||||
|
||||
// TestWS_QuerySTTOnly 验证只有 STT 结果、无 LLM 输出的场景。
|
||||
func TestWS_QuerySTTOnly(t *testing.T) {
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "测试语音",
|
||||
// LLMDeltas 为空 → 不发送 llm_chunk
|
||||
// TTSAudios 为空 → 不发送 tts_audio
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-stt-only",
|
||||
Audio: audioB64,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 应收到 stt_result
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"])
|
||||
assert.Equal(t, "测试语音", stt["text"])
|
||||
|
||||
// 应收到 llm_done(即使没有 chunk)
|
||||
done := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_done", done["type"])
|
||||
assert.Equal(t, "", done["full_text"])
|
||||
}
|
||||
|
||||
// TestWS_UnknownMessageType 验证未知消息类型返回 error。
|
||||
func TestWS_UnknownMessageType(t *testing.T) {
|
||||
srv, wsURL := setupTestServer(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
err := conn.WriteJSON(map[string]string{"type": "unknown_type"})
|
||||
require.NoError(t, err)
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
// TestWS_InvalidJSON 验证无效 JSON 返回 error。
|
||||
func TestWS_InvalidJSON(t *testing.T) {
|
||||
srv, wsURL := setupTestServer(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
err := conn.WriteMessage(websocket.TextMessage, []byte("not-json"))
|
||||
require.NoError(t, err)
|
||||
|
||||
errMsg := readJSON(t, conn)
|
||||
assert.Equal(t, "error", errMsg["type"])
|
||||
assert.Equal(t, "INVALID_MESSAGE", errMsg["code"])
|
||||
}
|
||||
|
||||
// TestWS_MultipleQueries 验证同一连接上可以发送多次 query。
|
||||
func TestWS_MultipleQueries(t *testing.T) {
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "识别结果",
|
||||
LLMDeltas: []string{"回复"},
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
reqID := "req-multi-" + string(rune('0'+i))
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: reqID,
|
||||
Audio: audioB64,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 每次应收到完整的响应序列
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"], "第 %d 次 query", i+1)
|
||||
|
||||
chunk := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk["type"], "第 %d 次 query", i+1)
|
||||
|
||||
done := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_done", done["type"], "第 %d 次 query", i+1)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWS_Interrupt 验证 interrupt 取消正在进行的请求。
|
||||
func TestWS_Interrupt(t *testing.T) {
|
||||
// 使用较长延迟模拟慢请求
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "识别文本",
|
||||
LLMDeltas: []string{"第一句", "第二句", "第三句", "第四句", "第五句"},
|
||||
Delay: 200 * time.Millisecond,
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
// 发送 query
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-interrupt",
|
||||
Audio: audioB64,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 收到 stt_result
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"])
|
||||
|
||||
// 收到第一个 llm_chunk
|
||||
chunk1 := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk1["type"])
|
||||
|
||||
// 发送 interrupt
|
||||
err = conn.WriteJSON(map[string]string{"type": "interrupt"})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 等待 interrupt 生效
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// 验证连接仍然存活(可以发 ping 收 pong)
|
||||
require.NoError(t, conn.WriteJSON(map[string]string{"type": "ping"}))
|
||||
|
||||
var pong map[string]any
|
||||
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
|
||||
require.NoError(t, conn.ReadJSON(&pong), "interrupt 后连接应仍存活")
|
||||
assert.Equal(t, "pong", pong["type"])
|
||||
}
|
||||
|
||||
// TestWS_DisconnectCleanup 验证断开连接时清理资源。
|
||||
func TestWS_DisconnectCleanup(t *testing.T) {
|
||||
// 使用较长延迟模拟慢请求
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "识别文本",
|
||||
LLMDeltas: []string{"长回复第一部分", "长回复第二部分"},
|
||||
Delay: 500 * time.Millisecond,
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
// 发送 query
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-disconnect",
|
||||
Audio: audioB64,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// 收到 stt_result
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"])
|
||||
|
||||
// 关闭连接(模拟客户端断开)
|
||||
conn.Close()
|
||||
|
||||
// 等待一小段时间让服务器处理断开
|
||||
time.Sleep(300 * time.Millisecond)
|
||||
|
||||
// 如果没有 panic 或 goroutine 泄漏,测试通过
|
||||
// (Go test 的 -race 检测器会捕获数据竞争)
|
||||
}
|
||||
|
||||
// TestWS_SessionCreated 验证每次连接都创建新会话。
|
||||
func TestWS_SessionCreated(t *testing.T) {
|
||||
srv, wsURL := setupTestServer(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
// 第一次连接
|
||||
conn1 := connectWS(t, wsURL)
|
||||
msg1 := readJSON(t, conn1)
|
||||
sid1 := msg1["session_id"].(string)
|
||||
conn1.Close()
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// 第二次连接
|
||||
conn2 := connectWS(t, wsURL)
|
||||
sid2 := readJSON(t, conn2)["session_id"].(string)
|
||||
|
||||
assert.NotEmpty(t, sid1)
|
||||
assert.NotEmpty(t, sid2)
|
||||
assert.NotEqual(t, sid1, sid2, "两次连接应创建不同的会话")
|
||||
}
|
||||
|
||||
// TestWS_QueryWithoutImage 验证不带图片的 query。
|
||||
func TestWS_QueryWithoutImage(t *testing.T) {
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "纯语音输入",
|
||||
LLMDeltas: []string{"收到"},
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-no-image",
|
||||
Audio: audioB64,
|
||||
// Image 为空
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"])
|
||||
assert.Equal(t, "纯语音输入", stt["text"])
|
||||
|
||||
chunk := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk["type"])
|
||||
|
||||
done := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_done", done["type"])
|
||||
}
|
||||
|
||||
// TestWS_QueryWithTTSDisabled 验证 TTS 未启用时不应收到 tts_audio。
|
||||
func TestWS_QueryWithTTSDisabled(t *testing.T) {
|
||||
// MockOrchestrator 的 TTSAudios 为空 → 不发送 tts_audio
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "语音",
|
||||
LLMDeltas: []string{"回复"},
|
||||
// TTSAudios 留空
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
defer srv.Close()
|
||||
|
||||
conn := connectWS(t, wsURL)
|
||||
_ = readJSON(t, conn) // connected
|
||||
|
||||
audioB64 := base64.StdEncoding.EncodeToString([]byte("audio"))
|
||||
|
||||
err := conn.WriteJSON(models.WsQuery{
|
||||
Type: "query",
|
||||
RequestID: "req-no-tts",
|
||||
Audio: audioB64,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stt := readJSON(t, conn)
|
||||
assert.Equal(t, "stt_result", stt["type"])
|
||||
|
||||
chunk := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_chunk", chunk["type"])
|
||||
|
||||
done := readJSON(t, conn)
|
||||
assert.Equal(t, "llm_done", done["type"])
|
||||
|
||||
// 不应有 tts_audio 消息;设置短超时验证
|
||||
conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond))
|
||||
var extra map[string]any
|
||||
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")
|
||||
}
|
||||
5
backend/migrations/001_users.down.sql
Normal file
5
backend/migrations/001_users.down.sql
Normal file
@@ -0,0 +1,5 @@
|
||||
-- 删除 Refresh Token 表(自动删除相关索引)
|
||||
DROP TABLE IF EXISTS refresh_tokens;
|
||||
|
||||
-- 删除用户表(自动删除相关索引)
|
||||
DROP TABLE IF EXISTS users;
|
||||
42
backend/migrations/001_users.up.sql
Normal file
42
backend/migrations/001_users.up.sql
Normal 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 创建时间';
|
||||
1
backend/migrations/002_messages.down.sql
Normal file
1
backend/migrations/002_messages.down.sql
Normal file
@@ -0,0 +1 @@
|
||||
DROP TABLE IF EXISTS messages;
|
||||
28
backend/migrations/002_messages.up.sql
Normal file
28
backend/migrations/002_messages.up.sql
Normal 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 '消息创建时间';
|
||||
1
backend/migrations/003_sessions.down.sql
Normal file
1
backend/migrations/003_sessions.down.sql
Normal file
@@ -0,0 +1 @@
|
||||
DROP TABLE IF EXISTS sessions;
|
||||
22
backend/migrations/003_sessions.up.sql
Normal file
22
backend/migrations/003_sessions.up.sql
Normal file
@@ -0,0 +1,22 @@
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id UUID PRIMARY KEY,
|
||||
user_id UUID NOT NULL,
|
||||
title VARCHAR(100) NOT NULL DEFAULT '新对话',
|
||||
config JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
|
||||
CONSTRAINT check_title_length CHECK (char_length(title) >= 1 AND char_length(title) <= 100)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions (user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_user_updated ON sessions (user_id, updated_at DESC);
|
||||
|
||||
-- 表和列注释
|
||||
COMMENT ON TABLE sessions IS '会话表,存储用户的对话会话信息';
|
||||
COMMENT ON COLUMN sessions.id IS '会话唯一标识符 (UUID)';
|
||||
COMMENT ON COLUMN sessions.user_id IS '所属用户 ID,关联 users 表';
|
||||
COMMENT ON COLUMN sessions.title IS '会话标题,默认为"新对话",长度 1-100 字符';
|
||||
COMMENT ON COLUMN sessions.config IS '会话配置 (JSONB),包含: tts_enabled (布尔), detail_level (''low''/''high''), language (语言代码), scenario (情景 ID)';
|
||||
COMMENT ON COLUMN sessions.created_at IS '会话创建时间';
|
||||
COMMENT ON COLUMN sessions.updated_at IS '会话最后更新时间';
|
||||
6
backend/migrations/004_user_scenarios.down.sql
Normal file
6
backend/migrations/004_user_scenarios.down.sql
Normal file
@@ -0,0 +1,6 @@
|
||||
-- 004_user_scenarios.down.sql
|
||||
-- 回滚用户自建情景表
|
||||
|
||||
DROP INDEX IF EXISTS idx_user_scenarios_created_at;
|
||||
DROP INDEX IF EXISTS idx_user_scenarios_user_id;
|
||||
DROP TABLE IF EXISTS user_scenarios;
|
||||
40
backend/migrations/004_user_scenarios.up.sql
Normal file
40
backend/migrations/004_user_scenarios.up.sql
Normal file
@@ -0,0 +1,40 @@
|
||||
-- 004_user_scenarios.up.sql
|
||||
-- 用户自建情景表
|
||||
|
||||
CREATE TABLE user_scenarios (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
name VARCHAR(50) NOT NULL,
|
||||
icon VARCHAR(20) DEFAULT '✨',
|
||||
description VARCHAR(100),
|
||||
prompt TEXT NOT NULL,
|
||||
greeting VARCHAR(500),
|
||||
language VARCHAR(10) DEFAULT 'zh-CN',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
|
||||
CONSTRAINT unique_user_scenario UNIQUE(user_id, name),
|
||||
CONSTRAINT check_name_length CHECK (char_length(name) >= 2 AND char_length(name) <= 50),
|
||||
CONSTRAINT check_description_length CHECK (description IS NULL OR char_length(description) <= 100),
|
||||
CONSTRAINT check_prompt_length CHECK (char_length(prompt) >= 10),
|
||||
CONSTRAINT check_greeting_length CHECK (greeting IS NULL OR char_length(greeting) <= 500)
|
||||
);
|
||||
|
||||
-- 为用户 ID 创建索引,加速查询
|
||||
CREATE INDEX idx_user_scenarios_user_id ON user_scenarios(user_id);
|
||||
|
||||
-- 为创建时间创建索引,用于排序
|
||||
CREATE INDEX idx_user_scenarios_created_at ON user_scenarios(created_at DESC);
|
||||
|
||||
-- 表和列注释
|
||||
COMMENT ON TABLE user_scenarios IS '用户自建情景表,存储用户创建的 AI 对话情景配置';
|
||||
COMMENT ON COLUMN user_scenarios.id IS '情景唯一标识符 (UUID)';
|
||||
COMMENT ON COLUMN user_scenarios.user_id IS '所属用户 ID,外键关联 users 表,用户删除时级联删除';
|
||||
COMMENT ON COLUMN user_scenarios.name IS '情景名称 (2-50 字符),如"创意写作导师"';
|
||||
COMMENT ON COLUMN user_scenarios.icon IS 'Emoji 图标 (最多 20 字符),支持复合 Emoji,如"🎨"';
|
||||
COMMENT ON COLUMN user_scenarios.description IS '简短描述 (最多 100 字符),可选,显示在情景卡片上';
|
||||
COMMENT ON COLUMN user_scenarios.prompt IS '角色 System Prompt (最少 10 字符,无上限),定义 AI 行为和对话风格';
|
||||
COMMENT ON COLUMN user_scenarios.greeting IS '首句引导 (最多 500 字符),可选,AI 的开场白';
|
||||
COMMENT ON COLUMN user_scenarios.language IS '默认语言代码 (如 zh-CN、en-US、ja-JP)';
|
||||
COMMENT ON COLUMN user_scenarios.created_at IS '情景创建时间';
|
||||
COMMENT ON COLUMN user_scenarios.updated_at IS '情景最后更新时间';
|
||||
9
backend/migrations/embed.go
Normal file
9
backend/migrations/embed.go
Normal file
@@ -0,0 +1,9 @@
|
||||
// Package migrations 提供数据库迁移 SQL 文件的嵌入式访问。
|
||||
package migrations
|
||||
|
||||
import "embed"
|
||||
|
||||
// FS 包含所有迁移 SQL 文件。
|
||||
//
|
||||
//go:embed *.sql
|
||||
var FS embed.FS
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user