diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml index 01c00bb962..342f48dabb 100644 --- a/.github/workflows/backend-ci.yml +++ b/.github/workflows/backend-ci.yml @@ -17,6 +17,7 @@ jobs: go-version-file: backend/go.mod check-latest: false cache: true + cache-dependency-path: backend/go.sum - name: Verify Go version run: | go version | grep -q 'go1.26.1' @@ -36,6 +37,7 @@ jobs: go-version-file: backend/go.mod check-latest: false cache: true + cache-dependency-path: backend/go.sum - name: Verify Go version run: | go version | grep -q 'go1.26.1' diff --git a/.gitignore b/.gitignore index 297c1d6f03..f87dde6dc5 100644 --- a/.gitignore +++ b/.gitignore @@ -78,6 +78,7 @@ Desktop.ini # =================== tmp/ temp/ +logs/ *.tmp *.temp *.log @@ -127,9 +128,17 @@ deploy/docker-compose.override.yml .gocache/ vite.config.js docs/* +!docs/ACCOUNT_SCHEDULING_FLOW.md .serena/ + +# =================== +# 压测工具 +# =================== +tools/loadtest/ +# Antigravity Manager +Antigravity-Manager/ +antigravity_projectid_fix.patch .codex/ frontend/coverage/ aicodex output/ - diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000000..1a98b0f15e --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,1303 @@ +# PR 默认语义(最高优先级) + +- 未特别说明时,文档或对话中的“PR”一律指**提交到上游仓库** `Wei-Shaw/sub2api:main` +- 如果只是合并回我们自己的仓库,必须明确表述为: + - “内部同步 PR” + - “合并回我们的 main” + - “fork 内部 PR” +- 禁止将“上游 PR”和“我们自己仓库内的同步 PR”混用为同一个概念 + +## 本地依赖联调 + +- 本地 `go-sora2api` 仓库固定路径:`C:\Users\16790\GolandProjects\go-sora2api` +- 需要联调 `go-sora2api` 时,优先使用 `backend/go.mod` 的 `replace` 指向该本地路径,而不是使用 `git submodule` +- 联调完成后,如需提交或部署,再切换为 fork 仓库的明确 tag 或 commit + +--- +# Sub2API 开发说明 + +## 版本管理策略 + +### 版本号规则 + +我们在官方版本号后面添加自己的小版本号: + +- 官方版本:`v0.1.68` +- 我们的版本:`v0.1.68.1`、`v0.1.68.2`(递增) + +### 分支策略 + +| 分支 | 说明 | +|------|------| +| `main` | 我们的主分支,包含所有定制功能 | +| `release/custom-X.Y.Z` | 基于官方 `vX.Y.Z` 的发布分支 | +| `upstream/main` | 上游官方仓库 | + +--- + +## 发布流程(基于新官方版本) + +当官方发布新版本(如 `v0.1.69`)时: + +### 1. 同步上游并创建发布分支 + +```bash +# 获取上游最新代码 +git fetch upstream --tags + +# 基于官方标签创建新的发布分支 +git checkout v0.1.69 -b release/custom-0.1.69 + +# 合并我们的 main 分支(包含所有定制功能) +git merge main --no-edit + +# 解决可能的冲突后继续 +``` + +### 2. 更新版本号并打标签 + +```bash +# 更新版本号文件 +echo "0.1.69.1" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.69.1" + +# 打上我们自己的标签 +git tag v0.1.69.1 + +# 推送分支和标签 +git push origin release/custom-0.1.69 +git push origin v0.1.69.1 +``` + +### 3. 更新 main 分支 + +```bash +# 将发布分支合并回 main,保持 main 包含最新定制功能 +git checkout main +git merge release/custom-0.1.69 +git push origin main +``` + +--- + +## 热修复发布(在现有版本上修复) + +当需要在当前版本上发布修复时: + +```bash +# 在当前发布分支上修复 +git checkout release/custom-0.1.68 +# ... 进行修复 ... +git commit -m "fix: 修复描述" + +# 递增小版本号 +echo "0.1.68.2" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.68.2" + +# 打标签并推送 +git tag v0.1.68.2 +git push origin release/custom-0.1.68 +git push origin v0.1.68.2 + +# 同步修复到 main +git checkout main +git cherry-pick +git push origin main +``` + +--- + +## 服务器部署流程 + +### 前置条件 + +- 本地已配置 SSH 别名 `clicodeplus` 连接到生产服务器(运行服务 + 构建镜像) +- 生产服务器部署目录:`/root/sub2api`(正式)、`/root/sub2api-beta`(测试)、`/root/sub2api-star`(Star) +- 生产服务器使用 Docker Compose 部署 +- **镜像在生产服务器本机构建**,使用资源限制的 `limited-builder` 构建器(3 核 CPU、4G 内存),避免构建占满服务器资源影响线上服务 + +### 服务器角色说明 + +| 服务器 | SSH 别名 | 职责 | +|--------|----------|------| +| 生产服务器 | `clicodeplus` | 拉取代码、构建镜像、运行服务、部署验证 | +| 数据库服务器 | `db-clicodeplus` | PostgreSQL 16 + Redis 7,所有环境共用 | + +> 数据库服务器运维手册:`db-clicodeplus:/root/README.md` + +### 构建器说明 + +生产服务器上配置了资源限制的 Docker buildx 构建器 `limited-builder`,**所有构建操作必须使用此构建器**: + +- **构建器名称**:`limited-builder` +- **驱动**:`docker-container`(独立容器运行 BuildKit) +- **资源限制**:3 核 CPU、4G 内存(服务器共 6 核 8G,预留一半给线上服务) +- **容器名**:`buildx_buildkit_limited-builder0` + +```bash +# 构建命令格式(必须指定 --builder) +ssh clicodeplus "cd /root/sub2api && docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile ." + +# 查看构建器状态 +ssh clicodeplus "docker buildx inspect limited-builder" + +# 如果构建器容器被意外删除,重新创建: +ssh clicodeplus "docker buildx create --name limited-builder --driver docker-container --driver-opt 'default-load=true' && docker buildx inspect --builder limited-builder --bootstrap && docker update --cpus=3 --memory=4g --memory-swap=4g buildx_buildkit_limited-builder0" +``` + +### 部署环境说明 + +| 环境 | 目录(生产服务器) | 端口 | 数据库 | Redis DB | 容器名 | +|------|------|------|--------|----------|--------| +| 正式 | `/root/sub2api` | 8080 | `sub2api` | 0 | `sub2api` | +| Beta | `/root/sub2api-beta` | 8084 | `beta` | 2 | `sub2api-beta` | +| OpenAI | `/root/sub2api-openai` | 8083 | `openai` | 3 | `sub2api-openai` | +| Star | `/root/sub2api-star` | 8086 | `star` | 4 | `sub2api-star` | + +### 外部数据库与 Redis + +所有环境(正式、Beta、OpenAI、Star)共用 `db.clicodeplus.com` 上的 **PostgreSQL 16** 和 **Redis 7**,不使用容器内数据库或 Redis。 + +**PostgreSQL**(端口 5432,TLS 加密,scram-sha-256 认证): + +| 环境 | 用户名 | 数据库 | +|------|--------|--------| +| 正式 | `sub2api` | `sub2api` | +| Beta | `beta` | `beta` | +| OpenAI | `openai` | `openai` | +| Star | `star` | `star` | + +**Redis**(端口 6379,密码认证): + +| 环境 | DB | +|------|-----| +| 正式 | 0 | +| Beta | 2 | +| OpenAI | 3 | +| Star | 4 | + +**配置方式**: +- 数据库通过 `.env` 中的 `DATABASE_HOST`、`DATABASE_SSLMODE`、`POSTGRES_USER`、`POSTGRES_PASSWORD`、`POSTGRES_DB` 配置 +- Redis 通过 `docker-compose.override.yml` 覆盖 `REDIS_HOST`(因主 compose 文件硬编码为 `redis`),密码通过 `.env` 中的 `REDIS_PASSWORD` 配置 +- 各环境的 `docker-compose.override.yml` 已通过 `depends_on: !reset {}` 和 `redis: profiles: [disabled]` 去掉了对容器 Redis 的依赖 + +#### 数据库操作命令 + +通过 SSH 在服务器上执行数据库操作: + +```bash +# 正式环境 - 查询迁移记录 +ssh clicodeplus "source /root/sub2api/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c 'SELECT * FROM schema_migrations ORDER BY applied_at DESC LIMIT 5;'" + +# Beta 环境 - 查询迁移记录 +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c 'SELECT * FROM schema_migrations ORDER BY applied_at DESC LIMIT 5;'" + +# Beta 环境 - 清除指定迁移记录(重新执行迁移) +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c \"DELETE FROM schema_migrations WHERE filename LIKE '%049%';\"" + +# Beta 环境 - 更新账号数据 +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c \"UPDATE accounts SET credentials = credentials - 'model_mapping' WHERE platform = 'antigravity';\"" +``` + +> **注意**:使用 `source .env` 加载环境变量,避免在命令行中暴露密码。 + +### 部署步骤 + +**重要:每次部署都必须递增版本号!** + +#### 0. 递增版本号并推送(本地操作) + +每次部署前,先在本地递增小版本号并确保推送成功: + +```bash +# 查看当前版本号 +cat backend/cmd/server/VERSION +# 假设当前是 0.1.69.1 + +# 递增版本号 +echo "0.1.69.2" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.69.2" +git push origin release/custom-0.1.69 + +# ⚠️ 确认推送成功(必须看到分支更新输出,不能有 rejected 错误) +``` + +> **检查点**:如果有其他未提交的改动,应先 commit 并 push,确保 release 分支上的所有代码都已推送到远程。 + +#### 1. 生产服务器拉取代码 + +```bash +# 拉取最新代码并切换分支 +ssh clicodeplus "cd /root/sub2api && git fetch fork && git checkout -B release/custom-0.1.69 fork/release/custom-0.1.69" + +# ⚠️ 验证版本号与步骤 0 一致 +ssh clicodeplus "cat /root/sub2api/backend/cmd/server/VERSION" +``` + +#### 2. 生产服务器构建镜像(使用 limited-builder) + +```bash +ssh clicodeplus "cd /root/sub2api && docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile ." + +# ⚠️ 必须看到构建成功输出,如果失败需要先排查问题 +``` + +> **常见构建问题**: +> - 构建器未启动 → `docker buildx inspect --builder limited-builder --bootstrap` +> - 磁盘空间不足 → `docker system prune -f` 清理无用镜像 +> - 构建器被删除 → 参见上方「构建器说明」重新创建 + +#### 3. 更新镜像标签并重启 + +```bash +# 更新镜像标签并重启 +ssh clicodeplus "docker tag sub2api:latest weishaw/sub2api:latest" +ssh clicodeplus "cd /root/sub2api/deploy && docker compose up -d --force-recreate sub2api" +``` + +#### 4. 验证部署 + +```bash +# 查看启动日志 +ssh clicodeplus "docker logs sub2api --tail 20" + +# 确认版本号(必须与步骤 0 中设置的版本号一致) +ssh clicodeplus "cat /root/sub2api/backend/cmd/server/VERSION" + +# 检查容器状态(必须显示 healthy) +ssh clicodeplus "docker ps | grep sub2api" +``` + +--- + +## Beta 并行部署(不影响现网) + +目标:在同一台服务器上并行启动一个 beta 实例(例如端口 `8084`),**严禁改动/重启**现网实例(默认目录 `/root/sub2api`)。 + +### 设计原则 + +- **新目录**:beta 使用独立目录,例如 `/root/sub2api-beta`。 +- **敏感信息只放 `.env`**:beta 的数据库密码、JWT_SECRET 等只写入 `/root/sub2api-beta/deploy/.env`,不要提交到 git。 +- **独立 Compose Project**:通过 `docker compose -p sub2api-beta ...` 启动,确保 network/volume 隔离。 +- **独立端口**:通过 `.env` 的 `SERVER_PORT` 映射宿主机端口(例如 `8084:8080`)。 + +### 前置检查 + +```bash +# 1) 确保 8084 未被占用 +ssh clicodeplus "ss -ltnp | grep :8084 || echo '8084 is free'" + +# 2) 确认现网容器还在(只读检查) +ssh clicodeplus "docker ps --format 'table {{.Names}}\t{{.Image}}\t{{.Ports}}' | sed -n '1,200p'" +``` + +### 首次部署步骤 + +> **构建说明**:正式和 beta 通过不同的镜像标签区分(`sub2api:latest` 用于正式,`sub2api:beta` 用于测试),均在生产服务器本机使用 `limited-builder` 构建。 + +```bash +# 1) 在生产服务器上拉取代码并构建 beta 镜像 +ssh clicodeplus "cd /root/sub2api-beta && git fetch --all --tags && git checkout -f release/custom-0.1.71 && git reset --hard origin/release/custom-0.1.71" +ssh clicodeplus "cd /root/sub2api-beta && docker buildx build --builder limited-builder --no-cache --load -t sub2api:beta -f Dockerfile ." + +# 2) 在生产服务器上准备 beta 环境 +ssh clicodeplus + +# 克隆代码(仅用于 deploy 配置和版本号确认,不在此构建) +cd /root +git clone https://github.com/touwaeriol/sub2api.git sub2api-beta +cd /root/sub2api-beta +git checkout release/custom-0.1.71 + +# 4) 准备 beta 的 .env(敏感信息只写这里) +cd /root/sub2api-beta/deploy + +# 推荐:从现网 .env 复制,保证除 DB 名/用户/端口外完全一致 +cp -f /root/sub2api/deploy/.env ./.env + +# 仅修改以下三项(其他保持不变) +perl -pi -e 's/^SERVER_PORT=.*/SERVER_PORT=8084/' ./.env +perl -pi -e 's/^POSTGRES_USER=.*/POSTGRES_USER=beta/' ./.env +perl -pi -e 's/^POSTGRES_DB=.*/POSTGRES_DB=beta/' ./.env + +# 5) 写 compose override(避免与现网容器名冲突,镜像使用本机构建的 sub2api:beta,Redis 使用外部服务) +cat > docker-compose.override.yml <<'YAML' +services: + sub2api: + image: sub2api:beta + container_name: sub2api-beta + environment: + - DATABASE_HOST=${DATABASE_HOST:-postgres} + - DATABASE_SSLMODE=${DATABASE_SSLMODE:-disable} + - REDIS_HOST=db.clicodeplus.com + depends_on: !reset {} + redis: + profiles: + - disabled +YAML + +# 6) 启动 beta(独立 project,确保不影响现网) +cd /root/sub2api-beta/deploy +docker compose -p sub2api-beta --env-file .env -f docker-compose.yml -f docker-compose.override.yml up -d + +# 7) 验证 beta +curl -fsS http://127.0.0.1:8084/health +docker logs sub2api-beta --tail 50 +``` + +### 数据库配置约定(beta) + +- 数据库地址/SSL/密码:与现网一致(从现网 `.env` 复制即可),均指向 `db.clicodeplus.com`。 +- 仅修改: + - `POSTGRES_USER=beta` + - `POSTGRES_DB=beta` + - `REDIS_DB=2` + +注意:需要数据库侧已存在 `beta` 用户与 `beta` 数据库,并授予权限;否则容器会启动失败并不断重启。 + +### 更新 beta(本机构建 + 仅重启 beta 容器) + +```bash +# 1) 生产服务器拉取代码并构建镜像 +ssh clicodeplus "cd /root/sub2api-beta && git fetch --all --tags && git checkout -f release/custom-0.1.71 && git reset --hard origin/release/custom-0.1.71" +ssh clicodeplus "cd /root/sub2api-beta && docker buildx build --builder limited-builder --no-cache --load -t sub2api:beta -f Dockerfile ." +# ⚠️ 必须看到构建成功输出 + +# 2) 重启 beta 容器并验证 +ssh clicodeplus "cd /root/sub2api-beta/deploy && docker compose -p sub2api-beta --env-file .env -f docker-compose.yml -f docker-compose.override.yml up -d --no-deps --force-recreate sub2api" +ssh clicodeplus "sleep 5 && curl -fsS http://127.0.0.1:8084/health" +ssh clicodeplus "cat /root/sub2api-beta/backend/cmd/server/VERSION" +``` + +### 停止/回滚 beta(只影响 beta) + +```bash +ssh clicodeplus "cd /root/sub2api-beta/deploy && docker compose -p sub2api-beta -f docker-compose.yml -f docker-compose.override.yml down" +``` + +--- + +## 服务器首次部署 + +### 1. 生产服务器:克隆代码并配置环境 + +```bash +ssh clicodeplus +cd /root +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api + +# 添加 fork 仓库 +git remote add fork https://github.com/touwaeriol/sub2api.git +git fetch fork +git checkout -B release/custom-0.1.69 fork/release/custom-0.1.69 + +# 配置环境变量 +cd deploy +cp .env.example .env +vim .env # 配置 DATABASE_HOST=db.clicodeplus.com, POSTGRES_PASSWORD, REDIS_PASSWORD, JWT_SECRET 等 + +# 创建 override 文件(Redis 指向外部服务,去掉容器 Redis 依赖) +cat > docker-compose.override.yml <<'YAML' +services: + sub2api: + environment: + - REDIS_HOST=db.clicodeplus.com + depends_on: !reset {} + redis: + profiles: + - disabled +YAML +``` + +### 2. 生产服务器:创建构建器并构建镜像 + +```bash +# 创建资源限制的构建器(首次执行一次即可) +docker buildx create --name limited-builder --driver docker-container --driver-opt "default-load=true" +docker buildx inspect --builder limited-builder --bootstrap +docker update --cpus=3 --memory=4g --memory-swap=4g buildx_buildkit_limited-builder0 + +# 构建镜像 +cd /root/sub2api +docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile . + +# 更新镜像标签并启动 +docker tag sub2api:latest weishaw/sub2api:latest +cd /root/sub2api/deploy && docker compose up -d +``` + +### 3. 验证部署 + +```bash +# 查看应用日志 +docker logs sub2api --tail 50 + +# 检查健康状态 +curl http://localhost:8080/health + +# 确认版本号 +cat /root/sub2api/backend/cmd/server/VERSION +``` + +### 4. 常用运维命令 + +```bash +# 查看实时日志 +docker logs -f sub2api + +# 重启服务 +docker compose restart sub2api + +# 停止所有服务 +docker compose down + +# 停止并删除数据卷(慎用!会删除数据库数据) +docker compose down -v + +# 查看资源使用情况 +docker stats sub2api +``` + +--- + +## Admin API 接口文档 + +### ⚠️ API 操作流程规范 + +当收到操作正式环境 Web 界面的新需求,但文档中未记录对应 API 接口时,**必须按以下流程执行**: + +1. **探索接口**:通过代码库搜索路由定义(`backend/internal/server/routes/`)、Handler(`backend/internal/handler/admin/`)和请求结构体,确定正确的 API 端点、请求方法、请求体格式 +2. **更新文档**:将新发现的接口补充到本文档的 Admin API 接口文档章节中,包含端点、参数说明和 curl 示例 +3. **执行操作**:根据最新文档中记录的接口完成用户需求 + +> **目的**:避免每次遇到相同需求都重复探索代码库,确保 API 文档持续完善,后续操作可直接查阅文档执行。 + +--- + +### 认证方式 + +所有 Admin API 通过 `x-api-key` 请求头传递 Admin API Key 认证。 + +``` +x-api-key: admin-xxx +``` + +> **使用说明**:Admin API Key 统一存放在项目根目录 `.env` 文件的 `ADMIN_API_KEY` 变量中(该文件已被 `.gitignore` 排除,不会提交到代码库)。操作前先从 `.env` 读取密钥;若密钥失效(返回 401),应提示用户提供新的密钥并更新到 `.env` 中。Token 格式为 `admin-` + 64 位十六进制字符,在管理后台 `设置 > Admin API Key` 中生成。**请勿将实际 token 写入文档或代码中。** + +### 环境地址 + +| 环境 | 基础地址 | 说明 | +|------|----------|------| +| 正式 | `https://clicodeplus.com` | 生产环境 | +| Beta | `http://<服务器IP>:8084` | 仅内网访问 | +| OpenAI | `http://<服务器IP>:8083` | 仅内网访问 | +| Star | `https://hyntoken.com` | 独立环境 | + +> 以下接口文档中,`${BASE}` 代表环境基础地址,`${KEY}` 代表 `.env` 中的 `ADMIN_API_KEY`。操作前执行 `source .env` 或 `export KEY=$ADMIN_API_KEY` 加载。 + +--- + +### 1. 账号管理 + +#### 1.1 获取账号列表 + +``` +GET /api/v1/admin/accounts +``` + +**查询参数**: + +| 参数 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `platform` | string | 否 | 平台筛选:`antigravity` / `anthropic` / `openai` / `gemini` | +| `type` | string | 否 | 账号类型:`oauth` / `api_key` / `cookie` | +| `status` | string | 否 | 状态:`active` / `disabled` / `error` | +| `search` | string | 否 | 搜索关键词(名称、备注) | +| `page` | int | 否 | 页码,默认 1 | +| `page_size` | int | 否 | 每页数量,默认 20 | + +```bash +curl -s "${BASE}/api/v1/admin/accounts?platform=antigravity&page=1&page_size=100" \ + -H "x-api-key: ${KEY}" +``` + +**响应**: +```json +{ + "code": 0, + "message": "success", + "data": { + "items": [{"id": 1, "name": "xxx@gmail.com", "platform": "antigravity", "status": "active", ...}], + "total": 66 + } +} +``` + +#### 1.2 获取账号详情 + +``` +GET /api/v1/admin/accounts/:id +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1" -H "x-api-key: ${KEY}" +``` + +#### 1.3 测试账号连接 + +``` +POST /api/v1/admin/accounts/:id/test +``` + +**请求体**(JSON,可选): + +| 字段 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `model_id` | string | 否 | 指定测试模型,如 `claude-opus-4-6`;不传则使用默认模型 | + +**响应格式**:SSE(Server-Sent Events)流 + +```bash +curl -N -X POST "${BASE}/api/v1/admin/accounts/1/test" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"model_id": "claude-opus-4-6"}' +``` + +**SSE 事件类型**: + +| type | 字段 | 说明 | +|------|------|------| +| `test_start` | `model` | 测试开始,返回测试模型名 | +| `content` | `text` | 模型响应内容(流式文本片段) | +| `test_end` | `success`, `error` | 测试结束,`success=true` 表示成功 | +| `error` | `text` | 错误信息 | + +#### 1.4 清除账号限流 + +``` +POST /api/v1/admin/accounts/:id/clear-rate-limit +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/clear-rate-limit" \ + -H "x-api-key: ${KEY}" +``` + +#### 1.5 清除账号错误状态 + +``` +POST /api/v1/admin/accounts/:id/clear-error +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/clear-error" \ + -H "x-api-key: ${KEY}" +``` + +#### 1.6 获取账号可用模型 + +``` +GET /api/v1/admin/accounts/:id/models +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/models" -H "x-api-key: ${KEY}" +``` + +#### 1.7 刷新 OAuth Token + +``` +POST /api/v1/admin/accounts/:id/refresh +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/refresh" -H "x-api-key: ${KEY}" +``` + +#### 1.8 刷新账号等级 + +``` +POST /api/v1/admin/accounts/:id/refresh-tier +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/refresh-tier" -H "x-api-key: ${KEY}" +``` + +#### 1.9 获取账号统计 + +``` +GET /api/v1/admin/accounts/:id/stats +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/stats" -H "x-api-key: ${KEY}" +``` + +#### 1.10 获取账号用量 + +``` +GET /api/v1/admin/accounts/:id/usage +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/usage" -H "x-api-key: ${KEY}" +``` + +#### 1.11 更新单个账号 + +``` +PUT /api/v1/admin/accounts/:id +``` + +**请求体**(JSON,所有字段均为可选,仅传需要更新的字段): + +| 字段 | 类型 | 说明 | +|------|------|------| +| `name` | string | 账号名称 | +| `notes` | *string | 备注 | +| `type` | string | 类型:`oauth` / `setup-token` / `apikey` / `upstream` | +| `credentials` | object | 凭证信息 | +| `extra` | object | 额外配置 | +| `proxy_id` | *int64 | 代理 ID | +| `concurrency` | *int | 并发数 | +| `priority` | *int | 优先级(默认 50) | +| `rate_multiplier` | *float64 | 速率倍数 | +| `status` | string | 状态:`active` / `inactive` | +| `group_ids` | *[]int64 | 分组 ID 列表 | +| `expires_at` | *int64 | 过期时间戳 | +| `auto_pause_on_expired` | *bool | 过期后自动暂停 | + +> 使用指针类型(`*`)的字段可以区分"未提供"和"设置为零值"。 + +```bash +# 示例:更新账号优先级为 100 +curl -X PUT "${BASE}/api/v1/admin/accounts/1" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"priority": 100}' +``` + +#### 1.12 批量更新账号 + +``` +POST /api/v1/admin/accounts/bulk-update +``` + +**请求体**(JSON): + +| 字段 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `account_ids` | []int64 | **是** | 要更新的账号 ID 列表 | +| `priority` | *int | 否 | 优先级 | +| `concurrency` | *int | 否 | 并发数 | +| `rate_multiplier` | *float64 | 否 | 速率倍数 | +| `status` | string | 否 | 状态:`active` / `inactive` / `error` | +| `schedulable` | *bool | 否 | 是否可调度 | +| `group_ids` | *[]int64 | 否 | 分组 ID 列表 | +| `proxy_id` | *int64 | 否 | 代理 ID | +| `credentials` | object | 否 | 凭证信息(批量覆盖) | +| `extra` | object | 否 | 额外配置(批量覆盖) | + +```bash +# 示例:批量设置多个账号优先级为 100 +curl -X POST "${BASE}/api/v1/admin/accounts/bulk-update" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"account_ids": [1, 2, 3], "priority": 100}' +``` + +#### 1.13 批量测试账号(脚本) + +批量测试指定平台所有账号的指定模型连通性: + +```bash +# 用户需提供:BASE(环境地址)、KEY(admin token)、MODEL(测试模型) +ACCOUNT_IDS=$(curl -s "${BASE}/api/v1/admin/accounts?platform=antigravity&page=1&page_size=100" \ + -H "x-api-key: ${KEY}" | python3 -c " +import json, sys +data = json.load(sys.stdin) +for item in data['data']['items']: + print(f\"{item['id']}|{item['name']}\") +") + +while IFS='|' read -r ID NAME; do + echo "测试账号 ID=${ID} (${NAME})..." + RESPONSE=$(curl -s --max-time 60 -N \ + -X POST "${BASE}/api/v1/admin/accounts/${ID}/test" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d "{\"model_id\": \"${MODEL}\"}" 2>&1) + if echo "$RESPONSE" | grep -q '"success":true'; then + echo " ✅ 成功" + elif echo "$RESPONSE" | grep -q '"type":"content"'; then + echo " ✅ 成功(有内容响应)" + else + ERROR_MSG=$(echo "$RESPONSE" | grep -o '"error":"[^"]*"' | tail -1) + echo " ❌ 失败: ${ERROR_MSG}" + fi +done <<< "$ACCOUNT_IDS" +``` + +--- + +### 2. 运维监控 + +#### 2.1 并发统计 + +``` +GET /api/v1/admin/ops/concurrency +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/concurrency" -H "x-api-key: ${KEY}" +``` + +#### 2.2 账号可用性 + +``` +GET /api/v1/admin/ops/account-availability +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/account-availability" -H "x-api-key: ${KEY}" +``` + +#### 2.3 实时流量摘要 + +``` +GET /api/v1/admin/ops/realtime-traffic +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/realtime-traffic" -H "x-api-key: ${KEY}" +``` + +#### 2.4 请求错误列表 + +``` +GET /api/v1/admin/ops/request-errors +``` + +**查询参数**:`page`、`page_size` + +```bash +curl -s "${BASE}/api/v1/admin/ops/request-errors?page=1&page_size=50" \ + -H "x-api-key: ${KEY}" +``` + +#### 2.5 上游错误列表 + +``` +GET /api/v1/admin/ops/upstream-errors +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/upstream-errors?page=1&page_size=50" \ + -H "x-api-key: ${KEY}" +``` + +#### 2.6 仪表板概览 + +``` +GET /api/v1/admin/ops/dashboard/overview +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/dashboard/overview" -H "x-api-key: ${KEY}" +``` + +--- + +### 3. 系统设置 + +#### 3.1 获取系统设置 + +``` +GET /api/v1/admin/settings +``` + +```bash +curl -s "${BASE}/api/v1/admin/settings" -H "x-api-key: ${KEY}" +``` + +#### 3.2 更新系统设置 + +``` +PUT /api/v1/admin/settings +``` + +```bash +curl -X PUT "${BASE}/api/v1/admin/settings" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{ ... }' +``` + +#### 3.3 Admin API Key 状态(脱敏) + +``` +GET /api/v1/admin/settings/admin-api-key +``` + +```bash +curl -s "${BASE}/api/v1/admin/settings/admin-api-key" -H "x-api-key: ${KEY}" +``` + +--- + +### 4. 用户管理 + +#### 4.1 用户列表 + +``` +GET /api/v1/admin/users +``` + +```bash +curl -s "${BASE}/api/v1/admin/users?page=1&page_size=20" -H "x-api-key: ${KEY}" +``` + +#### 4.2 用户详情 + +``` +GET /api/v1/admin/users/:id +``` + +```bash +curl -s "${BASE}/api/v1/admin/users/1" -H "x-api-key: ${KEY}" +``` + +#### 4.3 更新用户余额 + +``` +POST /api/v1/admin/users/:id/balance +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/users/1/balance" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"amount": 100, "reason": "充值"}' +``` + +--- + +### 5. 分组管理 + +#### 5.1 分组列表 + +``` +GET /api/v1/admin/groups +``` + +```bash +curl -s "${BASE}/api/v1/admin/groups" -H "x-api-key: ${KEY}" +``` + +#### 5.2 所有分组(不分页) + +``` +GET /api/v1/admin/groups/all +``` + +```bash +curl -s "${BASE}/api/v1/admin/groups/all" -H "x-api-key: ${KEY}" +``` + +--- + +## 注意事项 + +1. **前端必须打包进镜像**:使用 `docker buildx build --builder limited-builder` 在生产服务器(`clicodeplus`)本机构建,Dockerfile 会自动编译前端并 embed 到后端二进制中 + +2. **镜像标签**:docker-compose.yml 使用 `weishaw/sub2api:latest`,本地构建后需要 `docker tag` 覆盖 + +3. **Windows 换行符问题**:已通过 `.gitattributes` 解决,确保 `*.sql` 文件始终使用 LF + +4. **版本号管理**:每次发布必须更新 `backend/cmd/server/VERSION` 并打标签 + +5. **合并冲突**:合并上游新版本时,重点关注以下文件可能的冲突: + - `backend/internal/service/antigravity_gateway_service.go` + - `backend/internal/service/gateway_service.go` + - `backend/internal/pkg/antigravity/request_transformer.go` + +--- + +## Go 代码规范 + +### 1. 函数设计 + +#### 单一职责原则 +- **函数行数**:单个函数常规不应超过 **30 行**,超过时应拆分为子函数。若某段逻辑确实不可拆分(如复杂的状态机、协议解析等),可以例外,但需添加注释说明原因 +- **嵌套层级**:避免超过 3 层嵌套,使用 early return 减少嵌套 + +```go +// ❌ 不推荐:深层嵌套 +func process(data []Item) { + for _, item := range data { + if item.Valid { + if item.Type == "A" { + if item.Status == "active" { + // 业务逻辑... + } + } + } + } +} + +// ✅ 推荐:early return +func process(data []Item) { + for _, item := range data { + if !item.Valid { + continue + } + if item.Type != "A" { + continue + } + if item.Status != "active" { + continue + } + // 业务逻辑... + } +} +``` + +#### 复杂逻辑提取 +将复杂的条件判断或处理逻辑提取为独立函数: + +```go +// ❌ 不推荐:内联复杂逻辑 +if resp.StatusCode == 429 || resp.StatusCode == 503 { + // 80+ 行处理逻辑... +} + +// ✅ 推荐:提取为独立函数 +result := handleRateLimitResponse(resp, params) +switch result.action { +case actionRetry: + continue +case actionBreak: + return result.resp, nil +} +``` + +### 2. 重复代码消除 + +#### 配置获取模式 +将重复的配置获取逻辑提取为方法: + +```go +// ❌ 不推荐:重复代码 +logBody := s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBody +maxBytes := 2048 +if s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > 0 { + maxBytes = s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes +} + +// ✅ 推荐:提取为方法 +func (s *Service) getLogConfig() (logBody bool, maxBytes int) { + maxBytes = 2048 + if s.settingService == nil || s.settingService.cfg == nil { + return false, maxBytes + } + cfg := s.settingService.cfg.Gateway + if cfg.LogUpstreamErrorBodyMaxBytes > 0 { + maxBytes = cfg.LogUpstreamErrorBodyMaxBytes + } + return cfg.LogUpstreamErrorBody, maxBytes +} +``` + +### 3. 常量管理 + +#### 避免魔法数字 +所有硬编码的数值都应定义为常量: + +```go +// ❌ 不推荐 +if retryDelay >= 10*time.Second { + resetAt := time.Now().Add(30 * time.Second) +} + +// ✅ 推荐 +const ( + rateLimitThreshold = 10 * time.Second + defaultRateLimitDuration = 30 * time.Second +) + +if retryDelay >= rateLimitThreshold { + resetAt := time.Now().Add(defaultRateLimitDuration) +} +``` + +#### 注释引用常量名 +在注释中引用常量名而非硬编码值: + +```go +// ❌ 不推荐 +// < 10s: 等待后重试 + +// ✅ 推荐 +// < rateLimitThreshold: 等待后重试 +``` + +### 4. 错误处理 + +#### 使用结构化日志 +优先使用 `slog` 进行结构化日志记录: + +```go +// ❌ 不推荐 +log.Printf("%s status=%d model_rate_limit_failed model=%s error=%v", prefix, statusCode, modelName, err) + +// ✅ 推荐 +slog.Error("failed to set model rate limit", + "prefix", prefix, + "status_code", statusCode, + "model", modelName, + "error", err, +) +``` + +### 5. 测试规范 + +#### Mock 函数签名同步 +修改函数签名时,必须同步更新所有测试中的 mock 函数: + +```go +// 如果修改了 handleError 签名 +handleError func(..., groupID int64, sessionHash string) *Result + +// 必须同步更新测试中的 mock +handleError: func(..., groupID int64, sessionHash string) *Result { + return nil +}, +``` + +#### 测试构建标签 +统一使用测试构建标签: + +```go +//go:build unit + +package service +``` + +### 6. 时间格式解析 + +#### 使用标准库 +优先使用 `time.ParseDuration`,支持所有 Go duration 格式: + +```go +// ❌ 不推荐:手动限制格式 +if !strings.HasSuffix(delay, "s") || strings.Contains(delay, "m") { + continue +} + +// ✅ 推荐:使用标准库 +dur, err := time.ParseDuration(delay) // 支持 "0.5s", "4m50s", "1h30m" 等 +``` + +### 7. 接口设计 + +#### 接口隔离原则 +定义最小化接口,只包含必需的方法: + +```go +// ❌ 不推荐:使用过于宽泛的接口 +type AccountRepository interface { + // 20+ 个方法... +} + +// ✅ 推荐:定义最小化接口 +type ModelRateLimiter interface { + SetModelRateLimit(ctx context.Context, id int64, modelKey string, resetAt time.Time) error +} +``` + +### 8. 并发安全 + +#### 共享数据保护 +访问可能被并发修改的数据时,确保线程安全: + +```go +// 如果 Account.Extra 可能被并发修改 +// 需要使用互斥锁或原子操作保护读取 +func (a *Account) GetRateLimitRemainingTime(model string) time.Duration { + a.mu.RLock() + defer a.mu.RUnlock() + // 读取 Extra 字段... +} +``` + +### 9. 命名规范 + +#### 一致的命名风格 +- 常量使用 camelCase:`rateLimitThreshold` +- 类型使用 PascalCase:`AntigravityQuotaScope` +- 同一概念使用统一命名:`Threshold` 或 `Limit`,不要混用 + +```go +// ❌ 不推荐:命名不一致 +antigravitySmartRetryMinWait // 使用 Min +antigravityRateLimitThreshold // 使用 Threshold + +// ✅ 推荐:统一风格 +antigravityMinRetryWait +antigravityRateLimitThreshold +``` + +### 10. 代码审查清单 + +在提交代码前,检查以下项目: + +- [ ] 函数是否超过 30 行?(不可拆分的逻辑除外,需注释说明) +- [ ] 嵌套是否超过 3 层? +- [ ] 是否有重复代码可以提取? +- [ ] 是否使用了魔法数字? +- [ ] Mock 函数签名是否与实际函数一致? +- [ ] 测试是否覆盖了新增逻辑? +- [ ] 日志是否包含足够的上下文信息? +- [ ] 是否考虑了并发安全? + +--- + +## CI 检查与发布门禁 + +### GitHub Actions 检查项 + +本项目有 4 个 CI 任务,**任何代码推送或发布前都必须全部通过**: + +| Workflow | Job | 说明 | 本地验证命令 | +|----------|-----|------|-------------| +| CI | `test` | 单元测试 + 集成测试 | `cd backend && make test-unit && make test-integration` | +| CI | `golangci-lint` | Go 代码静态检查(golangci-lint v2.7) | `cd backend && golangci-lint run --timeout=5m` | +| Security Scan | `backend-security` | govulncheck + gosec 安全扫描 | `cd backend && govulncheck ./... && gosec -severity high -confidence high ./...` | +| Security Scan | `frontend-security` | pnpm audit 前端依赖安全检查 | `cd frontend && pnpm audit --prod --audit-level=high` | + +### 向上游提交 PR + +PR 目标是上游官方仓库,**只包含通用功能改动**(bug fix、新功能、性能优化等)。 + +**以下文件禁止出现在 PR 中**(属于我们 fork 的定制化内容): +- `CLAUDE.md`、`AGENTS.md` — 我们的开发文档 +- `backend/cmd/server/VERSION` — 我们的版本号文件 +- UI 定制改动(GitHub 链接移除、微信客服按钮、首页定制等) +- 部署配置(`deploy/` 目录下的定制修改) + +**PR 流程**: +1. 从我们的当前开发分支(如 `release/custom-0.1.93`)或对应功能分支创建 PR 分支,只包含要提交给上游的通用改动 +2. 推送分支后,**等待 4 个 CI job 全部通过** +3. 确认通过后再创建 PR +4. 使用 `gh run list --repo touwaeriol/sub2api --branch ` 检查状态 + +### 自有分支推送(release/custom-X.Y.Z / 功能分支 / main) + +推送到我们自己的 `release/custom-X.Y.Z`、其他开发功能分支或 `main` 分支时,包含所有改动(定制化 + 通用功能)。 + +**推送前必须在本地执行全部 CI 检查**(不要等 GitHub Actions): + +```bash +# 确保 Go 工具链可用(macOS homebrew) +export PATH="/opt/homebrew/bin:$HOME/go/bin:$PATH" + +# 1. 单元测试(必须) +cd backend && make test-unit + +# 2. 集成测试(推荐,需要 Docker) +make test-integration + +# 3. golangci-lint 静态检查(必须) +golangci-lint run --timeout=5m + +# 4. gofmt 格式检查(必须) +gofmt -l ./... +# 如果有输出,运行 gofmt -w 修复 +``` + +**推送后确认**: +1. 使用 `gh run list --repo touwaeriol/sub2api --branch ` 检查 GitHub Actions 状态 +2. 确认 CI 和 Security Scan 两个 workflow 的 4 个 job 全部绿色 ✅ +3. 任何 job 失败必须立即修复,**禁止在 CI 未通过的状态下继续后续操作** + +### 发布版本 + +1. 本地执行上述全部 CI 检查通过 +2. 递增 `backend/cmd/server/VERSION`,提交并推送 +3. 推送后确认 GitHub Actions 的 4 个 CI job 全部通过 +4. **CI 未通过时禁止部署** — 必须先修复问题 +5. 使用 `gh run list --repo touwaeriol/sub2api --limit 10` 确认状态 + +### 常见 CI 失败原因及修复 +- **gofmt**:struct 字段对齐不一致 → 运行 `gofmt -w ` 修复 +- **golangci-lint**:未使用的变量/导入 → 删除或使用 `_` 忽略 +- **test 失败**:mock 函数签名不一致 → 同步更新 mock +- **gosec**:安全漏洞 → 根据提示修复或添加例外 + +--- + +## PR 描述格式规范 + +所有 PR 描述使用中英文同步(先中文、后英文),包含以下三个部分: + +### 模板 + +```markdown +## 背景 / Background + +<一两句说明问题现状或触发原因> + + + +--- + +## 目的 / Purpose + +<本次改动要解决的问题或达到的目标> + + + +--- + +## 改动内容 / Changes + +### 后端 / Backend + +- **改动点 1**:说明 +- **改动点 2**:说明 + +--- + +- **Change 1**: description +- **Change 2**: description + +### 前端 / Frontend + +- **改动点 1**:说明 +- **改动点 2**:说明 + +--- + +- **Change 1**: description +- **Change 2**: description + +--- + +## 截图 / Screenshot(可选) + +ASCII 示意图或实际截图 +``` + +### 规范要点 + +- **标题**:使用 conventional commits 格式,如 `feat(scope): description` +- **中英文顺序**:同一段落先中文后英文,用空行分隔,不用 `---` 分割同段内容 +- **改动分类**:按 Backend / Frontend / Config 等模块分组,先列中文要点再列英文要点 +- **截图/示意图**:有 UI 变动时必须附上,可用 ASCII 示意布局 +- **目标分支**:提交到 `touwaeriol/sub2api` 的 `main` 分支 diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000000..1a98b0f15e --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,1303 @@ +# PR 默认语义(最高优先级) + +- 未特别说明时,文档或对话中的“PR”一律指**提交到上游仓库** `Wei-Shaw/sub2api:main` +- 如果只是合并回我们自己的仓库,必须明确表述为: + - “内部同步 PR” + - “合并回我们的 main” + - “fork 内部 PR” +- 禁止将“上游 PR”和“我们自己仓库内的同步 PR”混用为同一个概念 + +## 本地依赖联调 + +- 本地 `go-sora2api` 仓库固定路径:`C:\Users\16790\GolandProjects\go-sora2api` +- 需要联调 `go-sora2api` 时,优先使用 `backend/go.mod` 的 `replace` 指向该本地路径,而不是使用 `git submodule` +- 联调完成后,如需提交或部署,再切换为 fork 仓库的明确 tag 或 commit + +--- +# Sub2API 开发说明 + +## 版本管理策略 + +### 版本号规则 + +我们在官方版本号后面添加自己的小版本号: + +- 官方版本:`v0.1.68` +- 我们的版本:`v0.1.68.1`、`v0.1.68.2`(递增) + +### 分支策略 + +| 分支 | 说明 | +|------|------| +| `main` | 我们的主分支,包含所有定制功能 | +| `release/custom-X.Y.Z` | 基于官方 `vX.Y.Z` 的发布分支 | +| `upstream/main` | 上游官方仓库 | + +--- + +## 发布流程(基于新官方版本) + +当官方发布新版本(如 `v0.1.69`)时: + +### 1. 同步上游并创建发布分支 + +```bash +# 获取上游最新代码 +git fetch upstream --tags + +# 基于官方标签创建新的发布分支 +git checkout v0.1.69 -b release/custom-0.1.69 + +# 合并我们的 main 分支(包含所有定制功能) +git merge main --no-edit + +# 解决可能的冲突后继续 +``` + +### 2. 更新版本号并打标签 + +```bash +# 更新版本号文件 +echo "0.1.69.1" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.69.1" + +# 打上我们自己的标签 +git tag v0.1.69.1 + +# 推送分支和标签 +git push origin release/custom-0.1.69 +git push origin v0.1.69.1 +``` + +### 3. 更新 main 分支 + +```bash +# 将发布分支合并回 main,保持 main 包含最新定制功能 +git checkout main +git merge release/custom-0.1.69 +git push origin main +``` + +--- + +## 热修复发布(在现有版本上修复) + +当需要在当前版本上发布修复时: + +```bash +# 在当前发布分支上修复 +git checkout release/custom-0.1.68 +# ... 进行修复 ... +git commit -m "fix: 修复描述" + +# 递增小版本号 +echo "0.1.68.2" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.68.2" + +# 打标签并推送 +git tag v0.1.68.2 +git push origin release/custom-0.1.68 +git push origin v0.1.68.2 + +# 同步修复到 main +git checkout main +git cherry-pick +git push origin main +``` + +--- + +## 服务器部署流程 + +### 前置条件 + +- 本地已配置 SSH 别名 `clicodeplus` 连接到生产服务器(运行服务 + 构建镜像) +- 生产服务器部署目录:`/root/sub2api`(正式)、`/root/sub2api-beta`(测试)、`/root/sub2api-star`(Star) +- 生产服务器使用 Docker Compose 部署 +- **镜像在生产服务器本机构建**,使用资源限制的 `limited-builder` 构建器(3 核 CPU、4G 内存),避免构建占满服务器资源影响线上服务 + +### 服务器角色说明 + +| 服务器 | SSH 别名 | 职责 | +|--------|----------|------| +| 生产服务器 | `clicodeplus` | 拉取代码、构建镜像、运行服务、部署验证 | +| 数据库服务器 | `db-clicodeplus` | PostgreSQL 16 + Redis 7,所有环境共用 | + +> 数据库服务器运维手册:`db-clicodeplus:/root/README.md` + +### 构建器说明 + +生产服务器上配置了资源限制的 Docker buildx 构建器 `limited-builder`,**所有构建操作必须使用此构建器**: + +- **构建器名称**:`limited-builder` +- **驱动**:`docker-container`(独立容器运行 BuildKit) +- **资源限制**:3 核 CPU、4G 内存(服务器共 6 核 8G,预留一半给线上服务) +- **容器名**:`buildx_buildkit_limited-builder0` + +```bash +# 构建命令格式(必须指定 --builder) +ssh clicodeplus "cd /root/sub2api && docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile ." + +# 查看构建器状态 +ssh clicodeplus "docker buildx inspect limited-builder" + +# 如果构建器容器被意外删除,重新创建: +ssh clicodeplus "docker buildx create --name limited-builder --driver docker-container --driver-opt 'default-load=true' && docker buildx inspect --builder limited-builder --bootstrap && docker update --cpus=3 --memory=4g --memory-swap=4g buildx_buildkit_limited-builder0" +``` + +### 部署环境说明 + +| 环境 | 目录(生产服务器) | 端口 | 数据库 | Redis DB | 容器名 | +|------|------|------|--------|----------|--------| +| 正式 | `/root/sub2api` | 8080 | `sub2api` | 0 | `sub2api` | +| Beta | `/root/sub2api-beta` | 8084 | `beta` | 2 | `sub2api-beta` | +| OpenAI | `/root/sub2api-openai` | 8083 | `openai` | 3 | `sub2api-openai` | +| Star | `/root/sub2api-star` | 8086 | `star` | 4 | `sub2api-star` | + +### 外部数据库与 Redis + +所有环境(正式、Beta、OpenAI、Star)共用 `db.clicodeplus.com` 上的 **PostgreSQL 16** 和 **Redis 7**,不使用容器内数据库或 Redis。 + +**PostgreSQL**(端口 5432,TLS 加密,scram-sha-256 认证): + +| 环境 | 用户名 | 数据库 | +|------|--------|--------| +| 正式 | `sub2api` | `sub2api` | +| Beta | `beta` | `beta` | +| OpenAI | `openai` | `openai` | +| Star | `star` | `star` | + +**Redis**(端口 6379,密码认证): + +| 环境 | DB | +|------|-----| +| 正式 | 0 | +| Beta | 2 | +| OpenAI | 3 | +| Star | 4 | + +**配置方式**: +- 数据库通过 `.env` 中的 `DATABASE_HOST`、`DATABASE_SSLMODE`、`POSTGRES_USER`、`POSTGRES_PASSWORD`、`POSTGRES_DB` 配置 +- Redis 通过 `docker-compose.override.yml` 覆盖 `REDIS_HOST`(因主 compose 文件硬编码为 `redis`),密码通过 `.env` 中的 `REDIS_PASSWORD` 配置 +- 各环境的 `docker-compose.override.yml` 已通过 `depends_on: !reset {}` 和 `redis: profiles: [disabled]` 去掉了对容器 Redis 的依赖 + +#### 数据库操作命令 + +通过 SSH 在服务器上执行数据库操作: + +```bash +# 正式环境 - 查询迁移记录 +ssh clicodeplus "source /root/sub2api/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c 'SELECT * FROM schema_migrations ORDER BY applied_at DESC LIMIT 5;'" + +# Beta 环境 - 查询迁移记录 +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c 'SELECT * FROM schema_migrations ORDER BY applied_at DESC LIMIT 5;'" + +# Beta 环境 - 清除指定迁移记录(重新执行迁移) +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c \"DELETE FROM schema_migrations WHERE filename LIKE '%049%';\"" + +# Beta 环境 - 更新账号数据 +ssh clicodeplus "source /root/sub2api-beta/deploy/.env && PGPASSWORD=\"\$POSTGRES_PASSWORD\" psql -h \$DATABASE_HOST -U \$POSTGRES_USER -d \$POSTGRES_DB -c \"UPDATE accounts SET credentials = credentials - 'model_mapping' WHERE platform = 'antigravity';\"" +``` + +> **注意**:使用 `source .env` 加载环境变量,避免在命令行中暴露密码。 + +### 部署步骤 + +**重要:每次部署都必须递增版本号!** + +#### 0. 递增版本号并推送(本地操作) + +每次部署前,先在本地递增小版本号并确保推送成功: + +```bash +# 查看当前版本号 +cat backend/cmd/server/VERSION +# 假设当前是 0.1.69.1 + +# 递增版本号 +echo "0.1.69.2" > backend/cmd/server/VERSION +git add backend/cmd/server/VERSION +git commit -m "chore: bump version to 0.1.69.2" +git push origin release/custom-0.1.69 + +# ⚠️ 确认推送成功(必须看到分支更新输出,不能有 rejected 错误) +``` + +> **检查点**:如果有其他未提交的改动,应先 commit 并 push,确保 release 分支上的所有代码都已推送到远程。 + +#### 1. 生产服务器拉取代码 + +```bash +# 拉取最新代码并切换分支 +ssh clicodeplus "cd /root/sub2api && git fetch fork && git checkout -B release/custom-0.1.69 fork/release/custom-0.1.69" + +# ⚠️ 验证版本号与步骤 0 一致 +ssh clicodeplus "cat /root/sub2api/backend/cmd/server/VERSION" +``` + +#### 2. 生产服务器构建镜像(使用 limited-builder) + +```bash +ssh clicodeplus "cd /root/sub2api && docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile ." + +# ⚠️ 必须看到构建成功输出,如果失败需要先排查问题 +``` + +> **常见构建问题**: +> - 构建器未启动 → `docker buildx inspect --builder limited-builder --bootstrap` +> - 磁盘空间不足 → `docker system prune -f` 清理无用镜像 +> - 构建器被删除 → 参见上方「构建器说明」重新创建 + +#### 3. 更新镜像标签并重启 + +```bash +# 更新镜像标签并重启 +ssh clicodeplus "docker tag sub2api:latest weishaw/sub2api:latest" +ssh clicodeplus "cd /root/sub2api/deploy && docker compose up -d --force-recreate sub2api" +``` + +#### 4. 验证部署 + +```bash +# 查看启动日志 +ssh clicodeplus "docker logs sub2api --tail 20" + +# 确认版本号(必须与步骤 0 中设置的版本号一致) +ssh clicodeplus "cat /root/sub2api/backend/cmd/server/VERSION" + +# 检查容器状态(必须显示 healthy) +ssh clicodeplus "docker ps | grep sub2api" +``` + +--- + +## Beta 并行部署(不影响现网) + +目标:在同一台服务器上并行启动一个 beta 实例(例如端口 `8084`),**严禁改动/重启**现网实例(默认目录 `/root/sub2api`)。 + +### 设计原则 + +- **新目录**:beta 使用独立目录,例如 `/root/sub2api-beta`。 +- **敏感信息只放 `.env`**:beta 的数据库密码、JWT_SECRET 等只写入 `/root/sub2api-beta/deploy/.env`,不要提交到 git。 +- **独立 Compose Project**:通过 `docker compose -p sub2api-beta ...` 启动,确保 network/volume 隔离。 +- **独立端口**:通过 `.env` 的 `SERVER_PORT` 映射宿主机端口(例如 `8084:8080`)。 + +### 前置检查 + +```bash +# 1) 确保 8084 未被占用 +ssh clicodeplus "ss -ltnp | grep :8084 || echo '8084 is free'" + +# 2) 确认现网容器还在(只读检查) +ssh clicodeplus "docker ps --format 'table {{.Names}}\t{{.Image}}\t{{.Ports}}' | sed -n '1,200p'" +``` + +### 首次部署步骤 + +> **构建说明**:正式和 beta 通过不同的镜像标签区分(`sub2api:latest` 用于正式,`sub2api:beta` 用于测试),均在生产服务器本机使用 `limited-builder` 构建。 + +```bash +# 1) 在生产服务器上拉取代码并构建 beta 镜像 +ssh clicodeplus "cd /root/sub2api-beta && git fetch --all --tags && git checkout -f release/custom-0.1.71 && git reset --hard origin/release/custom-0.1.71" +ssh clicodeplus "cd /root/sub2api-beta && docker buildx build --builder limited-builder --no-cache --load -t sub2api:beta -f Dockerfile ." + +# 2) 在生产服务器上准备 beta 环境 +ssh clicodeplus + +# 克隆代码(仅用于 deploy 配置和版本号确认,不在此构建) +cd /root +git clone https://github.com/touwaeriol/sub2api.git sub2api-beta +cd /root/sub2api-beta +git checkout release/custom-0.1.71 + +# 4) 准备 beta 的 .env(敏感信息只写这里) +cd /root/sub2api-beta/deploy + +# 推荐:从现网 .env 复制,保证除 DB 名/用户/端口外完全一致 +cp -f /root/sub2api/deploy/.env ./.env + +# 仅修改以下三项(其他保持不变) +perl -pi -e 's/^SERVER_PORT=.*/SERVER_PORT=8084/' ./.env +perl -pi -e 's/^POSTGRES_USER=.*/POSTGRES_USER=beta/' ./.env +perl -pi -e 's/^POSTGRES_DB=.*/POSTGRES_DB=beta/' ./.env + +# 5) 写 compose override(避免与现网容器名冲突,镜像使用本机构建的 sub2api:beta,Redis 使用外部服务) +cat > docker-compose.override.yml <<'YAML' +services: + sub2api: + image: sub2api:beta + container_name: sub2api-beta + environment: + - DATABASE_HOST=${DATABASE_HOST:-postgres} + - DATABASE_SSLMODE=${DATABASE_SSLMODE:-disable} + - REDIS_HOST=db.clicodeplus.com + depends_on: !reset {} + redis: + profiles: + - disabled +YAML + +# 6) 启动 beta(独立 project,确保不影响现网) +cd /root/sub2api-beta/deploy +docker compose -p sub2api-beta --env-file .env -f docker-compose.yml -f docker-compose.override.yml up -d + +# 7) 验证 beta +curl -fsS http://127.0.0.1:8084/health +docker logs sub2api-beta --tail 50 +``` + +### 数据库配置约定(beta) + +- 数据库地址/SSL/密码:与现网一致(从现网 `.env` 复制即可),均指向 `db.clicodeplus.com`。 +- 仅修改: + - `POSTGRES_USER=beta` + - `POSTGRES_DB=beta` + - `REDIS_DB=2` + +注意:需要数据库侧已存在 `beta` 用户与 `beta` 数据库,并授予权限;否则容器会启动失败并不断重启。 + +### 更新 beta(本机构建 + 仅重启 beta 容器) + +```bash +# 1) 生产服务器拉取代码并构建镜像 +ssh clicodeplus "cd /root/sub2api-beta && git fetch --all --tags && git checkout -f release/custom-0.1.71 && git reset --hard origin/release/custom-0.1.71" +ssh clicodeplus "cd /root/sub2api-beta && docker buildx build --builder limited-builder --no-cache --load -t sub2api:beta -f Dockerfile ." +# ⚠️ 必须看到构建成功输出 + +# 2) 重启 beta 容器并验证 +ssh clicodeplus "cd /root/sub2api-beta/deploy && docker compose -p sub2api-beta --env-file .env -f docker-compose.yml -f docker-compose.override.yml up -d --no-deps --force-recreate sub2api" +ssh clicodeplus "sleep 5 && curl -fsS http://127.0.0.1:8084/health" +ssh clicodeplus "cat /root/sub2api-beta/backend/cmd/server/VERSION" +``` + +### 停止/回滚 beta(只影响 beta) + +```bash +ssh clicodeplus "cd /root/sub2api-beta/deploy && docker compose -p sub2api-beta -f docker-compose.yml -f docker-compose.override.yml down" +``` + +--- + +## 服务器首次部署 + +### 1. 生产服务器:克隆代码并配置环境 + +```bash +ssh clicodeplus +cd /root +git clone https://github.com/Wei-Shaw/sub2api.git +cd sub2api + +# 添加 fork 仓库 +git remote add fork https://github.com/touwaeriol/sub2api.git +git fetch fork +git checkout -B release/custom-0.1.69 fork/release/custom-0.1.69 + +# 配置环境变量 +cd deploy +cp .env.example .env +vim .env # 配置 DATABASE_HOST=db.clicodeplus.com, POSTGRES_PASSWORD, REDIS_PASSWORD, JWT_SECRET 等 + +# 创建 override 文件(Redis 指向外部服务,去掉容器 Redis 依赖) +cat > docker-compose.override.yml <<'YAML' +services: + sub2api: + environment: + - REDIS_HOST=db.clicodeplus.com + depends_on: !reset {} + redis: + profiles: + - disabled +YAML +``` + +### 2. 生产服务器:创建构建器并构建镜像 + +```bash +# 创建资源限制的构建器(首次执行一次即可) +docker buildx create --name limited-builder --driver docker-container --driver-opt "default-load=true" +docker buildx inspect --builder limited-builder --bootstrap +docker update --cpus=3 --memory=4g --memory-swap=4g buildx_buildkit_limited-builder0 + +# 构建镜像 +cd /root/sub2api +docker buildx build --builder limited-builder --no-cache --load -t sub2api:latest -f Dockerfile . + +# 更新镜像标签并启动 +docker tag sub2api:latest weishaw/sub2api:latest +cd /root/sub2api/deploy && docker compose up -d +``` + +### 3. 验证部署 + +```bash +# 查看应用日志 +docker logs sub2api --tail 50 + +# 检查健康状态 +curl http://localhost:8080/health + +# 确认版本号 +cat /root/sub2api/backend/cmd/server/VERSION +``` + +### 4. 常用运维命令 + +```bash +# 查看实时日志 +docker logs -f sub2api + +# 重启服务 +docker compose restart sub2api + +# 停止所有服务 +docker compose down + +# 停止并删除数据卷(慎用!会删除数据库数据) +docker compose down -v + +# 查看资源使用情况 +docker stats sub2api +``` + +--- + +## Admin API 接口文档 + +### ⚠️ API 操作流程规范 + +当收到操作正式环境 Web 界面的新需求,但文档中未记录对应 API 接口时,**必须按以下流程执行**: + +1. **探索接口**:通过代码库搜索路由定义(`backend/internal/server/routes/`)、Handler(`backend/internal/handler/admin/`)和请求结构体,确定正确的 API 端点、请求方法、请求体格式 +2. **更新文档**:将新发现的接口补充到本文档的 Admin API 接口文档章节中,包含端点、参数说明和 curl 示例 +3. **执行操作**:根据最新文档中记录的接口完成用户需求 + +> **目的**:避免每次遇到相同需求都重复探索代码库,确保 API 文档持续完善,后续操作可直接查阅文档执行。 + +--- + +### 认证方式 + +所有 Admin API 通过 `x-api-key` 请求头传递 Admin API Key 认证。 + +``` +x-api-key: admin-xxx +``` + +> **使用说明**:Admin API Key 统一存放在项目根目录 `.env` 文件的 `ADMIN_API_KEY` 变量中(该文件已被 `.gitignore` 排除,不会提交到代码库)。操作前先从 `.env` 读取密钥;若密钥失效(返回 401),应提示用户提供新的密钥并更新到 `.env` 中。Token 格式为 `admin-` + 64 位十六进制字符,在管理后台 `设置 > Admin API Key` 中生成。**请勿将实际 token 写入文档或代码中。** + +### 环境地址 + +| 环境 | 基础地址 | 说明 | +|------|----------|------| +| 正式 | `https://clicodeplus.com` | 生产环境 | +| Beta | `http://<服务器IP>:8084` | 仅内网访问 | +| OpenAI | `http://<服务器IP>:8083` | 仅内网访问 | +| Star | `https://hyntoken.com` | 独立环境 | + +> 以下接口文档中,`${BASE}` 代表环境基础地址,`${KEY}` 代表 `.env` 中的 `ADMIN_API_KEY`。操作前执行 `source .env` 或 `export KEY=$ADMIN_API_KEY` 加载。 + +--- + +### 1. 账号管理 + +#### 1.1 获取账号列表 + +``` +GET /api/v1/admin/accounts +``` + +**查询参数**: + +| 参数 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `platform` | string | 否 | 平台筛选:`antigravity` / `anthropic` / `openai` / `gemini` | +| `type` | string | 否 | 账号类型:`oauth` / `api_key` / `cookie` | +| `status` | string | 否 | 状态:`active` / `disabled` / `error` | +| `search` | string | 否 | 搜索关键词(名称、备注) | +| `page` | int | 否 | 页码,默认 1 | +| `page_size` | int | 否 | 每页数量,默认 20 | + +```bash +curl -s "${BASE}/api/v1/admin/accounts?platform=antigravity&page=1&page_size=100" \ + -H "x-api-key: ${KEY}" +``` + +**响应**: +```json +{ + "code": 0, + "message": "success", + "data": { + "items": [{"id": 1, "name": "xxx@gmail.com", "platform": "antigravity", "status": "active", ...}], + "total": 66 + } +} +``` + +#### 1.2 获取账号详情 + +``` +GET /api/v1/admin/accounts/:id +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1" -H "x-api-key: ${KEY}" +``` + +#### 1.3 测试账号连接 + +``` +POST /api/v1/admin/accounts/:id/test +``` + +**请求体**(JSON,可选): + +| 字段 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `model_id` | string | 否 | 指定测试模型,如 `claude-opus-4-6`;不传则使用默认模型 | + +**响应格式**:SSE(Server-Sent Events)流 + +```bash +curl -N -X POST "${BASE}/api/v1/admin/accounts/1/test" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"model_id": "claude-opus-4-6"}' +``` + +**SSE 事件类型**: + +| type | 字段 | 说明 | +|------|------|------| +| `test_start` | `model` | 测试开始,返回测试模型名 | +| `content` | `text` | 模型响应内容(流式文本片段) | +| `test_end` | `success`, `error` | 测试结束,`success=true` 表示成功 | +| `error` | `text` | 错误信息 | + +#### 1.4 清除账号限流 + +``` +POST /api/v1/admin/accounts/:id/clear-rate-limit +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/clear-rate-limit" \ + -H "x-api-key: ${KEY}" +``` + +#### 1.5 清除账号错误状态 + +``` +POST /api/v1/admin/accounts/:id/clear-error +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/clear-error" \ + -H "x-api-key: ${KEY}" +``` + +#### 1.6 获取账号可用模型 + +``` +GET /api/v1/admin/accounts/:id/models +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/models" -H "x-api-key: ${KEY}" +``` + +#### 1.7 刷新 OAuth Token + +``` +POST /api/v1/admin/accounts/:id/refresh +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/refresh" -H "x-api-key: ${KEY}" +``` + +#### 1.8 刷新账号等级 + +``` +POST /api/v1/admin/accounts/:id/refresh-tier +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/accounts/1/refresh-tier" -H "x-api-key: ${KEY}" +``` + +#### 1.9 获取账号统计 + +``` +GET /api/v1/admin/accounts/:id/stats +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/stats" -H "x-api-key: ${KEY}" +``` + +#### 1.10 获取账号用量 + +``` +GET /api/v1/admin/accounts/:id/usage +``` + +```bash +curl -s "${BASE}/api/v1/admin/accounts/1/usage" -H "x-api-key: ${KEY}" +``` + +#### 1.11 更新单个账号 + +``` +PUT /api/v1/admin/accounts/:id +``` + +**请求体**(JSON,所有字段均为可选,仅传需要更新的字段): + +| 字段 | 类型 | 说明 | +|------|------|------| +| `name` | string | 账号名称 | +| `notes` | *string | 备注 | +| `type` | string | 类型:`oauth` / `setup-token` / `apikey` / `upstream` | +| `credentials` | object | 凭证信息 | +| `extra` | object | 额外配置 | +| `proxy_id` | *int64 | 代理 ID | +| `concurrency` | *int | 并发数 | +| `priority` | *int | 优先级(默认 50) | +| `rate_multiplier` | *float64 | 速率倍数 | +| `status` | string | 状态:`active` / `inactive` | +| `group_ids` | *[]int64 | 分组 ID 列表 | +| `expires_at` | *int64 | 过期时间戳 | +| `auto_pause_on_expired` | *bool | 过期后自动暂停 | + +> 使用指针类型(`*`)的字段可以区分"未提供"和"设置为零值"。 + +```bash +# 示例:更新账号优先级为 100 +curl -X PUT "${BASE}/api/v1/admin/accounts/1" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"priority": 100}' +``` + +#### 1.12 批量更新账号 + +``` +POST /api/v1/admin/accounts/bulk-update +``` + +**请求体**(JSON): + +| 字段 | 类型 | 必填 | 说明 | +|------|------|------|------| +| `account_ids` | []int64 | **是** | 要更新的账号 ID 列表 | +| `priority` | *int | 否 | 优先级 | +| `concurrency` | *int | 否 | 并发数 | +| `rate_multiplier` | *float64 | 否 | 速率倍数 | +| `status` | string | 否 | 状态:`active` / `inactive` / `error` | +| `schedulable` | *bool | 否 | 是否可调度 | +| `group_ids` | *[]int64 | 否 | 分组 ID 列表 | +| `proxy_id` | *int64 | 否 | 代理 ID | +| `credentials` | object | 否 | 凭证信息(批量覆盖) | +| `extra` | object | 否 | 额外配置(批量覆盖) | + +```bash +# 示例:批量设置多个账号优先级为 100 +curl -X POST "${BASE}/api/v1/admin/accounts/bulk-update" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"account_ids": [1, 2, 3], "priority": 100}' +``` + +#### 1.13 批量测试账号(脚本) + +批量测试指定平台所有账号的指定模型连通性: + +```bash +# 用户需提供:BASE(环境地址)、KEY(admin token)、MODEL(测试模型) +ACCOUNT_IDS=$(curl -s "${BASE}/api/v1/admin/accounts?platform=antigravity&page=1&page_size=100" \ + -H "x-api-key: ${KEY}" | python3 -c " +import json, sys +data = json.load(sys.stdin) +for item in data['data']['items']: + print(f\"{item['id']}|{item['name']}\") +") + +while IFS='|' read -r ID NAME; do + echo "测试账号 ID=${ID} (${NAME})..." + RESPONSE=$(curl -s --max-time 60 -N \ + -X POST "${BASE}/api/v1/admin/accounts/${ID}/test" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d "{\"model_id\": \"${MODEL}\"}" 2>&1) + if echo "$RESPONSE" | grep -q '"success":true'; then + echo " ✅ 成功" + elif echo "$RESPONSE" | grep -q '"type":"content"'; then + echo " ✅ 成功(有内容响应)" + else + ERROR_MSG=$(echo "$RESPONSE" | grep -o '"error":"[^"]*"' | tail -1) + echo " ❌ 失败: ${ERROR_MSG}" + fi +done <<< "$ACCOUNT_IDS" +``` + +--- + +### 2. 运维监控 + +#### 2.1 并发统计 + +``` +GET /api/v1/admin/ops/concurrency +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/concurrency" -H "x-api-key: ${KEY}" +``` + +#### 2.2 账号可用性 + +``` +GET /api/v1/admin/ops/account-availability +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/account-availability" -H "x-api-key: ${KEY}" +``` + +#### 2.3 实时流量摘要 + +``` +GET /api/v1/admin/ops/realtime-traffic +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/realtime-traffic" -H "x-api-key: ${KEY}" +``` + +#### 2.4 请求错误列表 + +``` +GET /api/v1/admin/ops/request-errors +``` + +**查询参数**:`page`、`page_size` + +```bash +curl -s "${BASE}/api/v1/admin/ops/request-errors?page=1&page_size=50" \ + -H "x-api-key: ${KEY}" +``` + +#### 2.5 上游错误列表 + +``` +GET /api/v1/admin/ops/upstream-errors +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/upstream-errors?page=1&page_size=50" \ + -H "x-api-key: ${KEY}" +``` + +#### 2.6 仪表板概览 + +``` +GET /api/v1/admin/ops/dashboard/overview +``` + +```bash +curl -s "${BASE}/api/v1/admin/ops/dashboard/overview" -H "x-api-key: ${KEY}" +``` + +--- + +### 3. 系统设置 + +#### 3.1 获取系统设置 + +``` +GET /api/v1/admin/settings +``` + +```bash +curl -s "${BASE}/api/v1/admin/settings" -H "x-api-key: ${KEY}" +``` + +#### 3.2 更新系统设置 + +``` +PUT /api/v1/admin/settings +``` + +```bash +curl -X PUT "${BASE}/api/v1/admin/settings" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{ ... }' +``` + +#### 3.3 Admin API Key 状态(脱敏) + +``` +GET /api/v1/admin/settings/admin-api-key +``` + +```bash +curl -s "${BASE}/api/v1/admin/settings/admin-api-key" -H "x-api-key: ${KEY}" +``` + +--- + +### 4. 用户管理 + +#### 4.1 用户列表 + +``` +GET /api/v1/admin/users +``` + +```bash +curl -s "${BASE}/api/v1/admin/users?page=1&page_size=20" -H "x-api-key: ${KEY}" +``` + +#### 4.2 用户详情 + +``` +GET /api/v1/admin/users/:id +``` + +```bash +curl -s "${BASE}/api/v1/admin/users/1" -H "x-api-key: ${KEY}" +``` + +#### 4.3 更新用户余额 + +``` +POST /api/v1/admin/users/:id/balance +``` + +```bash +curl -X POST "${BASE}/api/v1/admin/users/1/balance" \ + -H "x-api-key: ${KEY}" \ + -H "Content-Type: application/json" \ + -d '{"amount": 100, "reason": "充值"}' +``` + +--- + +### 5. 分组管理 + +#### 5.1 分组列表 + +``` +GET /api/v1/admin/groups +``` + +```bash +curl -s "${BASE}/api/v1/admin/groups" -H "x-api-key: ${KEY}" +``` + +#### 5.2 所有分组(不分页) + +``` +GET /api/v1/admin/groups/all +``` + +```bash +curl -s "${BASE}/api/v1/admin/groups/all" -H "x-api-key: ${KEY}" +``` + +--- + +## 注意事项 + +1. **前端必须打包进镜像**:使用 `docker buildx build --builder limited-builder` 在生产服务器(`clicodeplus`)本机构建,Dockerfile 会自动编译前端并 embed 到后端二进制中 + +2. **镜像标签**:docker-compose.yml 使用 `weishaw/sub2api:latest`,本地构建后需要 `docker tag` 覆盖 + +3. **Windows 换行符问题**:已通过 `.gitattributes` 解决,确保 `*.sql` 文件始终使用 LF + +4. **版本号管理**:每次发布必须更新 `backend/cmd/server/VERSION` 并打标签 + +5. **合并冲突**:合并上游新版本时,重点关注以下文件可能的冲突: + - `backend/internal/service/antigravity_gateway_service.go` + - `backend/internal/service/gateway_service.go` + - `backend/internal/pkg/antigravity/request_transformer.go` + +--- + +## Go 代码规范 + +### 1. 函数设计 + +#### 单一职责原则 +- **函数行数**:单个函数常规不应超过 **30 行**,超过时应拆分为子函数。若某段逻辑确实不可拆分(如复杂的状态机、协议解析等),可以例外,但需添加注释说明原因 +- **嵌套层级**:避免超过 3 层嵌套,使用 early return 减少嵌套 + +```go +// ❌ 不推荐:深层嵌套 +func process(data []Item) { + for _, item := range data { + if item.Valid { + if item.Type == "A" { + if item.Status == "active" { + // 业务逻辑... + } + } + } + } +} + +// ✅ 推荐:early return +func process(data []Item) { + for _, item := range data { + if !item.Valid { + continue + } + if item.Type != "A" { + continue + } + if item.Status != "active" { + continue + } + // 业务逻辑... + } +} +``` + +#### 复杂逻辑提取 +将复杂的条件判断或处理逻辑提取为独立函数: + +```go +// ❌ 不推荐:内联复杂逻辑 +if resp.StatusCode == 429 || resp.StatusCode == 503 { + // 80+ 行处理逻辑... +} + +// ✅ 推荐:提取为独立函数 +result := handleRateLimitResponse(resp, params) +switch result.action { +case actionRetry: + continue +case actionBreak: + return result.resp, nil +} +``` + +### 2. 重复代码消除 + +#### 配置获取模式 +将重复的配置获取逻辑提取为方法: + +```go +// ❌ 不推荐:重复代码 +logBody := s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBody +maxBytes := 2048 +if s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > 0 { + maxBytes = s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes +} + +// ✅ 推荐:提取为方法 +func (s *Service) getLogConfig() (logBody bool, maxBytes int) { + maxBytes = 2048 + if s.settingService == nil || s.settingService.cfg == nil { + return false, maxBytes + } + cfg := s.settingService.cfg.Gateway + if cfg.LogUpstreamErrorBodyMaxBytes > 0 { + maxBytes = cfg.LogUpstreamErrorBodyMaxBytes + } + return cfg.LogUpstreamErrorBody, maxBytes +} +``` + +### 3. 常量管理 + +#### 避免魔法数字 +所有硬编码的数值都应定义为常量: + +```go +// ❌ 不推荐 +if retryDelay >= 10*time.Second { + resetAt := time.Now().Add(30 * time.Second) +} + +// ✅ 推荐 +const ( + rateLimitThreshold = 10 * time.Second + defaultRateLimitDuration = 30 * time.Second +) + +if retryDelay >= rateLimitThreshold { + resetAt := time.Now().Add(defaultRateLimitDuration) +} +``` + +#### 注释引用常量名 +在注释中引用常量名而非硬编码值: + +```go +// ❌ 不推荐 +// < 10s: 等待后重试 + +// ✅ 推荐 +// < rateLimitThreshold: 等待后重试 +``` + +### 4. 错误处理 + +#### 使用结构化日志 +优先使用 `slog` 进行结构化日志记录: + +```go +// ❌ 不推荐 +log.Printf("%s status=%d model_rate_limit_failed model=%s error=%v", prefix, statusCode, modelName, err) + +// ✅ 推荐 +slog.Error("failed to set model rate limit", + "prefix", prefix, + "status_code", statusCode, + "model", modelName, + "error", err, +) +``` + +### 5. 测试规范 + +#### Mock 函数签名同步 +修改函数签名时,必须同步更新所有测试中的 mock 函数: + +```go +// 如果修改了 handleError 签名 +handleError func(..., groupID int64, sessionHash string) *Result + +// 必须同步更新测试中的 mock +handleError: func(..., groupID int64, sessionHash string) *Result { + return nil +}, +``` + +#### 测试构建标签 +统一使用测试构建标签: + +```go +//go:build unit + +package service +``` + +### 6. 时间格式解析 + +#### 使用标准库 +优先使用 `time.ParseDuration`,支持所有 Go duration 格式: + +```go +// ❌ 不推荐:手动限制格式 +if !strings.HasSuffix(delay, "s") || strings.Contains(delay, "m") { + continue +} + +// ✅ 推荐:使用标准库 +dur, err := time.ParseDuration(delay) // 支持 "0.5s", "4m50s", "1h30m" 等 +``` + +### 7. 接口设计 + +#### 接口隔离原则 +定义最小化接口,只包含必需的方法: + +```go +// ❌ 不推荐:使用过于宽泛的接口 +type AccountRepository interface { + // 20+ 个方法... +} + +// ✅ 推荐:定义最小化接口 +type ModelRateLimiter interface { + SetModelRateLimit(ctx context.Context, id int64, modelKey string, resetAt time.Time) error +} +``` + +### 8. 并发安全 + +#### 共享数据保护 +访问可能被并发修改的数据时,确保线程安全: + +```go +// 如果 Account.Extra 可能被并发修改 +// 需要使用互斥锁或原子操作保护读取 +func (a *Account) GetRateLimitRemainingTime(model string) time.Duration { + a.mu.RLock() + defer a.mu.RUnlock() + // 读取 Extra 字段... +} +``` + +### 9. 命名规范 + +#### 一致的命名风格 +- 常量使用 camelCase:`rateLimitThreshold` +- 类型使用 PascalCase:`AntigravityQuotaScope` +- 同一概念使用统一命名:`Threshold` 或 `Limit`,不要混用 + +```go +// ❌ 不推荐:命名不一致 +antigravitySmartRetryMinWait // 使用 Min +antigravityRateLimitThreshold // 使用 Threshold + +// ✅ 推荐:统一风格 +antigravityMinRetryWait +antigravityRateLimitThreshold +``` + +### 10. 代码审查清单 + +在提交代码前,检查以下项目: + +- [ ] 函数是否超过 30 行?(不可拆分的逻辑除外,需注释说明) +- [ ] 嵌套是否超过 3 层? +- [ ] 是否有重复代码可以提取? +- [ ] 是否使用了魔法数字? +- [ ] Mock 函数签名是否与实际函数一致? +- [ ] 测试是否覆盖了新增逻辑? +- [ ] 日志是否包含足够的上下文信息? +- [ ] 是否考虑了并发安全? + +--- + +## CI 检查与发布门禁 + +### GitHub Actions 检查项 + +本项目有 4 个 CI 任务,**任何代码推送或发布前都必须全部通过**: + +| Workflow | Job | 说明 | 本地验证命令 | +|----------|-----|------|-------------| +| CI | `test` | 单元测试 + 集成测试 | `cd backend && make test-unit && make test-integration` | +| CI | `golangci-lint` | Go 代码静态检查(golangci-lint v2.7) | `cd backend && golangci-lint run --timeout=5m` | +| Security Scan | `backend-security` | govulncheck + gosec 安全扫描 | `cd backend && govulncheck ./... && gosec -severity high -confidence high ./...` | +| Security Scan | `frontend-security` | pnpm audit 前端依赖安全检查 | `cd frontend && pnpm audit --prod --audit-level=high` | + +### 向上游提交 PR + +PR 目标是上游官方仓库,**只包含通用功能改动**(bug fix、新功能、性能优化等)。 + +**以下文件禁止出现在 PR 中**(属于我们 fork 的定制化内容): +- `CLAUDE.md`、`AGENTS.md` — 我们的开发文档 +- `backend/cmd/server/VERSION` — 我们的版本号文件 +- UI 定制改动(GitHub 链接移除、微信客服按钮、首页定制等) +- 部署配置(`deploy/` 目录下的定制修改) + +**PR 流程**: +1. 从我们的当前开发分支(如 `release/custom-0.1.93`)或对应功能分支创建 PR 分支,只包含要提交给上游的通用改动 +2. 推送分支后,**等待 4 个 CI job 全部通过** +3. 确认通过后再创建 PR +4. 使用 `gh run list --repo touwaeriol/sub2api --branch ` 检查状态 + +### 自有分支推送(release/custom-X.Y.Z / 功能分支 / main) + +推送到我们自己的 `release/custom-X.Y.Z`、其他开发功能分支或 `main` 分支时,包含所有改动(定制化 + 通用功能)。 + +**推送前必须在本地执行全部 CI 检查**(不要等 GitHub Actions): + +```bash +# 确保 Go 工具链可用(macOS homebrew) +export PATH="/opt/homebrew/bin:$HOME/go/bin:$PATH" + +# 1. 单元测试(必须) +cd backend && make test-unit + +# 2. 集成测试(推荐,需要 Docker) +make test-integration + +# 3. golangci-lint 静态检查(必须) +golangci-lint run --timeout=5m + +# 4. gofmt 格式检查(必须) +gofmt -l ./... +# 如果有输出,运行 gofmt -w 修复 +``` + +**推送后确认**: +1. 使用 `gh run list --repo touwaeriol/sub2api --branch ` 检查 GitHub Actions 状态 +2. 确认 CI 和 Security Scan 两个 workflow 的 4 个 job 全部绿色 ✅ +3. 任何 job 失败必须立即修复,**禁止在 CI 未通过的状态下继续后续操作** + +### 发布版本 + +1. 本地执行上述全部 CI 检查通过 +2. 递增 `backend/cmd/server/VERSION`,提交并推送 +3. 推送后确认 GitHub Actions 的 4 个 CI job 全部通过 +4. **CI 未通过时禁止部署** — 必须先修复问题 +5. 使用 `gh run list --repo touwaeriol/sub2api --limit 10` 确认状态 + +### 常见 CI 失败原因及修复 +- **gofmt**:struct 字段对齐不一致 → 运行 `gofmt -w ` 修复 +- **golangci-lint**:未使用的变量/导入 → 删除或使用 `_` 忽略 +- **test 失败**:mock 函数签名不一致 → 同步更新 mock +- **gosec**:安全漏洞 → 根据提示修复或添加例外 + +--- + +## PR 描述格式规范 + +所有 PR 描述使用中英文同步(先中文、后英文),包含以下三个部分: + +### 模板 + +```markdown +## 背景 / Background + +<一两句说明问题现状或触发原因> + + + +--- + +## 目的 / Purpose + +<本次改动要解决的问题或达到的目标> + + + +--- + +## 改动内容 / Changes + +### 后端 / Backend + +- **改动点 1**:说明 +- **改动点 2**:说明 + +--- + +- **Change 1**: description +- **Change 2**: description + +### 前端 / Frontend + +- **改动点 1**:说明 +- **改动点 2**:说明 + +--- + +- **Change 1**: description +- **Change 2**: description + +--- + +## 截图 / Screenshot(可选) + +ASCII 示意图或实际截图 +``` + +### 规范要点 + +- **标题**:使用 conventional commits 格式,如 `feat(scope): description` +- **中英文顺序**:同一段落先中文后英文,用空行分隔,不用 `---` 分割同段内容 +- **改动分类**:按 Backend / Frontend / Config 等模块分组,先列中文要点再列英文要点 +- **截图/示意图**:有 UI 变动时必须附上,可用 ASCII 示意布局 +- **目标分支**:提交到 `touwaeriol/sub2api` 的 `main` 分支 diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 32844913ea..3b9a8d7662 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.88 \ No newline at end of file +0.1.102.5 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 48f15b5cb9..acd690c5bb 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -144,7 +144,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { rpmCache := repository.NewRPMCache(redisClient) groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache) groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService) - accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator) + accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, gatewayCache) adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService) dataManagementService := service.NewDataManagementService() dataManagementHandler := admin.NewDataManagementHandler(dataManagementService) @@ -177,11 +177,15 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository) opsService := service.NewOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink) soraS3Storage := service.NewSoraS3Storage(settingService) - settingService.SetOnS3UpdateCallback(soraS3Storage.RefreshClient) + soraGDriveStorage := service.NewSoraGDriveStorage(settingService) + soraStorageRouter := service.NewSoraStorageRouter(settingService, soraS3Storage, soraGDriveStorage) + settingService.SetOnS3UpdateCallback(soraStorageRouter.RefreshAll) soraGenerationRepository := repository.NewSoraGenerationRepository(db) soraQuotaService := service.NewSoraQuotaService(userRepository, groupRepository, settingService) - soraGenerationService := service.NewSoraGenerationService(soraGenerationRepository, soraS3Storage, soraQuotaService) - settingHandler := admin.NewSettingHandler(settingService, emailService, turnstileService, opsService, soraS3Storage) + soraGenerationService := service.NewSoraGenerationService(soraGenerationRepository, soraStorageRouter, soraQuotaService) + settingHandler := admin.NewSettingHandler(settingService, emailService, turnstileService, opsService, soraS3Storage, soraGDriveStorage, soraGenerationService) + soraGDriveOAuthService := service.NewSoraGDriveOAuthService(settingService) + gdriveOAuthHandler := admin.NewGDriveOAuthHandler(settingService, soraGDriveOAuthService, soraGDriveStorage) opsHandler := admin.NewOpsHandler(opsService) updateCache := repository.NewUpdateCache(redisClient) gitHubReleaseClient := repository.ProvideGitHubReleaseClient(configConfig) @@ -207,7 +211,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { scheduledTestResultRepository := repository.NewScheduledTestResultRepository(db) scheduledTestService := service.ProvideScheduledTestService(scheduledTestPlanRepository, scheduledTestResultRepository) scheduledTestHandler := admin.NewScheduledTestHandler(scheduledTestService) - adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, adminAPIKeyHandler, scheduledTestHandler) + adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, gdriveOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, adminAPIKeyHandler, scheduledTestHandler) usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig) userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) @@ -216,13 +220,18 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { soraSDKClient := service.ProvideSoraSDKClient(configConfig, httpUpstream, openAITokenProvider, accountRepository, soraAccountRepository) soraMediaStorage := service.ProvideSoraMediaStorage(configConfig) soraGatewayService := service.NewSoraGatewayService(soraSDKClient, rateLimitService, httpUpstream, configConfig) - soraClientHandler := handler.NewSoraClientHandler(soraGenerationService, soraQuotaService, soraS3Storage, soraGatewayService, gatewayService, soraMediaStorage, apiKeyService) + soraClientHandler := handler.NewSoraClientHandler(soraGenerationService, soraQuotaService, soraStorageRouter, soraGatewayService, gatewayService, soraMediaStorage, apiKeyService) + soraTaskRepository := repository.NewSoraTaskRepository(db) + soraTaskService := service.NewSoraTaskService(soraTaskRepository, accountRepository, soraSDKClient, httpUpstream) + soraTaskWorker := service.NewSoraTaskWorker(soraTaskService, accountRepository, soraStorageRouter, soraMediaStorage, 60*time.Second) + soraTaskWorker.Start() + soraVideosHandler := handler.NewSoraVideosHandler(soraTaskService, gatewayService, soraStorageRouter, soraMediaStorage, soraGatewayService) soraGatewayHandler := handler.NewSoraGatewayHandler(gatewayService, soraGatewayService, concurrencyService, billingCacheService, usageRecordWorkerPool, configConfig) handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo) totpHandler := handler.NewTotpHandler(totpService) idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig) idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig) - handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService) + handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, soraVideosHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService) jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService) adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService) apiKeyAuthMiddleware := middleware.NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, configConfig) @@ -238,7 +247,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { accountExpiryService := service.ProvideAccountExpiryService(accountRepository) subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository) scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, soraTaskWorker) application := &Application{ Server: httpServer, Cleanup: v, @@ -292,6 +301,7 @@ func provideCleanup( openAIGateway *service.OpenAIGatewayService, scheduledTestRunner *service.ScheduledTestRunnerService, backupSvc *service.BackupService, + soraTaskWorker *service.SoraTaskWorker, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -433,6 +443,12 @@ func provideCleanup( } return nil }}, + {"SoraTaskWorker", func() error { + if soraTaskWorker != nil { + soraTaskWorker.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 9d2a54b98b..51396c83e4 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -76,6 +76,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // openAIGateway nil, // scheduledTestRunner nil, // backupSvc + nil, // soraTaskWorker ) require.NotPanics(t, func() { diff --git a/backend/ent/group.go b/backend/ent/group.go index 3db54a643e..7ed4990538 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -62,26 +62,28 @@ type Group struct { SoraVideoPricePerRequestHd *float64 `json:"sora_video_price_per_request_hd,omitempty"` // SoraStorageQuotaBytes holds the value of the "sora_storage_quota_bytes" field. SoraStorageQuotaBytes int64 `json:"sora_storage_quota_bytes,omitempty"` - // 是否仅允许 Claude Code 客户端 + // allow Claude Code client only ClaudeCodeOnly bool `json:"claude_code_only,omitempty"` - // 非 Claude Code 请求降级使用的分组 ID + // fallback group for non-Claude-Code requests FallbackGroupID *int64 `json:"fallback_group_id,omitempty"` - // 无效请求兜底使用的分组 ID + // fallback group for invalid request FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request,omitempty"` - // 模型路由配置:模型模式 -> 优先账号ID列表 + // model routing config: pattern -> account ids ModelRouting map[string][]int64 `json:"model_routing,omitempty"` - // 是否启用模型路由配置 + // whether model routing is enabled ModelRoutingEnabled bool `json:"model_routing_enabled,omitempty"` - // 是否注入 MCP XML 调用协议提示词(仅 antigravity 平台) + // whether MCP XML prompt injection is enabled McpXMLInject bool `json:"mcp_xml_inject,omitempty"` - // 支持的模型系列:claude, gemini_text, gemini_image + // supported model scopes: claude, gemini_text, gemini_image SupportedModelScopes []string `json:"supported_model_scopes,omitempty"` - // 分组显示排序,数值越小越靠前 + // group display order, lower comes first SortOrder int `json:"sort_order,omitempty"` // 是否允许 /v1/messages 调度到此 OpenAI 分组 AllowMessagesDispatch bool `json:"allow_messages_dispatch,omitempty"` // 默认映射模型 ID,当账号级映射找不到时使用此值 DefaultMappedModel string `json:"default_mapped_model,omitempty"` + // simulate claude usage as claude-max style (1h cache write) + SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled,omitempty"` // Edges holds the relations/edges for other nodes in the graph. // The values are being populated by the GroupQuery when eager-loading is set. Edges GroupEdges `json:"edges"` @@ -190,7 +192,7 @@ func (*Group) scanValues(columns []string) ([]any, error) { switch columns[i] { case group.FieldModelRouting, group.FieldSupportedModelScopes: values[i] = new([]byte) - case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch: + case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldSimulateClaudeMaxEnabled: values[i] = new(sql.NullBool) case group.FieldRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldSoraImagePrice360, group.FieldSoraImagePrice540, group.FieldSoraVideoPricePerRequest, group.FieldSoraVideoPricePerRequestHd: values[i] = new(sql.NullFloat64) @@ -431,6 +433,12 @@ func (_m *Group) assignValues(columns []string, values []any) error { } else if value.Valid { _m.DefaultMappedModel = value.String } + case group.FieldSimulateClaudeMaxEnabled: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field simulate_claude_max_enabled", values[i]) + } else if value.Valid { + _m.SimulateClaudeMaxEnabled = value.Bool + } default: _m.selectValues.Set(columns[i], values[i]) } @@ -630,6 +638,9 @@ func (_m *Group) String() string { builder.WriteString(", ") builder.WriteString("default_mapped_model=") builder.WriteString(_m.DefaultMappedModel) + builder.WriteString(", ") + builder.WriteString("simulate_claude_max_enabled=") + builder.WriteString(fmt.Sprintf("%v", _m.SimulateClaudeMaxEnabled)) builder.WriteByte(')') return builder.String() } diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 2612b6cff2..970c7a85c3 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -79,6 +79,8 @@ const ( FieldAllowMessagesDispatch = "allow_messages_dispatch" // FieldDefaultMappedModel holds the string denoting the default_mapped_model field in the database. FieldDefaultMappedModel = "default_mapped_model" + // FieldSimulateClaudeMaxEnabled holds the string denoting the simulate_claude_max_enabled field in the database. + FieldSimulateClaudeMaxEnabled = "simulate_claude_max_enabled" // EdgeAPIKeys holds the string denoting the api_keys edge name in mutations. EdgeAPIKeys = "api_keys" // EdgeRedeemCodes holds the string denoting the redeem_codes edge name in mutations. @@ -186,6 +188,7 @@ var Columns = []string{ FieldSortOrder, FieldAllowMessagesDispatch, FieldDefaultMappedModel, + FieldSimulateClaudeMaxEnabled, } var ( @@ -259,6 +262,8 @@ var ( DefaultDefaultMappedModel string // DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. DefaultMappedModelValidator func(string) error + // DefaultSimulateClaudeMaxEnabled holds the default value on creation for the "simulate_claude_max_enabled" field. + DefaultSimulateClaudeMaxEnabled bool ) // OrderOption defines the ordering options for the Group queries. @@ -419,6 +424,11 @@ func ByDefaultMappedModel(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldDefaultMappedModel, opts...).ToFunc() } +// BySimulateClaudeMaxEnabled orders the results by the simulate_claude_max_enabled field. +func BySimulateClaudeMaxEnabled(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldSimulateClaudeMaxEnabled, opts...).ToFunc() +} + // ByAPIKeysCount orders the results by api_keys count. func ByAPIKeysCount(opts ...sql.OrderTermOption) OrderOption { return func(s *sql.Selector) { diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 5dd8759e5d..62c91d5af0 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -205,6 +205,11 @@ func DefaultMappedModel(v string) predicate.Group { return predicate.Group(sql.FieldEQ(FieldDefaultMappedModel, v)) } +// SimulateClaudeMaxEnabled applies equality check predicate on the "simulate_claude_max_enabled" field. It's identical to SimulateClaudeMaxEnabledEQ. +func SimulateClaudeMaxEnabled(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldSimulateClaudeMaxEnabled, v)) +} + // CreatedAtEQ applies the EQ predicate on the "created_at" field. func CreatedAtEQ(v time.Time) predicate.Group { return predicate.Group(sql.FieldEQ(FieldCreatedAt, v)) @@ -1555,6 +1560,16 @@ func DefaultMappedModelContainsFold(v string) predicate.Group { return predicate.Group(sql.FieldContainsFold(FieldDefaultMappedModel, v)) } +// SimulateClaudeMaxEnabledEQ applies the EQ predicate on the "simulate_claude_max_enabled" field. +func SimulateClaudeMaxEnabledEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldSimulateClaudeMaxEnabled, v)) +} + +// SimulateClaudeMaxEnabledNEQ applies the NEQ predicate on the "simulate_claude_max_enabled" field. +func SimulateClaudeMaxEnabledNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldSimulateClaudeMaxEnabled, v)) +} + // HasAPIKeys applies the HasEdge predicate on the "api_keys" edge. func HasAPIKeys() predicate.Group { return predicate.Group(func(s *sql.Selector) { diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index 6db5b97452..9418b02f7a 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -452,6 +452,20 @@ func (_c *GroupCreate) SetNillableDefaultMappedModel(v *string) *GroupCreate { return _c } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (_c *GroupCreate) SetSimulateClaudeMaxEnabled(v bool) *GroupCreate { + _c.mutation.SetSimulateClaudeMaxEnabled(v) + return _c +} + +// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil. +func (_c *GroupCreate) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupCreate { + if v != nil { + _c.SetSimulateClaudeMaxEnabled(*v) + } + return _c +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_c *GroupCreate) AddAPIKeyIDs(ids ...int64) *GroupCreate { _c.mutation.AddAPIKeyIDs(ids...) @@ -649,6 +663,10 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultDefaultMappedModel _c.mutation.SetDefaultMappedModel(v) } + if _, ok := _c.mutation.SimulateClaudeMaxEnabled(); !ok { + v := group.DefaultSimulateClaudeMaxEnabled + _c.mutation.SetSimulateClaudeMaxEnabled(v) + } return nil } @@ -730,6 +748,9 @@ func (_c *GroupCreate) check() error { return &ValidationError{Name: "default_mapped_model", err: fmt.Errorf(`ent: validator failed for field "Group.default_mapped_model": %w`, err)} } } + if _, ok := _c.mutation.SimulateClaudeMaxEnabled(); !ok { + return &ValidationError{Name: "simulate_claude_max_enabled", err: errors.New(`ent: missing required field "Group.simulate_claude_max_enabled"`)} + } return nil } @@ -885,6 +906,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) _node.DefaultMappedModel = value } + if value, ok := _c.mutation.SimulateClaudeMaxEnabled(); ok { + _spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value) + _node.SimulateClaudeMaxEnabled = value + } if nodes := _c.mutation.APIKeysIDs(); len(nodes) > 0 { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -1599,6 +1624,18 @@ func (u *GroupUpsert) UpdateDefaultMappedModel() *GroupUpsert { return u } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (u *GroupUpsert) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsert { + u.Set(group.FieldSimulateClaudeMaxEnabled, v) + return u +} + +// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create. +func (u *GroupUpsert) UpdateSimulateClaudeMaxEnabled() *GroupUpsert { + u.SetExcluded(group.FieldSimulateClaudeMaxEnabled) + return u +} + // UpdateNewValues updates the mutable fields using the new values that were set on create. // Using this option is equivalent to using: // @@ -2295,6 +2332,20 @@ func (u *GroupUpsertOne) UpdateDefaultMappedModel() *GroupUpsertOne { }) } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (u *GroupUpsertOne) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetSimulateClaudeMaxEnabled(v) + }) +} + +// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateSimulateClaudeMaxEnabled() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateSimulateClaudeMaxEnabled() + }) +} + // Exec executes the query. func (u *GroupUpsertOne) Exec(ctx context.Context) error { if len(u.create.conflict) == 0 { @@ -3157,6 +3208,20 @@ func (u *GroupUpsertBulk) UpdateDefaultMappedModel() *GroupUpsertBulk { }) } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (u *GroupUpsertBulk) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetSimulateClaudeMaxEnabled(v) + }) +} + +// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateSimulateClaudeMaxEnabled() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateSimulateClaudeMaxEnabled() + }) +} + // Exec executes the query. func (u *GroupUpsertBulk) Exec(ctx context.Context) error { if u.create.err != nil { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index b3698596f6..75955f7df5 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -653,6 +653,20 @@ func (_u *GroupUpdate) SetNillableDefaultMappedModel(v *string) *GroupUpdate { return _u } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (_u *GroupUpdate) SetSimulateClaudeMaxEnabled(v bool) *GroupUpdate { + _u.mutation.SetSimulateClaudeMaxEnabled(v) + return _u +} + +// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupUpdate { + if v != nil { + _u.SetSimulateClaudeMaxEnabled(*v) + } + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdate) AddAPIKeyIDs(ids ...int64) *GroupUpdate { _u.mutation.AddAPIKeyIDs(ids...) @@ -1149,6 +1163,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } + if value, ok := _u.mutation.SimulateClaudeMaxEnabled(); ok { + _spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -2081,6 +2098,20 @@ func (_u *GroupUpdateOne) SetNillableDefaultMappedModel(v *string) *GroupUpdateO return _u } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (_u *GroupUpdateOne) SetSimulateClaudeMaxEnabled(v bool) *GroupUpdateOne { + _u.mutation.SetSimulateClaudeMaxEnabled(v) + return _u +} + +// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetSimulateClaudeMaxEnabled(*v) + } + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdateOne) AddAPIKeyIDs(ids ...int64) *GroupUpdateOne { _u.mutation.AddAPIKeyIDs(ids...) @@ -2607,6 +2638,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } + if value, ok := _u.mutation.SimulateClaudeMaxEnabled(); ok { + _spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index acdd0d18b2..c6604fe1c8 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -410,6 +410,7 @@ var ( {Name: "sort_order", Type: field.TypeInt, Default: 0}, {Name: "allow_messages_dispatch", Type: field.TypeBool, Default: false}, {Name: "default_mapped_model", Type: field.TypeString, Size: 100, Default: ""}, + {Name: "simulate_claude_max_enabled", Type: field.TypeBool, Default: false}, } // GroupsTable holds the schema information for the "groups" table. GroupsTable = &schema.Table{ diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index ff58fa9eb2..2f43299ffb 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -8252,6 +8252,7 @@ type GroupMutation struct { addsort_order *int allow_messages_dispatch *bool default_mapped_model *string + simulate_claude_max_enabled *bool clearedFields map[string]struct{} api_keys map[int64]struct{} removedapi_keys map[int64]struct{} @@ -10068,6 +10069,42 @@ func (m *GroupMutation) ResetDefaultMappedModel() { m.default_mapped_model = nil } +// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field. +func (m *GroupMutation) SetSimulateClaudeMaxEnabled(b bool) { + m.simulate_claude_max_enabled = &b +} + +// SimulateClaudeMaxEnabled returns the value of the "simulate_claude_max_enabled" field in the mutation. +func (m *GroupMutation) SimulateClaudeMaxEnabled() (r bool, exists bool) { + v := m.simulate_claude_max_enabled + if v == nil { + return + } + return *v, true +} + +// OldSimulateClaudeMaxEnabled returns the old "simulate_claude_max_enabled" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldSimulateClaudeMaxEnabled(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldSimulateClaudeMaxEnabled is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldSimulateClaudeMaxEnabled requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldSimulateClaudeMaxEnabled: %w", err) + } + return oldValue.SimulateClaudeMaxEnabled, nil +} + +// ResetSimulateClaudeMaxEnabled resets all changes to the "simulate_claude_max_enabled" field. +func (m *GroupMutation) ResetSimulateClaudeMaxEnabled() { + m.simulate_claude_max_enabled = nil +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by ids. func (m *GroupMutation) AddAPIKeyIDs(ids ...int64) { if m.api_keys == nil { @@ -10426,7 +10463,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 32) + fields := make([]string, 0, 33) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -10523,6 +10560,9 @@ func (m *GroupMutation) Fields() []string { if m.default_mapped_model != nil { fields = append(fields, group.FieldDefaultMappedModel) } + if m.simulate_claude_max_enabled != nil { + fields = append(fields, group.FieldSimulateClaudeMaxEnabled) + } return fields } @@ -10595,6 +10635,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.AllowMessagesDispatch() case group.FieldDefaultMappedModel: return m.DefaultMappedModel() + case group.FieldSimulateClaudeMaxEnabled: + return m.SimulateClaudeMaxEnabled() } return nil, false } @@ -10668,6 +10710,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldAllowMessagesDispatch(ctx) case group.FieldDefaultMappedModel: return m.OldDefaultMappedModel(ctx) + case group.FieldSimulateClaudeMaxEnabled: + return m.OldSimulateClaudeMaxEnabled(ctx) } return nil, fmt.Errorf("unknown Group field %s", name) } @@ -10901,6 +10945,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetDefaultMappedModel(v) return nil + case group.FieldSimulateClaudeMaxEnabled: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetSimulateClaudeMaxEnabled(v) + return nil } return fmt.Errorf("unknown Group field %s", name) } @@ -11334,6 +11385,9 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldDefaultMappedModel: m.ResetDefaultMappedModel() return nil + case group.FieldSimulateClaudeMaxEnabled: + m.ResetSimulateClaudeMaxEnabled() + return nil } return fmt.Errorf("unknown Group field %s", name) } diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 2401e5538b..aa47756f1a 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -463,6 +463,10 @@ func init() { group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) + // groupDescSimulateClaudeMaxEnabled is the schema descriptor for simulate_claude_max_enabled field. + groupDescSimulateClaudeMaxEnabled := groupFields[29].Descriptor() + // group.DefaultSimulateClaudeMaxEnabled holds the default value on creation for the simulate_claude_max_enabled field. + group.DefaultSimulateClaudeMaxEnabled = groupDescSimulateClaudeMaxEnabled.Default.(bool) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() idempotencyrecordMixinFields0 := idempotencyrecordMixin[0].Fields() _ = idempotencyrecordMixinFields0 diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 0f5a7b14ba..0842a0f846 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -33,8 +33,6 @@ func (Group) Mixin() []ent.Mixin { func (Group) Fields() []ent.Field { return []ent.Field{ - // 唯一约束通过部分索引实现(WHERE deleted_at IS NULL),支持软删除后重用 - // 见迁移文件 016_soft_delete_partial_unique_indexes.sql field.String("name"). MaxLen(100). NotEmpty(), @@ -51,7 +49,6 @@ func (Group) Fields() []ent.Field { MaxLen(20). Default(domain.StatusActive), - // Subscription-related fields (added by migration 003) field.String("platform"). MaxLen(50). Default(domain.PlatformAnthropic), @@ -73,7 +70,6 @@ func (Group) Fields() []ent.Field { field.Int("default_validity_days"). Default(30), - // 图片生成计费配置(antigravity 和 gemini 平台使用) field.Float("image_price_1k"). Optional(). Nillable(). @@ -87,7 +83,6 @@ func (Group) Fields() []ent.Field { Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}), - // Sora 按次计费配置(阶段 1) field.Float("sora_image_price_360"). Optional(). Nillable(). @@ -109,45 +104,38 @@ func (Group) Fields() []ent.Field { field.Int64("sora_storage_quota_bytes"). Default(0), - // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). Default(false). - Comment("是否仅允许 Claude Code 客户端"), + Comment("allow Claude Code client only"), field.Int64("fallback_group_id"). Optional(). Nillable(). - Comment("非 Claude Code 请求降级使用的分组 ID"), + Comment("fallback group for non-Claude-Code requests"), field.Int64("fallback_group_id_on_invalid_request"). Optional(). Nillable(). - Comment("无效请求兜底使用的分组 ID"), + Comment("fallback group for invalid request"), - // 模型路由配置 (added by migration 040) field.JSON("model_routing", map[string][]int64{}). Optional(). SchemaType(map[string]string{dialect.Postgres: "jsonb"}). - Comment("模型路由配置:模型模式 -> 优先账号ID列表"), - - // 模型路由开关 (added by migration 041) + Comment("model routing config: pattern -> account ids"), field.Bool("model_routing_enabled"). Default(false). - Comment("是否启用模型路由配置"), + Comment("whether model routing is enabled"), - // MCP XML 协议注入开关 (added by migration 042) field.Bool("mcp_xml_inject"). Default(true). - Comment("是否注入 MCP XML 调用协议提示词(仅 antigravity 平台)"), + Comment("whether MCP XML prompt injection is enabled"), - // 支持的模型系列 (added by migration 046) field.JSON("supported_model_scopes", []string{}). Default([]string{"claude", "gemini_text", "gemini_image"}). SchemaType(map[string]string{dialect.Postgres: "jsonb"}). - Comment("支持的模型系列:claude, gemini_text, gemini_image"), + Comment("supported model scopes: claude, gemini_text, gemini_image"), - // 分组排序 (added by migration 052) field.Int("sort_order"). Default(0). - Comment("分组显示排序,数值越小越靠前"), + Comment("group display order, lower comes first"), // OpenAI Messages 调度配置 (added by migration 069) field.Bool("allow_messages_dispatch"). @@ -157,6 +145,9 @@ func (Group) Fields() []ent.Field { MaxLen(100). Default(""). Comment("默认映射模型 ID,当账号级映射找不到时使用此值"), + field.Bool("simulate_claude_max_enabled"). + Default(false). + Comment("simulate claude usage as claude-max style (1h cache write)"), } } @@ -172,14 +163,11 @@ func (Group) Edges() []ent.Edge { edge.From("allowed_users", User.Type). Ref("allowed_groups"). Through("user_allowed_groups", UserAllowedGroup.Type), - // 注意:fallback_group_id 直接作为字段使用,不定义 edge - // 这样允许多个分组指向同一个降级分组(M2O 关系) } } func (Group) Indexes() []ent.Index { return []ent.Index{ - // name 字段已在 Fields() 中声明 Unique(),无需重复索引 index.Fields("status"), index.Fields("platform"), index.Fields("subscription_type"), diff --git a/backend/go.mod b/backend/go.mod index 135cbd3eaf..0bef703da0 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -22,6 +22,8 @@ require ( github.com/imroc/req/v3 v3.57.0 github.com/lib/pq v1.10.9 github.com/patrickmn/go-cache v2.1.0+incompatible + github.com/pkoukk/tiktoken-go v0.1.8 + github.com/pkoukk/tiktoken-go-loader v0.0.2 github.com/pquerna/otp v1.5.0 github.com/redis/go-redis/v9 v9.17.2 github.com/refraction-networking/utls v1.8.2 @@ -37,8 +39,10 @@ require ( go.uber.org/zap v1.24.0 golang.org/x/crypto v0.48.0 golang.org/x/net v0.49.0 + golang.org/x/oauth2 v0.30.0 golang.org/x/sync v0.19.0 golang.org/x/term v0.40.0 + google.golang.org/api v0.153.0 gopkg.in/natefinch/lumberjack.v2 v2.2.1 gopkg.in/yaml.v3 v3.0.1 modernc.org/sqlite v1.44.3 @@ -46,6 +50,7 @@ require ( require ( ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9 // indirect + cloud.google.com/go/compute/metadata v0.7.0 // indirect dario.cat/mergo v1.0.2 // indirect github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 // indirect github.com/Microsoft/go-winio v0.6.2 // indirect @@ -87,6 +92,7 @@ require ( github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/distribution/reference v0.6.0 // indirect + github.com/dlclark/regexp2 v1.10.0 // indirect github.com/docker/docker v28.5.1+incompatible // indirect github.com/docker/go-connections v0.6.0 // indirect github.com/docker/go-units v0.5.0 // indirect @@ -105,8 +111,13 @@ require ( github.com/go-playground/universal-translator v0.18.1 // indirect github.com/go-playground/validator/v10 v10.14.0 // indirect github.com/goccy/go-json v0.10.2 // indirect + github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da // indirect + github.com/golang/protobuf v1.5.4 // indirect github.com/google/go-cmp v0.7.0 // indirect github.com/google/go-querystring v1.1.0 // indirect + github.com/google/s2a-go v0.1.7 // indirect + github.com/googleapis/enterprise-certificate-proxy v0.3.2 // indirect + github.com/googleapis/gax-go/v2 v2.12.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3 // indirect github.com/hashicorp/hcl v1.0.0 // indirect github.com/hashicorp/hcl/v2 v2.18.1 // indirect @@ -162,11 +173,11 @@ require ( github.com/yusufpapurcu/wmi v1.2.4 // indirect github.com/zclconf/go-cty v1.14.4 // indirect github.com/zclconf/go-cty-yaml v1.1.0 // indirect + go.opencensus.io v0.24.0 // indirect go.opentelemetry.io/auto/sdk v1.1.0 // indirect go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect go.opentelemetry.io/otel v1.37.0 // indirect go.opentelemetry.io/otel/metric v1.37.0 // indirect - go.opentelemetry.io/otel/sdk v1.37.0 // indirect go.opentelemetry.io/otel/trace v1.37.0 // indirect go.uber.org/atomic v1.10.0 // indirect go.uber.org/automaxprocs v1.6.0 // indirect @@ -176,6 +187,8 @@ require ( golang.org/x/mod v0.32.0 // indirect golang.org/x/sys v0.41.0 // indirect golang.org/x/text v0.34.0 // indirect + google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 // indirect google.golang.org/grpc v1.75.1 // indirect google.golang.org/protobuf v1.36.10 // indirect gopkg.in/ini.v1 v1.67.0 // indirect diff --git a/backend/go.sum b/backend/go.sum index 270be5f8fd..ace053ba71 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -1,5 +1,8 @@ ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9 h1:E0wvcUXTkgyN4wy4LGtNzMNGMytJN8afmIWXJVMi4cc= ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9/go.mod h1:Oe1xWPuu5q9LzyrWfbZmEZxFYeu4BHTyzfjeW2aZp/w= +cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= +cloud.google.com/go/compute/metadata v0.7.0 h1:PBWF+iiAerVNe8UCHxdOt6eHLVc3ydFeOCw78U8ytSU= +cloud.google.com/go/compute/metadata v0.7.0/go.mod h1:j5MvL9PprKL39t166CoB1uVHfQMs4tFQZZcKwksXUjo= dario.cat/mergo v1.0.2 h1:85+piFYR1tMbRrLcDwR18y4UKJ3aH1Tbzi24VRW1TK8= dario.cat/mergo v1.0.2/go.mod h1:E/hbnu0NxMFBjpMIE34DRGLWqDy0g5FuKDhCb31ngxA= entgo.io/ent v0.14.5 h1:Rj2WOYJtCkWyFo6a+5wB3EfBRP0rnx1fMk6gGA0UUe4= @@ -8,6 +11,7 @@ github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6 h1:He8af github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h1:8o94RPi1/7XTJvwPpRSzSUedZrtlirdB3r9Z20bi2f8= github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8= github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E= +github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU= github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU= github.com/DouDOU-start/go-sora2api v1.1.0 h1:PxWiukK77StiHxEngOFwT1rKUn9oTAJJTl07wQUXwiU= @@ -89,11 +93,14 @@ github.com/bytedance/sonic v1.9.1 h1:6iJ6NqdoxCDr6mbY8h18oSO+cShGSMRGCEo7F2h0x8s github.com/bytedance/sonic v1.9.1/go.mod h1:i736AoUSYt75HyZLoJW9ERYxcy6eaN6h4BZXU064P/U= github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8= github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE= +github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU= 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/chenzhuoyu/base64x v0.0.0-20211019084208-fb5309c8db06/go.mod h1:DH46F32mSOjUmXrMHnKwZdA8wcEefY7UVqBKYGjpdQY= github.com/chenzhuoyu/base64x v0.0.0-20221115062448-fe3a3abad311 h1:qSGYFH7+jGhDF8vLC+iwCD4WpbV1EBDSzWkJODFLams= github.com/chenzhuoyu/base64x v0.0.0-20221115062448-fe3a3abad311/go.mod h1:b583jCggY9gE99b6G5LEC39OIiVsWj+R97kbl5odCEk= +github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw= +github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc= github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g= github.com/coder/websocket v1.8.14/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg= github.com/containerd/errdefs v1.0.0 h1:tg5yIfIlQIrxYtu9ajqY42W3lpS19XqdxRQeEwYG8PI= @@ -120,6 +127,8 @@ github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/r github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk= github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E= +github.com/dlclark/regexp2 v1.10.0 h1:+/GIL799phkJqYW+3YbOd8LCcbHzT0Pbo8zl70MHsq0= +github.com/dlclark/regexp2 v1.10.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM= github.com/docker/docker v28.5.1+incompatible/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk= github.com/docker/go-connections v0.6.0 h1:LlMG9azAe1TqfR7sO+NJttz1gy6KO7VJBh+pMmjSD94= @@ -130,6 +139,10 @@ github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkp github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/ebitengine/purego v0.8.4 h1:CF7LEKg5FFOsASUj0+QwaXf8Ht6TlFxg09+S9wz0omw= github.com/ebitengine/purego v0.8.4/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ= +github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= +github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= +github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= +github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM= github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU= github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2Wg= @@ -167,7 +180,29 @@ 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/golang-jwt/jwt/v5 v5.2.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8= github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= +github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= +github.com/golang/groupcache v0.0.0-20200121045136-8c9f03a8e57e/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= +github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da h1:oI5xCqsCo564l8iNU+DwB5epxmsaqB+rhGL0m5jtYqE= +github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= +github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A= +github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= +github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= +github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8= +github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA= +github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs= +github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w= +github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0= +github.com/golang/protobuf v1.4.1/go.mod h1:U8fpvMrcmy5pZrNK1lt4xCsGvpyWQ/VVv6QDs8UjoX8= +github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M= +github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= +github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= +github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= +github.com/google/go-cmp v0.5.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.5.2/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= +github.com/google/go-cmp v0.5.3/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= @@ -176,10 +211,17 @@ github.com/google/go-querystring v1.1.0/go.mod h1:Kcdr2DB4koayq7X8pmAG4sNG59So17 github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/s2a-go v0.1.7 h1:60BLSyTrOV4/haCDW4zb1guZItoSq8foHCXrAnjBo/o= +github.com/google/s2a-go v0.1.7/go.mod h1:50CgR4k1jNlWBu4UfS4AcfhVe1r6pdZPygJ3R8F0Qdw= +github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4= github.com/google/wire v0.7.0/go.mod h1:n6YbUQD9cPKTnHXEBN2DXlOp/mVADhVErcMFb0v3J18= +github.com/googleapis/enterprise-certificate-proxy v0.3.2 h1:Vie5ybvEvT75RniqhfFxPRy3Bf7vr3h0cechB90XaQs= +github.com/googleapis/enterprise-certificate-proxy v0.3.2/go.mod h1:VLSiSSBs/ksPL8kq3OBOQ6WRI2QnaFynd1DCjZ62+V0= +github.com/googleapis/gax-go/v2 v2.12.0 h1:A+gCJKdRfqXkr+BIRGtZLibNXf0m1f9E4HG56etFpas= +github.com/googleapis/gax-go/v2 v2.12.0/go.mod h1:y+aIqrI5eb1YGMVJfuV3185Ts/D7qKpsEkdD5+I6QGU= 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/grpc-ecosystem/grpc-gateway/v2 v2.27.3 h1:NmZ1PKzSTQbuGHw9DGPFomqkkLWMC+vZCkfs+FHv1Vg= @@ -273,6 +315,10 @@ github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6 github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs= github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pkoukk/tiktoken-go v0.1.8 h1:85ENo+3FpWgAACBaEUVp+lctuTcYUO7BtmfhlN/QTRo= +github.com/pkoukk/tiktoken-go v0.1.8/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg= +github.com/pkoukk/tiktoken-go-loader v0.0.2 h1:LUKws63GV3pVHwH1srkBplBv+7URgmOmhSkRxsIvsK4= +github.com/pkoukk/tiktoken-go-loader v0.0.2/go.mod h1:4mIkYyZooFlnenDlormIo6cd5wrlUKNr97wp9nGgEKo= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= @@ -282,6 +328,7 @@ github.com/pquerna/otp v1.5.0 h1:NMMR+WrmaqXU4EzdGJEE1aUUI0AMRzsp96fFFWNPwxs= github.com/pquerna/otp v1.5.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg= github.com/prashantv/gostub v1.1.0 h1:BTyx3RfQjRHnUWaGF9oQos79AlQ5k8WNktv7VGvVH4g= github.com/prashantv/gostub v1.1.0/go.mod h1:A5zLQHz7ieHGG7is6LLXLz7I8+3LZzsrV0P1IAHhP5U= +github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA= github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/quic-go/quic-go v0.57.1 h1:25KAAR9QR8KZrCZRThWMKVAwGoiHIrNbT72ULHTuI10= @@ -370,6 +417,8 @@ github.com/zclconf/go-cty-yaml v1.1.0 h1:nP+jp0qPHv2IhUVqmQSzjvqAWcObN0KBkUl2rWB github.com/zclconf/go-cty-yaml v1.1.0/go.mod h1:9YLUH4g7lOhVWqUbctnVlZ5KLpg7JAprQNgxSZ1Gyxs= github.com/zeromicro/go-zero v1.9.4 h1:aRLFoISqAYijABtkbliQC5SsI5TbizJpQvoHc9xup8k= github.com/zeromicro/go-zero v1.9.4/go.mod h1:a17JOTch25SWxBcUgJZYps60hygK3pIYdw7nGwlcS38= +go.opencensus.io v0.24.0 h1:y73uSU6J157QMP2kn2r30vwW1A2W2WFwSCGnAVxeaD0= +go.opencensus.io v0.24.0/go.mod h1:vNK8G9p7aAivkbmorf4v+7Hgx+Zs0yY+0fOtgBfjQKo= go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA= go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A= go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk= @@ -384,6 +433,8 @@ go.opentelemetry.io/otel/metric v1.37.0 h1:mvwbQS5m0tbmqML4NqK+e3aDiO02vsf/Wgbsd go.opentelemetry.io/otel/metric v1.37.0/go.mod h1:04wGrZurHYKOc+RKeye86GwKiTb9FKm1WHtO+4EVr2E= go.opentelemetry.io/otel/sdk v1.37.0 h1:ItB0QUqnjesGRvNcmAcU0LyvkVyGJ2xftD29bWdDvKI= go.opentelemetry.io/otel/sdk v1.37.0/go.mod h1:VredYzxUvuo2q3WRcDnKDjbdvmO0sCzOvVAiY+yUkAg= +go.opentelemetry.io/otel/sdk/metric v1.37.0 h1:90lI228XrB9jCMuSdA0673aubgRobVZFhbjxHHspCPc= +go.opentelemetry.io/otel/sdk/metric v1.37.0/go.mod h1:cNen4ZWfiD37l5NhS+Keb5RXVWZWpRE+9WyVCpbo5ps= go.opentelemetry.io/otel/trace v1.37.0 h1:HLdcFNbRQBE2imdSEgm/kwqmQj1Or1l/7bW6mxVK7z4= go.opentelemetry.io/otel/trace v1.37.0/go.mod h1:TlgrlQ+PtQO5XFerSPUYG0JSgGyryXewPGyayAWSBS0= go.opentelemetry.io/proto/otlp v1.3.1 h1:TrMUixzpM0yuc/znrFTP9MMRh8trP93mkCiDVeXrui0= @@ -403,18 +454,40 @@ go.uber.org/zap v1.24.0/go.mod h1:2kMP+WWQ8aoFoedH3T2sq6iJ2yDWpHbP0f6MQbS9Gkg= golang.org/x/arch v0.0.0-20210923205945-b76863e36670/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8= golang.org/x/arch v0.3.0 h1:02VY4/ZcO/gBOH6PUaoiptASxtXU10jazRCP865E97k= golang.org/x/arch v0.3.0/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts= golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos= +golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= +golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= +golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= +golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c= golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU= +golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= +golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20201110031124-69a78807bb2b/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20211104170005-ce137452f963/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= +golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= +golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI= +golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU= +golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -430,22 +503,50 @@ golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.40.0 h1:36e4zGLqU4yhjlmxEaagx2KuYbJq3EwY8K943ZsHcvg= golang.org/x/term v0.40.0/go.mod h1:w2P8uVp06p2iyKKuvXIm7N/y0UCRt3UfJTfZ7oOpglM= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE= golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY= +golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= +golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc= golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 h1:wpZ8pe2x1Q3f2KyT5f8oP/fa9rHAKgFPr/HZdNuS+PQ= +gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= +gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= +google.golang.org/api v0.153.0 h1:N1AwGhielyKFaUqH07/ZSIQR3uNPcV7NVw0vj+j4iR4= +google.golang.org/api v0.153.0/go.mod h1:3qNJX5eOmhiWYc67jRA/3GsDw97UFb5ivv7Y2PrriAY= +google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM= +google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4= +google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc= +google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc= +google.golang.org/genproto v0.0.0-20200526211855-cb27e3aa2013/go.mod h1:NbSheEEYHJ7i3ixzK3sjbqSGDJWnxyFXZblF3eUsNvo= google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4 h1:8XJ4pajGwOlasW+L13MnEGA8W4115jJySQtVfS2/IBU= google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4/go.mod h1:NnuHhy+bxcg30o7FnVAZbXsPHUDQ9qKWAQKCD7VxFtk= google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 h1:i8QOKZfYg6AbGVZzUAY3LrNWCKF8O6zFisU9Wl9RER4= google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4/go.mod h1:HSkG/KdJWusxU1F6CNrwNDjBMgisKxGnc5dAZfT0mjQ= +google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= +google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg= +google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY= +google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk= +google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc= google.golang.org/grpc v1.75.1 h1:/ODCNEuf9VghjgO3rqLcfg8fiOP0nSluljWFlDxELLI= google.golang.org/grpc v1.75.1/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ= +google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= +google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= +google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= +google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE= +google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo= +google.golang.org/protobuf v1.22.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= +google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= +google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= +google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c= google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE= google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= @@ -460,6 +561,8 @@ gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q= gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA= +honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= +honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis= modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc= diff --git a/backend/internal/handler/admin/account_affinity_handler.go b/backend/internal/handler/admin/account_affinity_handler.go new file mode 100644 index 0000000000..62f4c3fe80 --- /dev/null +++ b/backend/internal/handler/admin/account_affinity_handler.go @@ -0,0 +1,254 @@ +package admin + +import ( + "context" + "log/slog" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +// 账号亲和调度 API 与契约说明见 docs/ACCOUNT_SCHEDULING_FLOW.md 。 + +// GetAffinityClients returns the list of affinity clients for an account with last active timestamps. +// GET /api/v1/admin/accounts/:id/affinity-clients +func (h *AccountHandler) GetAffinityClients(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + + account, err := h.adminService.GetAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + + if !account.IsAffinityEnabled() { + response.Success(c, []service.AffinityClient{}) + return + } + + if h.gatewayCache == nil || len(account.GroupIDs) == 0 { + response.Success(c, []service.AffinityClient{}) + return + } + + clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores( + c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL, + ) + if err != nil { + response.Success(c, []service.AffinityClient{}) + return + } + + response.Success(c, clients) +} + +// clearAccountAffinity 清除指定账号在所有分组的亲和记录。 +func (h *AccountHandler) clearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) { + if h.gatewayCache == nil || len(groupIDs) == 0 { + return + } + if err := h.gatewayCache.ClearAccountAffinity(ctx, accountID, groupIDs); err != nil { + // 清理失败不影响主流程,记录日志即可 + slog.Warn("clear account affinity failed", + "account_id", accountID, + "error", err, + ) + } +} + +// countUniqueUsersFromAffinityMembers 从 "{userID}/{clientID}" 格式的成员列表中计算唯一用户数。 +func countUniqueUsersFromAffinityMembers(members []string) int64 { + users := make(map[string]struct{}, len(members)) + for _, m := range members { + if idx := strings.Index(m, "/"); idx > 0 { + users[m[:idx]] = struct{}{} + } + } + return int64(len(users)) +} + +// AffinityDetailClient 亲和详情中的单个客户端信息 +type AffinityDetailClient struct { + ClientID string `json:"client_id"` + LastActive time.Time `json:"last_active"` + IsPinned bool `json:"is_pinned"` +} + +// AffinityDetailUser 亲和详情中按用户分组的信息 +type AffinityDetailUser struct { + UserID int64 `json:"user_id"` + UserEmail string `json:"user_email"` + ClientCount int `json:"client_count"` + IsPinned bool `json:"is_pinned"` + Clients []AffinityDetailClient `json:"clients"` +} + +// AffinityDetailsResponse 亲和详情响应 +type AffinityDetailsResponse struct { + Users []AffinityDetailUser `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` +} + +type affinityStateSnapshot struct { + enabled bool + groupIDs []int64 +} + +func (h *AccountHandler) captureAffinityStates(ctx context.Context, accountIDs []int64) map[int64]affinityStateSnapshot { + states := make(map[int64]affinityStateSnapshot, len(accountIDs)) + if h.gatewayCache == nil || len(accountIDs) == 0 { + return states + } + for _, accountID := range accountIDs { + account, err := h.adminService.GetAccount(ctx, accountID) + if err != nil || account == nil { + continue + } + states[accountID] = affinityStateSnapshot{ + enabled: account.IsAffinityEnabled(), + groupIDs: append([]int64(nil), account.GroupIDs...), + } + } + return states +} + +func (h *AccountHandler) clearAffinityCacheIfDisabled(ctx context.Context, accountID int64, oldState affinityStateSnapshot) { + if h.gatewayCache == nil || !oldState.enabled { + return + } + account, err := h.adminService.GetAccount(ctx, accountID) + if err != nil || account == nil || account.IsAffinityEnabled() { + return + } + groupIDs := oldState.groupIDs + if len(account.GroupIDs) > 0 { + groupIDs = mergeGroupIDs(groupIDs, account.GroupIDs) + } + h.clearAccountAffinity(ctx, accountID, groupIDs) +} + +func (h *AccountHandler) clearAffinityCacheForBulkIfDisabled(ctx context.Context, accountIDs []int64, oldStates map[int64]affinityStateSnapshot) { + for _, accountID := range accountIDs { + oldState, ok := oldStates[accountID] + if !ok { + continue + } + h.clearAffinityCacheIfDisabled(ctx, accountID, oldState) + } +} + +// GetAffinityDetails returns the affinity details grouped by user for an account. +// GET /api/v1/admin/accounts/:id/affinity-details +func (h *AccountHandler) GetAffinityDetails(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + + emptyResp := AffinityDetailsResponse{ + Users: []AffinityDetailUser{}, + TotalUsers: 0, + TotalClients: 0, + PinnedUsers: []int64{}, + } + + account, err := h.adminService.GetAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + + if !account.IsAffinityEnabled() { + response.Success(c, emptyResp) + return + } + + pinnedUsers := account.GetPinnedUsers() + if pinnedUsers == nil { + pinnedUsers = []int64{} + } + + if h.gatewayCache == nil || len(account.GroupIDs) == 0 { + emptyResp.PinnedUsers = pinnedUsers + response.Success(c, emptyResp) + return + } + + clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores( + c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL, + ) + if err != nil { + emptyResp.PinnedUsers = pinnedUsers + response.Success(c, emptyResp) + return + } + + pinnedSet := make(map[int64]struct{}, len(pinnedUsers)) + for _, uid := range pinnedUsers { + pinnedSet[uid] = struct{}{} + } + + // 按 UserID 分组 + userMap := make(map[int64]*AffinityDetailUser) + var userOrder []int64 + for _, cl := range clients { + u, ok := userMap[cl.UserID] + if !ok { + _, pinned := pinnedSet[cl.UserID] + u = &AffinityDetailUser{ + UserID: cl.UserID, + UserEmail: "", + ClientCount: 0, + IsPinned: pinned, + Clients: []AffinityDetailClient{}, + } + userMap[cl.UserID] = u + userOrder = append(userOrder, cl.UserID) + } + u.Clients = append(u.Clients, AffinityDetailClient{ + ClientID: cl.ClientID, + LastActive: cl.LastActive, + IsPinned: false, // 客户端级别暂无 pinned 概念 + }) + } + + // 关联用户邮箱(查询失败时保持空字符串,不影响主流程) + for _, uid := range userOrder { + user, uErr := h.adminService.GetUser(c.Request.Context(), uid) + if uErr != nil || user == nil { + continue + } + if u, ok := userMap[uid]; ok { + if user.Email != "" { + u.UserEmail = user.Email + } else { + u.UserEmail = user.Username + } + } + } + + users := make([]AffinityDetailUser, 0, len(userOrder)) + for _, uid := range userOrder { + u := userMap[uid] + u.ClientCount = len(u.Clients) + users = append(users, *u) + } + + response.Success(c, AffinityDetailsResponse{ + Users: users, + TotalUsers: len(users), + TotalClients: len(clients), + PinnedUsers: pinnedUsers, + }) +} diff --git a/backend/internal/handler/admin/account_data_handler_test.go b/backend/internal/handler/admin/account_data_handler_test.go index 285033a17d..e516b275f2 100644 --- a/backend/internal/handler/admin/account_data_handler_test.go +++ b/backend/internal/handler/admin/account_data_handler_test.go @@ -65,6 +65,7 @@ func setupAccountDataRouter() (*gin.Engine, *stubAdminService) { nil, nil, nil, + nil, ) router.GET("/api/v1/admin/accounts/data", h.ExportData) diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 3ef213e18d..ce48977014 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -57,6 +57,7 @@ type AccountHandler struct { sessionLimitCache service.SessionLimitCache rpmCache service.RPMCache tokenCacheInvalidator service.TokenCacheInvalidator + gatewayCache service.GatewayCache } // NewAccountHandler creates a new admin account handler @@ -74,6 +75,7 @@ func NewAccountHandler( sessionLimitCache service.SessionLimitCache, rpmCache service.RPMCache, tokenCacheInvalidator service.TokenCacheInvalidator, + gatewayCache service.GatewayCache, ) *AccountHandler { return &AccountHandler{ adminService: adminService, @@ -89,6 +91,7 @@ func NewAccountHandler( sessionLimitCache: sessionLimitCache, rpmCache: rpmCache, tokenCacheInvalidator: tokenCacheInvalidator, + gatewayCache: gatewayCache, } } @@ -206,6 +209,34 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac } } + // 亲和客户端数据(启用亲和的账号始终返回 count,即使为 0) + if account.IsAffinityEnabled() { + if h.gatewayCache != nil && len(account.GroupIDs) > 0 { + accountGroups := map[int64][]int64{account.ID: account.GroupIDs} + if clients, err := h.gatewayCache.GetAccountAffinityClientsBatch(ctx, accountGroups, service.ClientAffinityTTL); err == nil { + if cl, ok := clients[account.ID]; ok && len(cl) > 0 { + count := int64(len(cl)) + item.AffinityClientCount = &count + item.AffinityClients = cl + userCount := countUniqueUsersFromAffinityMembers(cl) + item.AffinityUserCount = &userCount + } else { + zero := int64(0) + item.AffinityClientCount = &zero + item.AffinityUserCount = &zero + } + } else { + zero := int64(0) + item.AffinityClientCount = &zero + item.AffinityUserCount = &zero + } + } else { + zero := int64(0) + item.AffinityClientCount = &zero + item.AffinityUserCount = &zero + } + } + return item } @@ -318,6 +349,21 @@ func (h *AccountHandler) List(c *gin.Context) { _ = g.Wait() } + // 获取亲和客户端数据(Redis Pipeline,低开销) + var affinityClients map[int64][]string + if h.gatewayCache != nil { + accountGroups := make(map[int64][]int64) + for i := range accounts { + acc := &accounts[i] + if acc.IsAffinityEnabled() && len(acc.GroupIDs) > 0 { + accountGroups[acc.ID] = acc.GroupIDs + } + } + if len(accountGroups) > 0 { + affinityClients, _ = h.gatewayCache.GetAccountAffinityClientsBatch(c.Request.Context(), accountGroups, service.ClientAffinityTTL) + } + } + // Build response with concurrency info result := make([]AccountWithConcurrency, len(accounts)) for i := range accounts { @@ -348,6 +394,22 @@ func (h *AccountHandler) List(c *gin.Context) { } } + // 注入亲和客户端数据到 DTO(启用亲和的账号始终返回 count,即使为 0) + if acc.IsAffinityEnabled() { + if clients, ok := affinityClients[acc.ID]; ok && len(clients) > 0 { + count := int64(len(clients)) + item.AffinityClientCount = &count + item.AffinityClients = clients + // 从成员列表中解析唯一用户数 + userCount := countUniqueUsersFromAffinityMembers(clients) + item.AffinityUserCount = &userCount + } else { + zero := int64(0) + item.AffinityClientCount = &zero + item.AffinityUserCount = &zero + } + } + result[i] = item } @@ -568,6 +630,8 @@ func (h *AccountHandler) Update(c *gin.Context) { // base_rpm 输入校验:负值归零,超过 10000 截断 sanitizeExtraBaseRPM(req.Extra) + oldStates := h.captureAffinityStates(c.Request.Context(), []int64{accountID}) + // 确定是否跳过混合渠道检查 skipCheck := req.ConfirmMixedChannelRisk != nil && *req.ConfirmMixedChannelRisk @@ -604,6 +668,8 @@ func (h *AccountHandler) Update(c *gin.Context) { return } + h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), []int64{accountID}, oldStates) + response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } @@ -1320,6 +1386,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) { return } + oldStates := h.captureAffinityStates(c.Request.Context(), req.AccountIDs) + result, err := h.adminService.BulkUpdateAccounts(c.Request.Context(), &service.BulkUpdateAccountsInput{ AccountIDs: req.AccountIDs, Name: req.Name, @@ -1341,6 +1409,12 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) { c.JSON(409, gin.H{ "error": "mixed_channel_warning", "message": mixedErr.Error(), + "details": gin.H{ + "group_id": mixedErr.GroupID, + "group_name": mixedErr.GroupName, + "current_platform": mixedErr.CurrentPlatform, + "other_platform": mixedErr.OtherPlatform, + }, }) return } @@ -1348,6 +1422,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) { return } + h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), req.AccountIDs, oldStates) + response.Success(c, result) } @@ -1560,6 +1636,25 @@ func (h *AccountHandler) ResetQuota(c *gin.Context) { response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } +// mergeGroupIDs 合并两个 groupID 切片并去重。 +func mergeGroupIDs(a, b []int64) []int64 { + seen := make(map[int64]struct{}, len(a)+len(b)) + result := make([]int64, 0, len(a)+len(b)) + for _, id := range a { + if _, ok := seen[id]; !ok { + seen[id] = struct{}{} + result = append(result, id) + } + } + for _, id := range b { + if _, ok := seen[id]; !ok { + seen[id] = struct{}{} + result = append(result, id) + } + } + return result +} + // GetTempUnschedulable handles getting temporary unschedulable status // GET /api/v1/admin/accounts/:id/temp-unschedulable func (h *AccountHandler) GetTempUnschedulable(c *gin.Context) { diff --git a/backend/internal/handler/admin/account_handler_affinity_cache_test.go b/backend/internal/handler/admin/account_handler_affinity_cache_test.go new file mode 100644 index 0000000000..3c22bb543a --- /dev/null +++ b/backend/internal/handler/admin/account_handler_affinity_cache_test.go @@ -0,0 +1,298 @@ +package admin + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type affinityCacheAdminService struct { + *stubAdminService + accounts map[int64]*service.Account + bulkResultIDs []int64 + bulkFailedCount int +} + +func (s *affinityCacheAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) { + if acc, ok := s.accounts[id]; ok { + cloned := *acc + return &cloned, nil + } + return s.stubAdminService.GetAccount(context.Background(), id) +} + +func (s *affinityCacheAdminService) UpdateAccount(_ context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) { + acc := s.accounts[id] + if acc == nil { + return s.stubAdminService.UpdateAccount(context.Background(), id, input) + } + if input.Extra != nil { + acc.Extra = input.Extra + } + if input.GroupIDs != nil { + acc.GroupIDs = *input.GroupIDs + } + cloned := *acc + return &cloned, nil +} + +func (s *affinityCacheAdminService) BulkUpdateAccounts(_ context.Context, input *service.BulkUpdateAccountsInput) (*service.BulkUpdateAccountsResult, error) { + for _, accountID := range input.AccountIDs { + acc := s.accounts[accountID] + if acc == nil { + continue + } + if input.Extra != nil { + acc.Extra = input.Extra + } + if input.GroupIDs != nil { + acc.GroupIDs = *input.GroupIDs + } + } + successIDs := append([]int64(nil), input.AccountIDs...) + failed := 0 + if len(s.bulkResultIDs) > 0 { + successIDs = append([]int64(nil), s.bulkResultIDs...) + failed = s.bulkFailedCount + } + return &service.BulkUpdateAccountsResult{ + Success: len(successIDs), + Failed: failed, + SuccessIDs: successIDs, + Results: []service.BulkUpdateAccountResult{}, + }, nil +} + +type affinityCacheClearRecorder struct { + clearCalls []struct { + accountID int64 + groupIDs []int64 + } +} + +func (r *affinityCacheClearRecorder) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) { + return 0, nil +} +func (r *affinityCacheClearRecorder) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { + return nil +} +func (r *affinityCacheClearRecorder) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (r *affinityCacheClearRecorder) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { + return nil, nil +} +func (r *affinityCacheClearRecorder) ClearAccountAffinity(_ context.Context, accountID int64, groupIDs []int64) error { + r.clearCalls = append(r.clearCalls, struct { + accountID int64 + groupIDs []int64 + }{ + accountID: accountID, + groupIDs: append([]int64(nil), groupIDs...), + }) + return nil +} +func (r *affinityCacheClearRecorder) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) { + return 0, 0, 0, nil +} + +func setupAffinityCacheRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache) + router.PUT("/api/v1/admin/accounts/:id", handler.Update) + router.POST("/api/v1/admin/accounts/bulk-update", handler.BulkUpdate) + return router +} + +func TestAccountHandlerUpdate_DisablingAffinityClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 3: { + ID: 3, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1, 2}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "extra": map[string]any{}, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPut, "/api/v1/admin/accounts/3", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 1) + require.Equal(t, int64(3), cache.clearCalls[0].accountID) + require.Equal(t, []int64{1, 2}, cache.clearCalls[0].groupIDs) +} + +func TestAccountHandlerBulkUpdate_DisablingAffinityClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{11}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + 2: { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeSetupToken, + Status: service.StatusActive, + GroupIDs: []int64{22}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1, 2}, + "extra": map[string]any{ + "client_affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 2) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{11}, cache.clearCalls[0].groupIDs) + require.Equal(t, int64(2), cache.clearCalls[1].accountID) + require.Equal(t, []int64{22}, cache.clearCalls[1].groupIDs) +} + +func TestAccountHandlerBulkUpdate_DisablingAffinityWithNewFlagClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{33}, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1}, + "extra": map[string]any{ + "affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 1) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{33}, cache.clearCalls[0].groupIDs) +} + +func TestAccountHandlerBulkUpdate_PartialFailureStillClearsDisabledAffinityCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{44}, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + 2: { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeSetupToken, + Status: service.StatusActive, + GroupIDs: []int64{55}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + bulkResultIDs: []int64{1}, + bulkFailedCount: 1, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1, 2}, + "group_ids": []int64{101}, + "extra": map[string]any{ + "affinity_enabled": false, + "client_affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 2) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{44, 101}, cache.clearCalls[0].groupIDs) + require.Equal(t, int64(2), cache.clearCalls[1].accountID) + require.Equal(t, []int64{55, 101}, cache.clearCalls[1].groupIDs) +} diff --git a/backend/internal/handler/admin/account_handler_affinity_details_test.go b/backend/internal/handler/admin/account_handler_affinity_details_test.go new file mode 100644 index 0000000000..0a4b52d579 --- /dev/null +++ b/backend/internal/handler/admin/account_handler_affinity_details_test.go @@ -0,0 +1,205 @@ +package admin + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type affinityDetailsAdminService struct { + *stubAdminService + account service.Account + usersByID map[int64]service.User +} + +func (s *affinityDetailsAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) { + if s.account.ID == id { + acc := s.account + return &acc, nil + } + return s.stubAdminService.GetAccount(context.Background(), id) +} + +func (s *affinityDetailsAdminService) GetUser(_ context.Context, id int64) (*service.User, error) { + if u, ok := s.usersByID[id]; ok { + user := u + return &user, nil + } + return s.stubAdminService.GetUser(context.Background(), id) +} + +type affinityDetailsGatewayCacheStub struct { + clients []service.AffinityClient +} + +func (s *affinityDetailsGatewayCacheStub) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) { + return 0, nil +} +func (s *affinityDetailsGatewayCacheStub) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (s *affinityDetailsGatewayCacheStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { + return s.clients, nil +} +func (s *affinityDetailsGatewayCacheStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) { + return 0, 0, 0, nil +} + +func setupAffinityDetailsRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache) + router.GET("/api/v1/admin/accounts/:id/affinity-details", handler.GetAffinityDetails) + return router +} + +func TestAccountHandlerGetAffinityDetails_ContractFields(t *testing.T) { + adminSvc := &affinityDetailsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 88, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1}, + Extra: map[string]any{ + "client_affinity_enabled": true, + "pinned_users": []any{float64(42), float64(99)}, + }, + }, + usersByID: map[int64]service.User{ + 42: {ID: 42, Email: "user42@example.com"}, + 7: {ID: 7, Username: "user7"}, + }, + } + cache := &affinityDetailsGatewayCacheStub{ + clients: []service.AffinityClient{ + {UserID: 42, ClientID: "client-a", LastActive: time.Now().Add(-time.Minute)}, + {UserID: 7, ClientID: "client-b", LastActive: time.Now().Add(-2 * time.Minute)}, + {UserID: 42, ClientID: "client-c", LastActive: time.Now().Add(-3 * time.Minute)}, + }, + } + router := setupAffinityDetailsRouter(adminSvc, cache) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/88/affinity-details", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Code int `json:"code"` + Data struct { + Users []struct { + UserID int64 `json:"user_id"` + UserEmail string `json:"user_email"` + ClientCount int `json:"client_count"` + IsPinned bool `json:"is_pinned"` + Clients []struct { + ClientID string `json:"client_id"` + } `json:"clients"` + } `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Equal(t, 2, resp.Data.TotalUsers) + require.Equal(t, 3, resp.Data.TotalClients) + require.Equal(t, []int64{42, 99}, resp.Data.PinnedUsers) + require.Len(t, resp.Data.Users, 2) + + usersByID := make(map[int64]struct { + UserEmail string + ClientCount int + IsPinned bool + ClientLen int + }, len(resp.Data.Users)) + for _, u := range resp.Data.Users { + usersByID[u.UserID] = struct { + UserEmail string + ClientCount int + IsPinned bool + ClientLen int + }{ + UserEmail: u.UserEmail, + ClientCount: u.ClientCount, + IsPinned: u.IsPinned, + ClientLen: len(u.Clients), + } + } + require.Equal(t, "user42@example.com", usersByID[42].UserEmail) + require.Equal(t, 2, usersByID[42].ClientCount) + require.True(t, usersByID[42].IsPinned) + require.Equal(t, 2, usersByID[42].ClientLen) + + require.Equal(t, "user7", usersByID[7].UserEmail) + require.Equal(t, 1, usersByID[7].ClientCount) + require.False(t, usersByID[7].IsPinned) + require.Equal(t, 1, usersByID[7].ClientLen) +} + +func TestAccountHandlerGetAffinityDetails_DisabledReturnsEmptyContract(t *testing.T) { + adminSvc := &affinityDetailsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 89, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1}, + }, + usersByID: map[int64]service.User{}, + } + router := setupAffinityDetailsRouter(adminSvc, &affinityDetailsGatewayCacheStub{}) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/89/affinity-details", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Code int `json:"code"` + Data struct { + Users []any `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Empty(t, resp.Data.Users) + require.Equal(t, 0, resp.Data.TotalUsers) + require.Equal(t, 0, resp.Data.TotalClients) + require.Empty(t, resp.Data.PinnedUsers) +} diff --git a/backend/internal/handler/admin/account_handler_available_models_test.go b/backend/internal/handler/admin/account_handler_available_models_test.go index c5f1e2d884..79f9689a92 100644 --- a/backend/internal/handler/admin/account_handler_available_models_test.go +++ b/backend/internal/handler/admin/account_handler_available_models_test.go @@ -28,7 +28,7 @@ func (s *availableModelsAdminService) GetAccount(_ context.Context, id int64) (* func setupAvailableModelsRouter(adminSvc service.AdminService) *gin.Engine { gin.SetMode(gin.TestMode) router := gin.New() - handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) router.GET("/api/v1/admin/accounts/:id/models", handler.GetAvailableModels) return router } diff --git a/backend/internal/handler/admin/account_handler_mixed_channel_test.go b/backend/internal/handler/admin/account_handler_mixed_channel_test.go index 24ec5bcfe1..5e45619064 100644 --- a/backend/internal/handler/admin/account_handler_mixed_channel_test.go +++ b/backend/internal/handler/admin/account_handler_mixed_channel_test.go @@ -15,7 +15,7 @@ import ( func setupAccountMixedChannelRouter(adminSvc *stubAdminService) *gin.Engine { gin.SetMode(gin.TestMode) router := gin.New() - accountHandler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + accountHandler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) router.POST("/api/v1/admin/accounts/check-mixed-channel", accountHandler.CheckMixedChannel) router.POST("/api/v1/admin/accounts", accountHandler.Create) router.PUT("/api/v1/admin/accounts/:id", accountHandler.Update) @@ -111,7 +111,7 @@ func TestAccountHandlerCreateMixedChannelConflictSimplifiedResponse(t *testing.T var resp map[string]any require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) require.Equal(t, "mixed_channel_warning", resp["error"]) - require.Contains(t, resp["message"], "mixed_channel_warning") + require.Contains(t, resp["message"], "claude-max") _, hasDetails := resp["details"] _, hasRequireConfirmation := resp["require_confirmation"] require.False(t, hasDetails) @@ -140,7 +140,7 @@ func TestAccountHandlerUpdateMixedChannelConflictSimplifiedResponse(t *testing.T var resp map[string]any require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) require.Equal(t, "mixed_channel_warning", resp["error"]) - require.Contains(t, resp["message"], "mixed_channel_warning") + require.Contains(t, resp["message"], "claude-max") _, hasDetails := resp["details"] _, hasRequireConfirmation := resp["require_confirmation"] require.False(t, hasDetails) diff --git a/backend/internal/handler/admin/account_handler_passthrough_test.go b/backend/internal/handler/admin/account_handler_passthrough_test.go index d86501c047..a11f2c5195 100644 --- a/backend/internal/handler/admin/account_handler_passthrough_test.go +++ b/backend/internal/handler/admin/account_handler_passthrough_test.go @@ -29,6 +29,7 @@ func TestAccountHandler_Create_AnthropicAPIKeyPassthroughExtraForwarded(t *testi nil, nil, nil, + nil, ) router := gin.New() diff --git a/backend/internal/handler/admin/batch_update_credentials_test.go b/backend/internal/handler/admin/batch_update_credentials_test.go index 0b1b669174..f2cd1e3a1c 100644 --- a/backend/internal/handler/admin/batch_update_credentials_test.go +++ b/backend/internal/handler/admin/batch_update_credentials_test.go @@ -36,7 +36,7 @@ func (f *failingAdminService) UpdateAccount(ctx context.Context, id int64, input func setupAccountHandlerWithService(adminSvc service.AdminService) (*gin.Engine, *AccountHandler) { gin.SetMode(gin.TestMode) router := gin.New() - handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) router.POST("/api/v1/admin/accounts/batch-update-credentials", handler.BatchUpdateCredentials) return router, handler } diff --git a/backend/internal/handler/admin/gdrive_oauth_handler.go b/backend/internal/handler/admin/gdrive_oauth_handler.go new file mode 100644 index 0000000000..bdf549066c --- /dev/null +++ b/backend/internal/handler/admin/gdrive_oauth_handler.go @@ -0,0 +1,148 @@ +package admin + +import ( + "log/slog" + "net/http" + + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +// GDriveOAuthHandler 处理 Google Drive OAuth 授权流程。 +type GDriveOAuthHandler struct { + settingService *service.SettingService + gdriveOAuth *service.SoraGDriveOAuthService + gdriveStorage *service.SoraGDriveStorage +} + +// NewGDriveOAuthHandler 创建 GDrive OAuth Handler。 +func NewGDriveOAuthHandler(settingService *service.SettingService, gdriveOAuth *service.SoraGDriveOAuthService, gdriveStorage *service.SoraGDriveStorage) *GDriveOAuthHandler { + return &GDriveOAuthHandler{ + settingService: settingService, + gdriveOAuth: gdriveOAuth, + gdriveStorage: gdriveStorage, + } +} + +// StartOAuthRequest 启动 OAuth 授权请求。 +type StartOAuthRequest struct { + ClientID string `json:"client_id" binding:"required"` + ClientSecret string `json:"client_secret" binding:"required"` + RedirectURI string `json:"redirect_uri" binding:"required"` +} + +// StartOAuth 生成 Google OAuth 授权 URL。 +// POST /api/v1/admin/settings/sora-storage/gdrive-oauth/start +func (h *GDriveOAuthHandler) StartOAuth(c *gin.Context) { + if h.gdriveOAuth == nil { + response.Error(c, http.StatusInternalServerError, "GDrive OAuth service not initialized") + return + } + + var req StartOAuthRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + + authURL, state, err := h.gdriveOAuth.GenerateAuthURL(req.ClientID, req.ClientSecret, req.RedirectURI) + if err != nil { + response.Error(c, http.StatusInternalServerError, "生成授权 URL 失败: "+err.Error()) + return + } + + response.Success(c, gin.H{ + "auth_url": authURL, + "state": state, + }) +} + +// OAuthCallbackRequest OAuth 回调请求。 +type OAuthCallbackRequest struct { + ClientID string `json:"client_id" binding:"required"` + ClientSecret string `json:"client_secret" binding:"required"` + RedirectURI string `json:"redirect_uri" binding:"required"` + Code string `json:"code" binding:"required"` + ProfileID string `json:"profile_id"` // 要保存到的 profile ID(可选) +} + +// OAuthCallback 用授权码换取 refresh_token 并保存到 profile。 +// POST /api/v1/admin/settings/sora-storage/gdrive-oauth/callback +func (h *GDriveOAuthHandler) OAuthCallback(c *gin.Context) { + if h.gdriveOAuth == nil { + response.Error(c, http.StatusInternalServerError, "GDrive OAuth service not initialized") + return + } + + var req OAuthCallbackRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + + refreshToken, err := h.gdriveOAuth.ExchangeCode(c.Request.Context(), req.ClientID, req.ClientSecret, req.RedirectURI, req.Code) + if err != nil { + slog.Error("[GDriveOAuth] exchange failed", + "client_id_len", len(req.ClientID), + "client_secret_len", len(req.ClientSecret), + "redirect_uri", req.RedirectURI, + "code_len", len(req.Code), + "error", err, + ) + response.Error(c, http.StatusBadRequest, "换取 refresh_token 失败: "+err.Error()) + return + } + + // 如果指定了 profile_id,自动保存 refresh_token 到 profile + if req.ProfileID != "" { + profiles, err := h.settingService.ListSoraS3Profiles(c.Request.Context()) + if err == nil { + for _, p := range profiles.Items { + if p.ProfileID == req.ProfileID { + _, _ = h.settingService.UpdateSoraS3Profile(c.Request.Context(), req.ProfileID, &service.SoraS3Profile{ + Name: p.Name, + Provider: p.Provider, + AccessMode: p.AccessMode, + Enabled: p.Enabled, + Endpoint: p.Endpoint, + Region: p.Region, + Bucket: p.Bucket, + AccessKeyID: p.AccessKeyID, + Prefix: p.Prefix, + ForcePathStyle: p.ForcePathStyle, + CDNURL: p.CDNURL, + DefaultStorageQuotaBytes: p.DefaultStorageQuotaBytes, + AuthType: p.AuthType, + ClientID: p.ClientID, + FolderID: p.FolderID, + RefreshToken: refreshToken, + }) + break + } + } + } + } + + response.Success(c, gin.H{ + "refresh_token": refreshToken, + "message": "OAuth 授权成功", + }) +} + +// TestGDriveStorage 测试 GDrive 存储的完整上传→下载→删除流程。 +// POST /api/v1/admin/settings/sora-storage/gdrive-test +func (h *GDriveOAuthHandler) TestGDriveStorage(c *gin.Context) { + if h.gdriveStorage == nil { + response.Error(c, http.StatusInternalServerError, "GDrive storage not initialized") + return + } + + result, err := h.gdriveStorage.TestFullCycle(c.Request.Context()) + if err != nil { + response.Error(c, http.StatusBadRequest, "GDrive 测试失败: "+err.Error()) + return + } + + response.Success(c, result) +} diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 459fd94911..c7ef313541 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -103,9 +103,10 @@ type CreateGroupRequest struct { FallbackGroupID *int64 `json:"fallback_group_id"` FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) - ModelRouting map[string][]int64 `json:"model_routing"` - ModelRoutingEnabled bool `json:"model_routing_enabled"` - MCPXMLInject *bool `json:"mcp_xml_inject"` + ModelRouting map[string][]int64 `json:"model_routing"` + ModelRoutingEnabled bool `json:"model_routing_enabled"` + MCPXMLInject *bool `json:"mcp_xml_inject"` + SimulateClaudeMaxEnabled *bool `json:"simulate_claude_max_enabled"` // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes []string `json:"supported_model_scopes"` // Sora 存储配额 @@ -141,9 +142,10 @@ type UpdateGroupRequest struct { FallbackGroupID *int64 `json:"fallback_group_id"` FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) - ModelRouting map[string][]int64 `json:"model_routing"` - ModelRoutingEnabled *bool `json:"model_routing_enabled"` - MCPXMLInject *bool `json:"mcp_xml_inject"` + ModelRouting map[string][]int64 `json:"model_routing"` + ModelRoutingEnabled *bool `json:"model_routing_enabled"` + MCPXMLInject *bool `json:"mcp_xml_inject"` + SimulateClaudeMaxEnabled *bool `json:"simulate_claude_max_enabled"` // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes *[]string `json:"supported_model_scopes"` // Sora 存储配额 @@ -264,6 +266,7 @@ func (h *GroupHandler) Create(c *gin.Context) { ModelRouting: req.ModelRouting, ModelRoutingEnabled: req.ModelRoutingEnabled, MCPXMLInject: req.MCPXMLInject, + SimulateClaudeMaxEnabled: req.SimulateClaudeMaxEnabled, SupportedModelScopes: req.SupportedModelScopes, SoraStorageQuotaBytes: req.SoraStorageQuotaBytes, AllowMessagesDispatch: req.AllowMessagesDispatch, @@ -317,6 +320,7 @@ func (h *GroupHandler) Update(c *gin.Context) { ModelRouting: req.ModelRouting, ModelRoutingEnabled: req.ModelRoutingEnabled, MCPXMLInject: req.MCPXMLInject, + SimulateClaudeMaxEnabled: req.SimulateClaudeMaxEnabled, SupportedModelScopes: req.SupportedModelScopes, SoraStorageQuotaBytes: req.SoraStorageQuotaBytes, AllowMessagesDispatch: req.AllowMessagesDispatch, diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 25456bb3ac..288103337d 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -37,21 +37,25 @@ func generateMenuItemID() (string, error) { // SettingHandler 系统设置处理器 type SettingHandler struct { - settingService *service.SettingService - emailService *service.EmailService - turnstileService *service.TurnstileService - opsService *service.OpsService - soraS3Storage *service.SoraS3Storage + settingService *service.SettingService + emailService *service.EmailService + turnstileService *service.TurnstileService + opsService *service.OpsService + soraS3Storage *service.SoraS3Storage + soraGDriveStorage *service.SoraGDriveStorage + soraGenerationService *service.SoraGenerationService } // NewSettingHandler 创建系统设置处理器 -func NewSettingHandler(settingService *service.SettingService, emailService *service.EmailService, turnstileService *service.TurnstileService, opsService *service.OpsService, soraS3Storage *service.SoraS3Storage) *SettingHandler { +func NewSettingHandler(settingService *service.SettingService, emailService *service.EmailService, turnstileService *service.TurnstileService, opsService *service.OpsService, soraS3Storage *service.SoraS3Storage, soraGDriveStorage *service.SoraGDriveStorage, soraGenerationService *service.SoraGenerationService) *SettingHandler { return &SettingHandler{ - settingService: settingService, - emailService: emailService, - turnstileService: turnstileService, - opsService: opsService, - soraS3Storage: soraS3Storage, + settingService: settingService, + emailService: emailService, + turnstileService: turnstileService, + opsService: opsService, + soraS3Storage: soraS3Storage, + soraGDriveStorage: soraGDriveStorage, + soraGenerationService: soraGenerationService, } } @@ -1070,6 +1074,8 @@ func toSoraS3ProfileDTO(profile service.SoraS3Profile) dto.SoraS3Profile { ProfileID: profile.ProfileID, Name: profile.Name, IsActive: profile.IsActive, + Provider: profile.GetProvider(), + AccessMode: profile.AccessMode, Enabled: profile.Enabled, Endpoint: profile.Endpoint, Region: profile.Region, @@ -1081,6 +1087,13 @@ func toSoraS3ProfileDTO(profile service.SoraS3Profile) dto.SoraS3Profile { CDNURL: profile.CDNURL, DefaultStorageQuotaBytes: profile.DefaultStorageQuotaBytes, UpdatedAt: profile.UpdatedAt, + // Google Drive 专属 + AuthType: profile.AuthType, + ClientID: profile.ClientID, + ClientSecretConfigured: profile.ClientSecretConfigured, + RefreshTokenConfigured: profile.RefreshTokenConfigured, + ServiceAccountConfigured: profile.ServiceAccountConfigured, + FolderID: profile.FolderID, } } @@ -1160,6 +1173,8 @@ type CreateSoraS3ProfileRequest struct { ProfileID string `json:"profile_id"` Name string `json:"name"` SetActive bool `json:"set_active"` + Provider string `json:"provider"` // "s3" / "gdrive" + AccessMode string `json:"access_mode"` // "direct" / "proxy" Enabled bool `json:"enabled"` Endpoint string `json:"endpoint"` Region string `json:"region"` @@ -1170,10 +1185,19 @@ type CreateSoraS3ProfileRequest struct { ForcePathStyle bool `json:"force_path_style"` CDNURL string `json:"cdn_url"` DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` + // Google Drive 专属 + AuthType string `json:"auth_type,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` + RefreshToken string `json:"refresh_token,omitempty"` + ServiceAccountJSON string `json:"service_account_json,omitempty"` + FolderID string `json:"folder_id,omitempty"` } type UpdateSoraS3ProfileRequest struct { Name string `json:"name"` + Provider string `json:"provider"` + AccessMode string `json:"access_mode"` Enabled bool `json:"enabled"` Endpoint string `json:"endpoint"` Region string `json:"region"` @@ -1184,6 +1208,13 @@ type UpdateSoraS3ProfileRequest struct { ForcePathStyle bool `json:"force_path_style"` CDNURL string `json:"cdn_url"` DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` + // Google Drive 专属 + AuthType string `json:"auth_type,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` + RefreshToken string `json:"refresh_token,omitempty"` + ServiceAccountJSON string `json:"service_account_json,omitempty"` + FolderID string `json:"folder_id,omitempty"` } // CreateSoraS3Profile 创建 Sora S3 配置 @@ -1206,14 +1237,23 @@ func (h *SettingHandler) CreateSoraS3Profile(c *gin.Context) { response.BadRequest(c, "Profile ID is required") return } - if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, false); err != nil { - response.BadRequest(c, err.Error()) - return + // S3 专属字段验证:仅当 provider 为 s3(或未指定)时校验 + provider := req.Provider + if provider == "" { + provider = "s3" + } + if provider == "s3" { + if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, false); err != nil { + response.BadRequest(c, err.Error()) + return + } } created, err := h.settingService.CreateSoraS3Profile(c.Request.Context(), &service.SoraS3Profile{ ProfileID: req.ProfileID, Name: req.Name, + Provider: req.Provider, + AccessMode: req.AccessMode, Enabled: req.Enabled, Endpoint: req.Endpoint, Region: req.Region, @@ -1224,6 +1264,13 @@ func (h *SettingHandler) CreateSoraS3Profile(c *gin.Context) { ForcePathStyle: req.ForcePathStyle, CDNURL: req.CDNURL, DefaultStorageQuotaBytes: req.DefaultStorageQuotaBytes, + // Google Drive 专属 + AuthType: req.AuthType, + ClientID: req.ClientID, + ClientSecret: req.ClientSecret, + RefreshToken: req.RefreshToken, + ServiceAccountJSON: req.ServiceAccountJSON, + FolderID: req.FolderID, }, req.SetActive) if err != nil { response.ErrorFrom(c, err) @@ -1266,13 +1313,25 @@ func (h *SettingHandler) UpdateSoraS3Profile(c *gin.Context) { response.ErrorFrom(c, service.ErrSoraS3ProfileNotFound) return } - if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, existing.SecretAccessKeyConfigured); err != nil { - response.BadRequest(c, err.Error()) - return + // S3 专属字段验证 + provider := req.Provider + if provider == "" && existing != nil { + provider = existing.GetProvider() + } + if provider == "" { + provider = "s3" + } + if provider == "s3" { + if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, existing.SecretAccessKeyConfigured); err != nil { + response.BadRequest(c, err.Error()) + return + } } updated, updateErr := h.settingService.UpdateSoraS3Profile(c.Request.Context(), profileID, &service.SoraS3Profile{ Name: req.Name, + Provider: req.Provider, + AccessMode: req.AccessMode, Enabled: req.Enabled, Endpoint: req.Endpoint, Region: req.Region, @@ -1283,6 +1342,13 @@ func (h *SettingHandler) UpdateSoraS3Profile(c *gin.Context) { ForcePathStyle: req.ForcePathStyle, CDNURL: req.CDNURL, DefaultStorageQuotaBytes: req.DefaultStorageQuotaBytes, + // Google Drive 专属 + AuthType: req.AuthType, + ClientID: req.ClientID, + ClientSecret: req.ClientSecret, + RefreshToken: req.RefreshToken, + ServiceAccountJSON: req.ServiceAccountJSON, + FolderID: req.FolderID, }) if updateErr != nil { response.ErrorFrom(c, updateErr) @@ -1583,3 +1649,44 @@ func (h *SettingHandler) UpdateStreamTimeoutSettings(c *gin.Context) { ThresholdWindowMinutes: updatedSettings.ThresholdWindowMinutes, }) } + +// GetGDriveQuota 获取 Google Drive 配额信息。 +// GET /api/v1/admin/settings/sora-storage/gdrive-quota +func (h *SettingHandler) GetGDriveQuota(c *gin.Context) { + if h.soraGDriveStorage == nil { + response.Error(c, http.StatusServiceUnavailable, "GDrive storage not configured") + return + } + quota, err := h.soraGDriveStorage.GetQuotaInfo(c.Request.Context()) + if err != nil { + response.Error(c, http.StatusInternalServerError, fmt.Sprintf("failed to get GDrive quota: %v", err)) + return + } + response.Success(c, quota) +} + +// GetStorageVideoStats 获取各存储类型的视频统计信息。 +// GET /api/v1/admin/settings/sora-storage/video-stats +func (h *SettingHandler) GetStorageVideoStats(c *gin.Context) { + if h.soraGenerationService == nil { + response.Error(c, http.StatusServiceUnavailable, "generation service not configured") + return + } + + storageTypes := []string{service.SoraStorageTypeS3, service.SoraStorageTypeGDrive} + result := make(map[string]*service.StorageVideoStats, len(storageTypes)) + + for _, st := range storageTypes { + completed, inProgress, err := h.soraGenerationService.CountByStorageType(c.Request.Context(), st) + if err != nil { + log.Printf("[SettingHandler] CountByStorageType(%s) error: %v", st, err) + continue + } + result[st] = &service.StorageVideoStats{ + Completed: completed, + InProgress: inProgress, + } + } + + response.Success(c, result) +} diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index d1d867ee88..7952a21377 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -135,16 +135,17 @@ func GroupFromServiceAdmin(g *service.Group) *AdminGroup { return nil } out := &AdminGroup{ - Group: groupFromServiceBase(g), - ModelRouting: g.ModelRouting, - ModelRoutingEnabled: g.ModelRoutingEnabled, - MCPXMLInject: g.MCPXMLInject, - DefaultMappedModel: g.DefaultMappedModel, - SupportedModelScopes: g.SupportedModelScopes, - AccountCount: g.AccountCount, - ActiveAccountCount: g.ActiveAccountCount, - RateLimitedAccountCount: g.RateLimitedAccountCount, - SortOrder: g.SortOrder, + Group: groupFromServiceBase(g), + ModelRouting: g.ModelRouting, + ModelRoutingEnabled: g.ModelRoutingEnabled, + MCPXMLInject: g.MCPXMLInject, + DefaultMappedModel: g.DefaultMappedModel, + SimulateClaudeMaxEnabled: g.SimulateClaudeMaxEnabled, + SupportedModelScopes: g.SupportedModelScopes, + AccountCount: g.AccountCount, + ActiveAccountCount: g.ActiveAccountCount, + RateLimitedAccountCount: g.RateLimitedAccountCount, + SortOrder: g.SortOrder, } if len(g.AccountGroups) > 0 { out.AccountGroups = make([]AccountGroup, 0, len(g.AccountGroups)) @@ -266,6 +267,17 @@ func AccountFromServiceShallow(a *service.Account) *Account { } } + // 客户端亲和调度(Anthropic 和 Antigravity 账号) + if a.IsAffinityEnabled() { + enabled := true + out.ClientAffinityEnabled = &enabled + allow := a.IsAffinityAllowSwitch() + out.AffinityAllowSwitch = &allow + if pinnedUsers := a.GetPinnedUsers(); len(pinnedUsers) > 0 { + out.PinnedUserIDs = pinnedUsers + } + } + // 提取账号配额限制(apikey / bedrock 类型有效) if a.IsAPIKeyOrBedrock() { if limit := a.GetQuotaLimit(); limit > 0 { diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index b953e33641..452ed749a5 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -133,11 +133,13 @@ type SoraS3Settings struct { DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` } -// SoraS3Profile Sora S3 存储配置项 DTO(响应用,不含敏感字段) +// SoraS3Profile Sora 存储配置项 DTO(响应用,不含敏感字段) type SoraS3Profile struct { ProfileID string `json:"profile_id"` Name string `json:"name"` IsActive bool `json:"is_active"` + Provider string `json:"provider"` // "s3" / "gdrive" + AccessMode string `json:"access_mode"` // "direct" / "proxy" Enabled bool `json:"enabled"` Endpoint string `json:"endpoint"` Region string `json:"region"` @@ -149,6 +151,14 @@ type SoraS3Profile struct { CDNURL string `json:"cdn_url"` DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` UpdatedAt string `json:"updated_at"` + + // --- Google Drive 专属 --- + AuthType string `json:"auth_type,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecretConfigured bool `json:"client_secret_configured"` + RefreshTokenConfigured bool `json:"refresh_token_configured"` + ServiceAccountConfigured bool `json:"service_account_configured"` + FolderID string `json:"folder_id,omitempty"` } // ListSoraS3ProfilesResponse Sora S3 配置列表响应 diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 7b3443be62..cddc430353 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -117,6 +117,8 @@ type AdminGroup struct { // MCP XML 协议注入(仅 antigravity 平台使用) MCPXMLInject bool `json:"mcp_xml_inject"` + // Claude usage 模拟开关(仅管理员可见) + SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled"` // OpenAI Messages 调度配置(仅 openai 平台使用) DefaultMappedModel string `json:"default_mapped_model"` @@ -197,6 +199,23 @@ type Account struct { CacheTTLOverrideEnabled *bool `json:"cache_ttl_override_enabled,omitempty"` CacheTTLOverrideTarget *string `json:"cache_ttl_override_target,omitempty"` + // 客户端亲和调度(Anthropic 和 Antigravity 账号有效) + // 启用后新会话会优先调度到客户端之前使用过的账号 + ClientAffinityEnabled *bool `json:"client_affinity_enabled,omitempty"` + + // 亲和允许切换(默认 true) + AffinityAllowSwitch *bool `json:"affinity_allow_switch,omitempty"` + + // 亲和用户数量(admin 列表端点注入) + AffinityUserCount *int64 `json:"affinity_user_count,omitempty"` + + // 指定亲和用户 ID 列表 + PinnedUserIDs []int64 `json:"pinned_user_ids,omitempty"` + + // 亲和客户端数据(仅 admin 列表端点注入,不由 mapper 填充) + AffinityClientCount *int64 `json:"affinity_client_count,omitempty"` + AffinityClients []string `json:"affinity_clients,omitempty"` + // API Key 账号配额限制 QuotaLimit *float64 `json:"quota_limit,omitempty"` QuotaUsed *float64 `json:"quota_used,omitempty"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 831029c48b..80d9d58dd4 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -291,7 +291,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { } for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "") // Gemini 不使用会话限制 + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制 if err != nil { if len(fs.FailedAccountIDs) == 0 { h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted) @@ -453,6 +453,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.submitUsageRecordTask(func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, + ParsedRequest: parsedReq, APIKey: apiKey, User: apiKey.User, Account: account, @@ -500,8 +501,16 @@ func (h *GatewayHandler) Messages(c *gin.Context) { for { // 选择支持该模型的账号 - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID) + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID, subject.UserID) if err != nil { + if errors.Is(err, service.ErrAffinityNoSwitch) { + h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "Affinity account unavailable and switching is disabled", streamStarted) + return + } + if errors.Is(err, service.ErrAffinityLimitExceeded) { + h.handleStreamingAwareError(c, http.StatusTooManyRequests, "api_error", "Affinity client limit exceeded", streamStarted) + return + } if len(fs.FailedAccountIDs) == 0 { h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted) return @@ -647,6 +656,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { // ===== 用户消息串行队列 END ===== // 转发请求 - 根据账号平台分流 + c.Set("parsed_request", parsedReq) var result *service.ForwardResult requestCtx := c.Request.Context() if fs.SwitchCount > 0 { @@ -772,6 +782,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.submitUsageRecordTask(func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, + ParsedRequest: parsedReq, APIKey: currentAPIKey, User: currentAPIKey.User, Account: account, diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index cfe809114b..4dc3b282b5 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -352,7 +352,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { } for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "") // Gemini 不使用会话限制 + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制 if err != nil { if len(fs.FailedAccountIDs) == 0 { googleError(c, http.StatusServiceUnavailable, "No available Gemini accounts: "+err.Error()) diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index 89d556cc1d..6cbebaeee4 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -29,6 +29,7 @@ type AdminHandlers struct { ErrorPassthrough *admin.ErrorPassthroughHandler APIKey *admin.AdminAPIKeyHandler ScheduledTest *admin.ScheduledTestHandler + GDriveOAuth *admin.GDriveOAuthHandler } // Handlers contains all HTTP handlers @@ -45,6 +46,7 @@ type Handlers struct { OpenAIGateway *OpenAIGatewayHandler SoraGateway *SoraGatewayHandler SoraClient *SoraClientHandler + SoraVideos *SoraVideosHandler Setting *SettingHandler Totp *TotpHandler } diff --git a/backend/internal/handler/sora_client_handler.go b/backend/internal/handler/sora_client_handler.go index 80acc83349..3cb85d18b6 100644 --- a/backend/internal/handler/sora_client_handler.go +++ b/backend/internal/handler/sora_client_handler.go @@ -31,7 +31,7 @@ const ( type SoraClientHandler struct { genService *service.SoraGenerationService quotaService *service.SoraQuotaService - s3Storage *service.SoraS3Storage + objectStorage service.SoraObjectStorage soraGatewayService *service.SoraGatewayService gatewayService *service.GatewayService mediaStorage *service.SoraMediaStorage @@ -48,7 +48,7 @@ type SoraClientHandler struct { func NewSoraClientHandler( genService *service.SoraGenerationService, quotaService *service.SoraQuotaService, - s3Storage *service.SoraS3Storage, + objectStorage service.SoraObjectStorage, soraGatewayService *service.SoraGatewayService, gatewayService *service.GatewayService, mediaStorage *service.SoraMediaStorage, @@ -57,7 +57,7 @@ func NewSoraClientHandler( return &SoraClientHandler{ genService: genService, quotaService: quotaService, - s3Storage: s3Storage, + objectStorage: objectStorage, soraGatewayService: soraGatewayService, gatewayService: gatewayService, mediaStorage: mediaStorage, @@ -291,11 +291,11 @@ func (h *SoraClientHandler) processGeneration(genID int64, userID int64, groupID return } - // 三层降级存储:S3 → 本地 → 上游临时 URL + // 三层降级存储:对象存储 → 本地 → 上游临时 URL storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation(ctx, userID, mediaType, mediaURL, mediaURLs) usageAdded := false - if (storageType == service.SoraStorageTypeS3 || storageType == service.SoraStorageTypeLocal) && fileSize > 0 && h.quotaService != nil { + if (service.IsObjectStorageType(storageType) || storageType == service.SoraStorageTypeLocal) && fileSize > 0 && h.quotaService != nil { if err := h.quotaService.AddUsage(ctx, userID, fileSize); err != nil { h.cleanupStoredMedia(ctx, storageType, s3Keys, storedURLs) var quotaErr *service.QuotaExceededError @@ -346,39 +346,41 @@ func (h *SoraClientHandler) storeMediaWithDegradation( urls = []string{mediaURL} } - // 第一层:尝试 S3 - if h.s3Storage != nil && h.s3Storage.Enabled(ctx) { + // 第一层:尝试对象存储(S3 / Google Drive) + if h.objectStorage != nil && h.objectStorage.Enabled(ctx) { keys := make([]string, 0, len(urls)) var totalSize int64 + var actualStorageType string allOK := true for _, u := range urls { - key, size, err := h.s3Storage.UploadFromURL(ctx, userID, u) + key, size, st, err := h.objectStorage.UploadFromURL(ctx, userID, u) if err != nil { - logger.LegacyPrintf("handler.sora_client", "[SoraClient] S3 上传失败 err=%v", err) + logger.LegacyPrintf("handler.sora_client", "[SoraClient] 对象存储上传失败 type=%s err=%v", h.objectStorage.StorageType(), err) allOK = false // 清理已上传的文件 if len(keys) > 0 { - _ = h.s3Storage.DeleteObjects(ctx, keys) + _ = h.objectStorage.DeleteObjects(ctx, keys) } break } keys = append(keys, key) totalSize += size + actualStorageType = st } if allOK && len(keys) > 0 { accessURLs := make([]string, 0, len(keys)) for _, key := range keys { - accessURL, err := h.s3Storage.GetAccessURL(ctx, key) + accessURL, err := h.objectStorage.GetAccessURL(ctx, key) if err != nil { - logger.LegacyPrintf("handler.sora_client", "[SoraClient] 生成 S3 访问 URL 失败 err=%v", err) - _ = h.s3Storage.DeleteObjects(ctx, keys) + logger.LegacyPrintf("handler.sora_client", "[SoraClient] 生成访问 URL 失败 type=%s err=%v", h.objectStorage.StorageType(), err) + _ = h.objectStorage.DeleteObjects(ctx, keys) allOK = false break } accessURLs = append(accessURLs, accessURL) } if allOK && len(accessURLs) > 0 { - return accessURLs[0], accessURLs, service.SoraStorageTypeS3, keys, totalSize + return accessURLs[0], accessURLs, actualStorageType, keys, totalSize } } } @@ -678,7 +680,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { return } - if h.s3Storage == nil || !h.s3Storage.Enabled(c.Request.Context()) { + if h.objectStorage == nil || !h.objectStorage.Enabled(c.Request.Context()) { response.Error(c, http.StatusServiceUnavailable, "云存储未配置,请联系管理员") return } @@ -697,24 +699,24 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { var totalSize int64 for _, sourceURL := range sourceURLs { - objectKey, fileSize, uploadErr := h.s3Storage.UploadFromURL(c.Request.Context(), userID, sourceURL) + objectKey, fileSize, _, uploadErr := h.objectStorage.UploadFromURL(c.Request.Context(), userID, sourceURL) if uploadErr != nil { if len(uploadedKeys) > 0 { - _ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys) + _ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys) } var upstreamErr *service.UpstreamDownloadError if errors.As(uploadErr, &upstreamErr) && (upstreamErr.StatusCode == http.StatusForbidden || upstreamErr.StatusCode == http.StatusNotFound) { response.Error(c, http.StatusGone, "媒体链接已过期,无法保存") return } - response.Error(c, http.StatusInternalServerError, "上传到 S3 失败: "+uploadErr.Error()) + response.Error(c, http.StatusInternalServerError, "上传到存储失败: "+uploadErr.Error()) return } - accessURL, err := h.s3Storage.GetAccessURL(c.Request.Context(), objectKey) + accessURL, err := h.objectStorage.GetAccessURL(c.Request.Context(), objectKey) if err != nil { uploadedKeys = append(uploadedKeys, objectKey) - _ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys) - response.Error(c, http.StatusInternalServerError, "生成 S3 访问链接失败: "+err.Error()) + _ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys) + response.Error(c, http.StatusInternalServerError, "生成访问链接失败: "+err.Error()) return } uploadedKeys = append(uploadedKeys, objectKey) @@ -725,7 +727,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { usageAdded := false if totalSize > 0 && h.quotaService != nil { if err := h.quotaService.AddUsage(c.Request.Context(), userID, totalSize); err != nil { - _ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys) + _ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys) var quotaErr *service.QuotaExceededError if errors.As(err, "aErr) { response.Error(c, http.StatusTooManyRequests, "存储配额已满,请删除不需要的作品释放空间") @@ -742,11 +744,11 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { id, accessURLs[0], accessURLs, - service.SoraStorageTypeS3, + h.objectStorage.StorageType(), uploadedKeys, totalSize, ); err != nil { - _ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys) + _ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys) if usageAdded && h.quotaService != nil { _ = h.quotaService.ReleaseUsage(c.Request.Context(), userID, totalSize) } @@ -755,7 +757,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { } response.Success(c, gin.H{ - "message": "已保存到 S3", + "message": "已保存到云存储", "object_key": uploadedKeys[0], "object_keys": uploadedKeys, }) @@ -764,28 +766,30 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) { // GetStorageStatus 返回存储状态。 // GET /api/v1/sora/storage-status func (h *SoraClientHandler) GetStorageStatus(c *gin.Context) { - s3Enabled := h.s3Storage != nil && h.s3Storage.Enabled(c.Request.Context()) - s3Healthy := false - if s3Enabled { - s3Healthy = h.s3Storage.IsHealthy(c.Request.Context()) + objectStorageEnabled := h.objectStorage != nil && h.objectStorage.Enabled(c.Request.Context()) + objectStorageHealthy := false + storageType := "" + if objectStorageEnabled { + objectStorageHealthy = h.objectStorage.IsHealthy(c.Request.Context()) + storageType = h.objectStorage.StorageType() } localEnabled := h.mediaStorage != nil && h.mediaStorage.Enabled() response.Success(c, gin.H{ - "s3_enabled": s3Enabled, - "s3_healthy": s3Healthy, + "s3_enabled": objectStorageEnabled, // 保留字段名向后兼容 + "s3_healthy": objectStorageHealthy, + "storage_type": storageType, "local_enabled": localEnabled, }) } func (h *SoraClientHandler) cleanupStoredMedia(ctx context.Context, storageType string, s3Keys []string, localPaths []string) { - switch storageType { - case service.SoraStorageTypeS3: - if h.s3Storage != nil && len(s3Keys) > 0 { - if err := h.s3Storage.DeleteObjects(ctx, s3Keys); err != nil { - logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理 S3 文件失败 keys=%v err=%v", s3Keys, err) + if service.IsObjectStorageType(storageType) { + if h.objectStorage != nil && len(s3Keys) > 0 { + if err := h.objectStorage.DeleteObjects(ctx, s3Keys); err != nil { + logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理存储文件失败 type=%s keys=%v err=%v", storageType, s3Keys, err) } } - case service.SoraStorageTypeLocal: + } else if storageType == service.SoraStorageTypeLocal { if h.mediaStorage != nil && len(localPaths) > 0 { if err := h.mediaStorage.DeleteByRelativePaths(localPaths); err != nil { logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理本地文件失败 paths=%v err=%v", localPaths, err) diff --git a/backend/internal/handler/sora_client_handler_test.go b/backend/internal/handler/sora_client_handler_test.go index dab17673eb..92e77d9251 100644 --- a/backend/internal/handler/sora_client_handler_test.go +++ b/backend/internal/handler/sora_client_handler_test.go @@ -124,6 +124,9 @@ func (r *stubSoraGenRepo) CountByUserAndStatus(_ context.Context, _ int64, _ []s } return r.countValue, nil } +func (r *stubSoraGenRepo) CountByStorageType(_ context.Context, _ string, _ []string) (int64, error) { + return 0, nil +} // ==================== 辅助函数 ==================== @@ -1641,7 +1644,7 @@ func TestStoreMediaWithDegradation_S3SuccessSingleURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation( context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil, @@ -1663,7 +1666,7 @@ func TestStoreMediaWithDegradation_S3SuccessMultiURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"} storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation( @@ -1688,7 +1691,7 @@ func TestStoreMediaWithDegradation_S3DownloadFails(t *testing.T) { defer badSource.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} _, _, storageType, _, _ := h.storeMediaWithDegradation( context.Background(), 1, "video", badSource.URL+"/missing.mp4", nil, @@ -1703,7 +1706,7 @@ func TestStoreMediaWithDegradation_S3FailsSingleURL(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} _, _, storageType, s3Keys, _ := h.storeMediaWithDegradation( context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil, @@ -1720,7 +1723,7 @@ func TestStoreMediaWithDegradation_S3PartialFailureCleanup(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"} _, _, storageType, s3Keys, _ := h.storeMediaWithDegradation( @@ -1804,8 +1807,8 @@ func TestStoreMediaWithDegradation_S3FailsFallbackToLocal(t *testing.T) { } mediaStorage := service.NewSoraMediaStorage(cfg) h := &SoraClientHandler{ - s3Storage: s3Storage, - mediaStorage: mediaStorage, + objectStorage: s3Storage, + mediaStorage: mediaStorage, } _, _, storageType, _, _ := h.storeMediaWithDegradation( @@ -1831,14 +1834,14 @@ func TestSaveToStorage_S3EnabledButUploadFails(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} h.SaveToStorage(c) require.Equal(t, http.StatusInternalServerError, rec.Code) resp := parseResponse(t, rec) - require.Contains(t, resp["message"], "S3") + require.Contains(t, resp["message"], "上传到存储失败") } func TestSaveToStorage_UpstreamURLExpired(t *testing.T) { @@ -1857,7 +1860,7 @@ func TestSaveToStorage_UpstreamURLExpired(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1881,7 +1884,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1889,7 +1892,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) { require.Equal(t, http.StatusOK, rec.Code) resp := parseResponse(t, rec) data := resp["data"].(map[string]any) - require.Contains(t, data["message"], "S3") + require.Contains(t, data["message"], "已保存到云存储") require.NotEmpty(t, data["object_key"]) // 验证记录已更新为 S3 存储 require.Equal(t, service.SoraStorageTypeS3, repo.gens[1].StorageType) @@ -1913,7 +1916,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess_MultiMediaURLs(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1949,7 +1952,7 @@ func TestSaveToStorage_S3EnabledUploadSuccessWithQuota(t *testing.T) { SoraStorageUsedBytes: 0, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1975,7 +1978,7 @@ func TestSaveToStorage_S3UploadSuccessMarkCompletedFails(t *testing.T) { repo.updateErr = fmt.Errorf("db error") s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -1991,7 +1994,7 @@ func TestGetStorageStatus_S3EnabledNotHealthy(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0) h.GetStorageStatus(c) @@ -2007,7 +2010,7 @@ func TestGetStorageStatus_S3EnabledHealthy(t *testing.T) { defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0) h.GetStorageStatus(c) @@ -2447,7 +2450,7 @@ func TestProcessGeneration_FullSuccessWithS3(t *testing.T) { genService: genService, gatewayService: gatewayService, soraGatewayService: soraGatewayService, - s3Storage: s3Storage, + objectStorage: s3Storage, quotaService: quotaService, } @@ -2497,7 +2500,7 @@ func TestProcessGeneration_MarkCompletedFails(t *testing.T) { // ==================== cleanupStoredMedia 直接测试 ==================== func TestCleanupStoredMedia_S3Path(t *testing.T) { - // S3 清理路径:s3Storage 为 nil 时不 panic + // S3 清理路径:objectStorage 为 nil 时不 panic h := &SoraClientHandler{} // 不应 panic h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil) @@ -2955,7 +2958,7 @@ func TestSaveToStorage_QuotaExceeded(t *testing.T) { SoraStorageUsedBytes: 10, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -2983,7 +2986,7 @@ func TestSaveToStorage_QuotaNonQuotaError(t *testing.T) { // 用户不存在 → GetByID 失败 → AddUsage 返回普通 error userRepo := newStubUserRepoForHandler() quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3006,7 +3009,7 @@ func TestSaveToStorage_EmptyMediaURLs(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3033,7 +3036,7 @@ func TestSaveToStorage_MultiURL_SecondUploadFails(t *testing.T) { } s3Storage := newS3StorageForHandler(fakeS3.URL) genService := service.NewSoraGenerationService(repo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3066,7 +3069,7 @@ func TestSaveToStorage_MarkCompletedFailsWithQuotaRollback(t *testing.T) { SoraStorageUsedBytes: 0, } quotaService := service.NewSoraQuotaService(userRepo, nil, nil) - h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService} + h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService} c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1) c.Params = gin.Params{{Key: "id", Value: "1"}} @@ -3080,7 +3083,7 @@ func TestCleanupStoredMedia_WithS3Storage_ActualDelete(t *testing.T) { fakeS3 := newFakeS3Server("ok") defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1", "key2"}, nil) } @@ -3089,7 +3092,7 @@ func TestCleanupStoredMedia_S3DeleteFails_LogOnly(t *testing.T) { fakeS3 := newFakeS3Server("fail") defer fakeS3.Close() s3Storage := newS3StorageForHandler(fakeS3.URL) - h := &SoraClientHandler{s3Storage: s3Storage} + h := &SoraClientHandler{objectStorage: s3Storage} h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil) } diff --git a/backend/internal/handler/sora_gateway_handler.go b/backend/internal/handler/sora_gateway_handler.go index dc301ce149..3cc6a0397e 100644 --- a/backend/internal/handler/sora_gateway_handler.go +++ b/backend/internal/handler/sora_gateway_handler.go @@ -225,7 +225,7 @@ func (h *SoraGatewayHandler) ChatCompletions(c *gin.Context) { var lastFailoverHeaders http.Header for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "") + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "", int64(0)) if err != nil { reqLog.Warn("sora.account_select_failed", zap.Error(err), diff --git a/backend/internal/handler/sora_videos_handler.go b/backend/internal/handler/sora_videos_handler.go new file mode 100644 index 0000000000..b2a9d79c81 --- /dev/null +++ b/backend/internal/handler/sora_videos_handler.go @@ -0,0 +1,391 @@ +package handler + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "unicode/utf8" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + + "github.com/gin-gonic/gin" +) + +// SoraVideosHandler handles Sora video/image async task API. +type SoraVideosHandler struct { + taskService *service.SoraTaskService + gatewayService *service.GatewayService + objectStorage service.SoraObjectStorage + mediaStorage *service.SoraMediaStorage + soraGatewayService *service.SoraGatewayService +} + +func NewSoraVideosHandler( + taskService *service.SoraTaskService, + gatewayService *service.GatewayService, + objectStorage service.SoraObjectStorage, + mediaStorage *service.SoraMediaStorage, + soraGatewayService *service.SoraGatewayService, +) *SoraVideosHandler { + if taskService == nil { + return nil + } + return &SoraVideosHandler{ + taskService: taskService, + gatewayService: gatewayService, + objectStorage: objectStorage, + mediaStorage: mediaStorage, + soraGatewayService: soraGatewayService, + } +} + +func (h *SoraVideosHandler) CreateVideo(c *gin.Context) { + apiKey, account, release, ok := h.selectAccount(c, "") + if !ok { + return + } + defer release() + + body, err := readBody(c) + if err != nil { + return + } + + var req service.CreateVideoRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Model == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "model is required") + return + } + + task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, &req, body) + if err != nil { + handleTaskCreateError(c, "CreateVideo", err) + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) GetVideo(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) RemixVideo(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + body, err := readBody(c) + if err != nil { + return + } + + var req service.RemixRequest + if err := json.Unmarshal(body, &req); err != nil || req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + + originalTask, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Original video task not found") + return + } + if originalTask.Status != service.SoraTaskCompleted { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Original video must be completed before remix") + return + } + + remixTargetID := originalTask.ShareID + if remixTargetID == "" { + remixTargetID = originalTask.UpstreamTaskID + } + + account, err := h.selectAccountByID(c, originalTask.AccountID) + if err != nil { + soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "Failed to get account") + return + } + + videoReq := &service.CreateVideoRequest{ + Model: originalTask.Model, + Prompt: req.Prompt, + RemixTargetID: remixTargetID, + } + reqBody, _ := json.Marshal(videoReq) + + task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, videoReq, reqBody) + if err != nil { + handleTaskCreateError(c, "RemixVideo", err) + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +// GetVideoContent returns video content based on storage configuration. +func (h *SoraVideosHandler) GetVideoContent(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found") + return + } + + switch task.Status { + case service.SoraTaskCompleted: + contentURL := h.resolveContentURL(c, task) + if contentURL == "" { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Video URL not available") + return + } + c.Redirect(http.StatusFound, contentURL) + + case service.SoraTaskFailed: + c.JSON(http.StatusGone, gin.H{ + "id": task.ID, + "object": task.ObjectType, + "status": task.Status, + "error": gin.H{ + "message": task.ErrorMessage, + "type": task.ErrorType, + }, + }) + + default: + c.JSON(http.StatusAccepted, service.TaskToResponse(task)) + } +} + +func (h *SoraVideosHandler) resolveContentURL(c *gin.Context, task *service.SoraTask) string { + if task.StoredKey == "" { + return task.VideoURL + } + + switch task.StorageType { + case "s3", "gdrive": + if h.objectStorage != nil { + accessURL, err := h.objectStorage.GetAccessURL(c.Request.Context(), task.StoredKey) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", + "[GetVideoContent] task=%s get access URL error: %v, fallback to upstream", task.ID, err) + return task.VideoURL + } + return accessURL + } + return task.VideoURL + + case "local": + return "/sora/media" + task.StoredKey + + default: + return task.VideoURL + } +} + +func (h *SoraVideosHandler) CreateImage(c *gin.Context) { + apiKey, account, release, ok := h.selectAccount(c, "") + if !ok { + return + } + defer release() + + body, err := readBody(c) + if err != nil { + return + } + + var req service.CreateImageRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + if req.Model == "" { + req.Model = inferImageModel(req.Size) + } + + task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, &req, body) + if err != nil { + handleTaskCreateError(c, "CreateImage", err) + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) EditImage(c *gin.Context) { + apiKey, account, release, ok := h.selectAccount(c, "") + if !ok { + return + } + defer release() + + body, err := readBody(c) + if err != nil { + return + } + + var req service.EditImageRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Image == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "image is required") + return + } + if req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + if req.Model == "" { + req.Model = inferImageModel(req.Size) + } + + imageReq := &service.CreateImageRequest{ + Model: req.Model, + Prompt: req.Prompt, + Image: req.Image, + Size: req.Size, + ResponseFormat: req.ResponseFormat, + N: 1, + } + + task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, imageReq, body) + if err != nil { + handleTaskCreateError(c, "EditImage", err) + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +// ── Internal helpers ── + +func (h *SoraVideosHandler) getAPIKey(c *gin.Context) (*service.APIKey, bool) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok { + soraErrorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return nil, false + } + return apiKey, true +} + +func (h *SoraVideosHandler) selectAccount(c *gin.Context, model string) (*service.APIKey, *service.Account, func(), bool) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return nil, nil, nil, false + } + + selection, err := h.gatewayService.SelectAccountWithLoadAwareness( + c.Request.Context(), apiKey.GroupID, "", model, nil, "", int64(0), + ) + if err != nil { + soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "No available accounts") + return nil, nil, nil, false + } + + releaseFunc := func() {} + if selection.ReleaseFunc != nil { + releaseFunc = selection.ReleaseFunc + } + return apiKey, selection.Account, releaseFunc, true +} + +func (h *SoraVideosHandler) selectAccountByID(c *gin.Context, accountID int64) (*service.Account, error) { + return h.taskService.GetAccountByID(c.Request.Context(), accountID) +} + +func readBody(c *gin.Context) ([]byte, error) { + body, err := c.GetRawData() + if err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return nil, err + } + if len(body) == 0 { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty") + return nil, fmt.Errorf("empty body") + } + if !utf8.Valid(body) { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body must be valid UTF-8") + return nil, fmt.Errorf("invalid utf-8") + } + return body, nil +} + +func soraErrorResponse(c *gin.Context, status int, errType, message string) { + c.JSON(status, gin.H{ + "error": gin.H{ + "message": message, + "type": errType, + }, + }) +} + +// handleTaskCreateError writes the appropriate error response, transparently +// forwarding upstream HTTP status codes when available. +func handleTaskCreateError(c *gin.Context, logTag string, err error) { + logger.LegacyPrintf("handler.sora_videos", "[%s] error: %v", logTag, err) + var ue *service.SoraUpstreamError + if errors.As(err, &ue) { + c.Data(ue.StatusCode, "application/json", ue.Body) + return + } + soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create task") +} + +func inferImageModel(size string) string { + switch size { + case "540x360": + return "gpt-image-landscape" + case "360x540": + return "gpt-image-portrait" + default: + return "gpt-image" + } +} diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index f3aadcf330..08bdaea1cc 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -20,6 +20,7 @@ func ProvideAdminHandlers( openaiOAuthHandler *admin.OpenAIOAuthHandler, geminiOAuthHandler *admin.GeminiOAuthHandler, antigravityOAuthHandler *admin.AntigravityOAuthHandler, + gdriveOAuthHandler *admin.GDriveOAuthHandler, proxyHandler *admin.ProxyHandler, redeemHandler *admin.RedeemHandler, promoHandler *admin.PromoHandler, @@ -57,6 +58,7 @@ func ProvideAdminHandlers( ErrorPassthrough: errorPassthroughHandler, APIKey: apiKeyHandler, ScheduledTest: scheduledTestHandler, + GDriveOAuth: gdriveOAuthHandler, } } @@ -84,6 +86,7 @@ func ProvideHandlers( openaiGatewayHandler *OpenAIGatewayHandler, soraGatewayHandler *SoraGatewayHandler, soraClientHandler *SoraClientHandler, + soraVideosHandler *SoraVideosHandler, settingHandler *SettingHandler, totpHandler *TotpHandler, _ *service.IdempotencyCoordinator, @@ -102,6 +105,7 @@ func ProvideHandlers( OpenAIGateway: openaiGatewayHandler, SoraGateway: soraGatewayHandler, SoraClient: soraClientHandler, + SoraVideos: soraVideosHandler, Setting: settingHandler, Totp: totpHandler, } @@ -147,6 +151,7 @@ var ProviderSet = wire.NewSet( admin.NewErrorPassthroughHandler, admin.NewAdminAPIKeyHandler, admin.NewScheduledTestHandler, + admin.NewGDriveOAuthHandler, // AdminHandlers and Handlers constructors ProvideAdminHandlers, diff --git a/backend/internal/pkg/antigravity/stream_transformer.go b/backend/internal/pkg/antigravity/stream_transformer.go index deed5f922e..ee600c8b78 100644 --- a/backend/internal/pkg/antigravity/stream_transformer.go +++ b/backend/internal/pkg/antigravity/stream_transformer.go @@ -18,6 +18,9 @@ const ( BlockTypeFunction ) +// UsageMapHook is a callback that can modify usage data before it's emitted in SSE events. +type UsageMapHook func(usageMap map[string]any) + // StreamingProcessor 流式响应处理器 type StreamingProcessor struct { blockType BlockType @@ -30,6 +33,7 @@ type StreamingProcessor struct { originalModel string webSearchQueries []string groundingChunks []GeminiGroundingChunk + usageMapHook UsageMapHook // 累计 usage inputTokens int @@ -45,6 +49,25 @@ func NewStreamingProcessor(originalModel string) *StreamingProcessor { } } +// SetUsageMapHook sets an optional hook that modifies usage maps before they are emitted. +func (p *StreamingProcessor) SetUsageMapHook(fn UsageMapHook) { + p.usageMapHook = fn +} + +func usageToMap(u ClaudeUsage) map[string]any { + m := map[string]any{ + "input_tokens": u.InputTokens, + "output_tokens": u.OutputTokens, + } + if u.CacheCreationInputTokens > 0 { + m["cache_creation_input_tokens"] = u.CacheCreationInputTokens + } + if u.CacheReadInputTokens > 0 { + m["cache_read_input_tokens"] = u.CacheReadInputTokens + } + return m +} + // ProcessLine 处理 SSE 行,返回 Claude SSE 事件 func (p *StreamingProcessor) ProcessLine(line string) []byte { line = strings.TrimSpace(line) @@ -168,6 +191,13 @@ func (p *StreamingProcessor) emitMessageStart(v1Resp *V1InternalResponse) []byte responseID = "msg_" + generateRandomID() } + var usageValue any = usage + if p.usageMapHook != nil { + usageMap := usageToMap(usage) + p.usageMapHook(usageMap) + usageValue = usageMap + } + message := map[string]any{ "id": responseID, "type": "message", @@ -176,7 +206,7 @@ func (p *StreamingProcessor) emitMessageStart(v1Resp *V1InternalResponse) []byte "model": p.originalModel, "stop_reason": nil, "stop_sequence": nil, - "usage": usage, + "usage": usageValue, } event := map[string]any{ @@ -487,13 +517,20 @@ func (p *StreamingProcessor) emitFinish(finishReason string) []byte { CacheReadInputTokens: p.cacheReadTokens, } + var usageValue any = usage + if p.usageMapHook != nil { + usageMap := usageToMap(usage) + p.usageMapHook(usageMap) + usageValue = usageMap + } + deltaEvent := map[string]any{ "type": "message_delta", "delta": map[string]any{ "stop_reason": stopReason, "stop_sequence": nil, }, - "usage": usage, + "usage": usageValue, } _, _ = result.Write(p.formatSSE("message_delta", deltaEvent)) diff --git a/backend/internal/pkg/logger/logger_test.go b/backend/internal/pkg/logger/logger_test.go index 74aae0613a..06a277a490 100644 --- a/backend/internal/pkg/logger/logger_test.go +++ b/backend/internal/pkg/logger/logger_test.go @@ -10,7 +10,13 @@ import ( ) func TestInit_DualOutput(t *testing.T) { - tmpDir := t.TempDir() + // Use os.MkdirTemp instead of t.TempDir to avoid cleanup failures + // when lumberjack holds file handles on Windows. + tmpDir, err := os.MkdirTemp("", "logger-test-*") + if err != nil { + t.Fatalf("create temp dir: %v", err) + } + t.Cleanup(func() { _ = os.RemoveAll(tmpDir) }) logPath := filepath.Join(tmpDir, "logs", "sub2api.log") origStdout := os.Stdout @@ -57,7 +63,9 @@ func TestInit_DualOutput(t *testing.T) { L().Info("dual-output-info") L().Warn("dual-output-warn") - Sync() + + // Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers). + // The log data is already in the pipe buffer; closing writers is sufficient. _ = stdoutW.Close() _ = stderrW.Close() @@ -166,7 +174,9 @@ func TestInit_CallerShouldPointToCallsite(t *testing.T) { } L().Info("caller-check") - Sync() + // Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers). + os.Stdout = origStdout + os.Stderr = origStderr _ = stdoutW.Close() logBytes, _ := io.ReadAll(stdoutR) diff --git a/backend/internal/pkg/logger/stdlog_bridge_test.go b/backend/internal/pkg/logger/stdlog_bridge_test.go index 4482a2ecd3..30d25b333f 100644 --- a/backend/internal/pkg/logger/stdlog_bridge_test.go +++ b/backend/internal/pkg/logger/stdlog_bridge_test.go @@ -77,7 +77,7 @@ func TestStdLogBridgeRoutesLevels(t *testing.T) { log.Printf("service started") log.Printf("Warning: queue full") log.Printf("Forward request failed: timeout") - Sync() + // Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers). _ = stdoutW.Close() _ = stderrW.Close() @@ -139,7 +139,7 @@ func TestLegacyPrintfRoutesLevels(t *testing.T) { LegacyPrintf("service.test", "request started") LegacyPrintf("service.test", "Warning: queue full") LegacyPrintf("service.test", "forward failed: timeout") - Sync() + // Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers). _ = stdoutW.Close() _ = stderrW.Close() diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 4c7f38a82c..a45a83a37d 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -164,6 +164,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldModelRoutingEnabled, group.FieldModelRouting, group.FieldMcpXMLInject, + group.FieldSimulateClaudeMaxEnabled, group.FieldSupportedModelScopes, group.FieldAllowMessagesDispatch, group.FieldDefaultMappedModel, @@ -645,6 +646,7 @@ func groupEntityToService(g *dbent.Group) *service.Group { ModelRouting: g.ModelRouting, ModelRoutingEnabled: g.ModelRoutingEnabled, MCPXMLInject: g.McpXMLInject, + SimulateClaudeMaxEnabled: g.SimulateClaudeMaxEnabled, SupportedModelScopes: g.SupportedModelScopes, SortOrder: g.SortOrder, AllowMessagesDispatch: g.AllowMessagesDispatch, diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index 58291b6652..e40baffdc0 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -2,14 +2,46 @@ package repository import ( "context" + _ "embed" "fmt" + "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/redis/go-redis/v9" ) -const stickySessionPrefix = "sticky_session:" +const ( + stickySessionPrefix = "sticky_session:" + affinityKeyPrefix = "affinity:" + affinityRevKeyPrefix = "affinity_rev:" +) + +var ( + //go:embed lua/get_affinity.lua + getAffinityLua string + //go:embed lua/update_affinity.lua + updateAffinityLua string + //go:embed lua/get_affinity_count.lua + getAffinityCountLua string + //go:embed lua/get_affinity_clients.lua + getAffinityClientsLua string + //go:embed lua/get_affinity_clients_with_scores.lua + getAffinityClientsWithScoresLua string + //go:embed lua/clear_account_affinity.lua + clearAccountAffinityLua string + //go:embed lua/get_affinity_multi_count.lua + getAffinityMultiCountLua string + + getAffinityScript = redis.NewScript(getAffinityLua) + updateAffinityScript = redis.NewScript(updateAffinityLua) + getAffinityCountScript = redis.NewScript(getAffinityCountLua) + getAffinityClientsScript = redis.NewScript(getAffinityClientsLua) + getAffinityClientsWithScoresScript = redis.NewScript(getAffinityClientsWithScoresLua) + clearAccountAffinityScript = redis.NewScript(clearAccountAffinityLua) + getAffinityMultiCountScript = redis.NewScript(getAffinityMultiCountLua) +) type gatewayCache struct { rdb *redis.Client @@ -19,6 +51,16 @@ func NewGatewayCache(rdb *redis.Client) service.GatewayCache { return &gatewayCache{rdb: rdb} } +// ensureScriptLoaded 确保 Lua 脚本已加载到 Redis 服务器的脚本缓存中。 +// Pipeline 中的 Script.Run 只发送 EVALSHA,如果 Redis 重启过导致脚本缓存丢失, +// EVALSHA 会返回 NOSCRIPT 错误。此方法提前加载脚本以避免该问题。 +func ensureScriptLoaded(ctx context.Context, rdb *redis.Client, script *redis.Script) { + exists, err := script.Exists(ctx, rdb).Result() + if err != nil || len(exists) == 0 || !exists[0] { + _ = script.Load(ctx, rdb).Err() + } +} + // buildSessionKey 构建 session key,包含 groupID 实现分组隔离 // 格式: sticky_session:{groupID}:{sessionHash} func buildSessionKey(groupID int64, sessionHash string) string { @@ -41,13 +83,281 @@ func (c *gatewayCache) RefreshSessionTTL(ctx context.Context, groupID int64, ses } // DeleteSessionAccountID 删除粘性会话与账号的绑定关系。 -// 当检测到绑定的账号不可用(如状态错误、禁用、不可调度等)时调用, -// 以便下次请求能够重新选择可用账号。 -// -// DeleteSessionAccountID removes the sticky session binding for the given session. -// Called when the bound account becomes unavailable (e.g., error status, disabled, -// or unschedulable), allowing subsequent requests to select a new available account. func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error { key := buildSessionKey(groupID, sessionHash) return c.rdb.Del(ctx, key).Err() } + +// buildAffinityKey 构建正向亲和 key(member → accounts) +// 格式: affinity:{groupID}:{userID}/{clientID} +func buildAffinityKey(groupID int64, userID int64, clientID string) string { + return fmt.Sprintf("%s%d:%s", affinityKeyPrefix, groupID, buildAffinityMember(userID, clientID)) +} + +// buildAffinityReverseKey 构建反向亲和 key(account → members) +// 格式: affinity_rev:{groupID}:{accountID} +func buildAffinityReverseKey(groupID int64, accountID int64) string { + return fmt.Sprintf("%s%d:%d", affinityRevKeyPrefix, groupID, accountID) +} + +// buildAffinityMember 构建亲和成员标识 +// 格式: {userID}/{clientID} +func buildAffinityMember(userID int64, clientID string) string { + return fmt.Sprintf("%d/%s", userID, clientID) +} + +// parseAffinityMember 解析亲和成员标识为 userID 和 clientID +func parseAffinityMember(member string) (userID int64, clientID string) { + idx := strings.IndexByte(member, '/') + if idx < 0 { + // 兼容旧格式(纯 clientID,无 userID 前缀) + return 0, member + } + userID, _ = strconv.ParseInt(member[:idx], 10, 64) + clientID = member[idx+1:] + return userID, clientID +} + +// GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员 +func (c *gatewayCache) GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) { + key := buildAffinityKey(groupID, userID, clientID) + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + result, err := getAffinityScript.Run(ctx, c.rdb, []string{key}, expireThreshold).StringSlice() + if err != nil { + if err == redis.Nil { + return nil, nil + } + return nil, err + } + + accountIDs := make([]int64, 0, len(result)) + for _, s := range result { + id, err := strconv.ParseInt(s, 10, 64) + if err != nil { + continue + } + accountIDs = append(accountIDs, id) + } + return accountIDs, nil +} + +// UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL) +func (c *gatewayCache) UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error { + fwdKey := buildAffinityKey(groupID, userID, clientID) + revKey := buildAffinityReverseKey(groupID, accountID) + now := time.Now().Unix() + ttlSeconds := int64(ttl.Seconds()) + expireThreshold := now - ttlSeconds + + member := buildAffinityMember(userID, clientID) + return updateAffinityScript.Run(ctx, c.rdb, []string{fwdKey, revKey}, + now, ttlSeconds, accountID, expireThreshold, member, + ).Err() +} + +// GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员) +func (c *gatewayCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) { + if len(accountIDs) == 0 { + return map[int64]int64{}, nil + } + + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + ensureScriptLoaded(ctx, c.rdb, getAffinityCountScript) + + pipe := c.rdb.Pipeline() + cmds := make([]*redis.Cmd, len(accountIDs)) + for i, accID := range accountIDs { + key := buildAffinityReverseKey(groupID, accID) + cmds[i] = getAffinityCountScript.Run(ctx, pipe, []string{key}, expireThreshold) + } + _, err := pipe.Exec(ctx) + if err != nil && err != redis.Nil { + return nil, err + } + + result := make(map[int64]int64, len(accountIDs)) + for i, accID := range accountIDs { + count, _ := cmds[i].Int64() + result[accID] = count + } + return result, nil +} + +// GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重)。 +// accountGroups: map[accountID][]groupID,对每个 (groupID, accountID) 组合查询反向索引。 +// 返回值成员格式为 {userID}/{clientID}。 +func (c *gatewayCache) GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error) { + if len(accountGroups) == 0 { + return map[int64][]string{}, nil + } + + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + // 构建所有 (accountID, groupID) 组合的查询 + type queryItem struct { + accountID int64 + groupID int64 + } + var queries []queryItem + for accID, groupIDs := range accountGroups { + for _, gID := range groupIDs { + queries = append(queries, queryItem{accountID: accID, groupID: gID}) + } + } + + ensureScriptLoaded(ctx, c.rdb, getAffinityClientsScript) + + pipe := c.rdb.Pipeline() + cmds := make([]*redis.Cmd, len(queries)) + for i, q := range queries { + key := buildAffinityReverseKey(q.groupID, q.accountID) + cmds[i] = getAffinityClientsScript.Run(ctx, pipe, []string{key}, expireThreshold) + } + _, err := pipe.Exec(ctx) + if err != nil && err != redis.Nil { + return nil, err + } + + // 合并结果:同一个 accountID 跨多个 group 的成员去重 + result := make(map[int64][]string, len(accountGroups)) + seen := make(map[int64]map[string]struct{}, len(accountGroups)) + for i, q := range queries { + members, _ := cmds[i].StringSlice() + if len(members) == 0 { + continue + } + if seen[q.accountID] == nil { + seen[q.accountID] = make(map[string]struct{}) + } + for _, member := range members { + if _, exists := seen[q.accountID][member]; !exists { + seen[q.accountID][member] = struct{}{} + result[q.accountID] = append(result[q.accountID], member) + } + } + } + return result, nil +} + +// GetAccountAffinityClientsWithScores 获取单个账号跨所有分组的亲和客户端列表(含最后活跃时间戳,去重取最近)。 +func (c *gatewayCache) GetAccountAffinityClientsWithScores( + ctx context.Context, + accountID int64, + groupIDs []int64, + ttl time.Duration, +) ([]service.AffinityClient, error) { + if len(groupIDs) == 0 { + return nil, nil + } + + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + ensureScriptLoaded(ctx, c.rdb, getAffinityClientsWithScoresScript) + + pipe := c.rdb.Pipeline() + cmds := make([]*redis.Cmd, len(groupIDs)) + for i, gID := range groupIDs { + key := buildAffinityReverseKey(gID, accountID) + cmds[i] = getAffinityClientsWithScoresScript.Run(ctx, pipe, []string{key}, expireThreshold) + } + _, err := pipe.Exec(ctx) + if err != nil && err != redis.Nil { + return nil, err + } + + // 合并跨组结果,同一 member 取最近的 lastActive + type memberInfo struct { + userID int64 + clientID string + ts int64 + } + seen := make(map[string]*memberInfo) // member string → info + for _, cmd := range cmds { + vals, _ := cmd.StringSlice() + // vals 格式: [member1, score1, member2, score2, ...] + for j := 0; j+1 < len(vals); j += 2 { + member := vals[j] + ts, _ := strconv.ParseInt(vals[j+1], 10, 64) + if existing, ok := seen[member]; !ok || ts > existing.ts { + uid, cid := parseAffinityMember(member) + seen[member] = &memberInfo{userID: uid, clientID: cid, ts: ts} + } + } + } + + result := make([]service.AffinityClient, 0, len(seen)) + for _, info := range seen { + result = append(result, service.AffinityClient{ + UserID: info.userID, + ClientID: info.clientID, + LastActive: time.Unix(info.ts, 0), + }) + } + + // 按最后活跃时间降序排序 + service.SortAffinityClients(result) + + return result, nil +} + +// ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引)。 +// 对每个 groupID 执行 Lua 脚本:读取反向索引获取所有成员, +// 从每个成员的正向索引中移除该账号,然后删除反向索引。 +func (c *gatewayCache) ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error { + if len(groupIDs) == 0 { + return nil + } + + ensureScriptLoaded(ctx, c.rdb, clearAccountAffinityScript) + + pipe := c.rdb.Pipeline() + for _, gID := range groupIDs { + revKey := buildAffinityReverseKey(gID, accountID) + clearAccountAffinityScript.Run(ctx, pipe, []string{revKey}, gID, accountID) + } + _, err := pipe.Exec(ctx) + if err != nil && err != redis.Nil { + return err + } + return nil +} + +// GetAffinityMultiCount 获取账号的多维度亲和计数。 +// 返回: uniqueUsers(独立用户数), uniqueClients(独立客户端数), perUserClients(目标用户的客户端数) +func (c *gatewayCache) GetAffinityMultiCount( + ctx context.Context, + groupID int64, + accountID int64, + targetUserID int64, + ttl time.Duration, +) (users, clients, perUser int64, err error) { + key := buildAffinityReverseKey(groupID, accountID) + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + targetUserStr := "" + if targetUserID > 0 { + targetUserStr = strconv.FormatInt(targetUserID, 10) + } + + result, err := getAffinityMultiCountScript.Run(ctx, c.rdb, []string{key}, expireThreshold, targetUserStr).Int64Slice() + if err != nil { + if err == redis.Nil { + return 0, 0, 0, nil + } + return 0, 0, 0, err + } + + if len(result) < 4 { + return 0, 0, 0, nil + } + + // result: {totalMembers, uniqueUsers, uniqueClients, perUserClients} + return result[1], result[2], result[3], nil +} diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 674c655b80..8d24bc6f14 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -61,7 +61,8 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetMcpXMLInject(groupIn.MCPXMLInject). SetSoraStorageQuotaBytes(groupIn.SoraStorageQuotaBytes). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). - SetDefaultMappedModel(groupIn.DefaultMappedModel) + SetDefaultMappedModel(groupIn.DefaultMappedModel). + SetSimulateClaudeMaxEnabled(groupIn.SimulateClaudeMaxEnabled) // 设置模型路由配置 if groupIn.ModelRouting != nil { @@ -130,7 +131,8 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetMcpXMLInject(groupIn.MCPXMLInject). SetSoraStorageQuotaBytes(groupIn.SoraStorageQuotaBytes). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). - SetDefaultMappedModel(groupIn.DefaultMappedModel) + SetDefaultMappedModel(groupIn.DefaultMappedModel). + SetSimulateClaudeMaxEnabled(groupIn.SimulateClaudeMaxEnabled) // 显式处理可空字段:nil 需要 clear,非 nil 需要 set。 if groupIn.DailyLimitUSD != nil { diff --git a/backend/internal/repository/lua/clear_account_affinity.lua b/backend/internal/repository/lua/clear_account_affinity.lua new file mode 100644 index 0000000000..d2f0af0338 --- /dev/null +++ b/backend/internal/repository/lua/clear_account_affinity.lua @@ -0,0 +1,29 @@ +-- 清除单个账号在指定分组的所有亲和记录(正向+反向) +-- KEYS[1] = affinity_rev:{groupID}:{accountID} (反向索引) +-- ARGV[1] = groupID (用于构建正向 key) +-- ARGV[2] = accountID (正向索引中要移除的成员) +-- 返回: 清理的成员数量 +local rev_key = KEYS[1] +local group_id = ARGV[1] +local account_id = ARGV[2] + +-- 获取反向索引中所有成员 ({userID}/{clientID}) +local members = redis.call('ZRANGE', rev_key, 0, -1) +if #members == 0 then + return 0 +end + +-- 从每个成员的正向索引中移除该账号 +for _, member in ipairs(members) do + local fwd_key = 'affinity:' .. group_id .. ':' .. member + redis.call('ZREM', fwd_key, account_id) + -- 如果正向索引为空,删除 key + if redis.call('ZCARD', fwd_key) == 0 then + redis.call('DEL', fwd_key) + end +end + +-- 删除反向索引 +redis.call('DEL', rev_key) + +return #members diff --git a/backend/internal/repository/lua/get_affinity.lua b/backend/internal/repository/lua/get_affinity.lua new file mode 100644 index 0000000000..9db85be2a9 --- /dev/null +++ b/backend/internal/repository/lua/get_affinity.lua @@ -0,0 +1,5 @@ +-- 清理过期成员后返回亲和账号列表(按最近使用降序) +-- KEYS[1] = affinity:{groupID}:{userID}/{clientID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) +return redis.call('ZREVRANGE', KEYS[1], 0, -1) diff --git a/backend/internal/repository/lua/get_affinity_clients.lua b/backend/internal/repository/lua/get_affinity_clients.lua new file mode 100644 index 0000000000..7290b63758 --- /dev/null +++ b/backend/internal/repository/lua/get_affinity_clients.lua @@ -0,0 +1,6 @@ +-- 清理过期成员后返回反向索引的成员列表(按最近使用降序) +-- 成员格式: {userID}/{clientID} +-- KEYS[1] = affinity_rev:{groupID}:{accountID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) +return redis.call('ZREVRANGE', KEYS[1], 0, -1) diff --git a/backend/internal/repository/lua/get_affinity_clients_with_scores.lua b/backend/internal/repository/lua/get_affinity_clients_with_scores.lua new file mode 100644 index 0000000000..bff6b616ed --- /dev/null +++ b/backend/internal/repository/lua/get_affinity_clients_with_scores.lua @@ -0,0 +1,7 @@ +-- 清理过期成员后返回反向索引的成员列表及其 score(最后活跃时间戳) +-- 成员格式: {userID}/{clientID} +-- KEYS[1] = affinity_rev:{groupID}:{accountID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +-- 返回: {member1, score1, member2, score2, ...}(按最近使用降序) +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) +return redis.call('ZREVRANGEBYSCORE', KEYS[1], '+inf', '-inf', 'WITHSCORES') diff --git a/backend/internal/repository/lua/get_affinity_count.lua b/backend/internal/repository/lua/get_affinity_count.lua new file mode 100644 index 0000000000..7cf9b5cb51 --- /dev/null +++ b/backend/internal/repository/lua/get_affinity_count.lua @@ -0,0 +1,5 @@ +-- 清理过期成员后返回反向索引的成员数量 +-- KEYS[1] = affinity_rev:{groupID}:{accountID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) +return redis.call('ZCARD', KEYS[1]) diff --git a/backend/internal/repository/lua/get_affinity_multi_count.lua b/backend/internal/repository/lua/get_affinity_multi_count.lua new file mode 100644 index 0000000000..166cf4f9a3 --- /dev/null +++ b/backend/internal/repository/lua/get_affinity_multi_count.lua @@ -0,0 +1,40 @@ +-- 从反向索引解析多维度计数(用户/客户端/每用户客户端) +-- KEYS[1] = affinity_rev:{groupID}:{accountID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +-- ARGV[2] = 目标 userID(传 "" 则不计算 perUserClients) +-- 返回: {totalMembers, uniqueUsers, uniqueClients, perUserClients} +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) + +local members = redis.call('ZRANGE', KEYS[1], 0, -1) +local total = #members +if total == 0 then + return {0, 0, 0, 0} +end + +local target_user = ARGV[2] +local users = {} +local clients = {} +local user_count = 0 +local client_count = 0 +local per_user_count = 0 + +for _, member in ipairs(members) do + local sep = string.find(member, '/', 1, true) + if sep then + local uid = string.sub(member, 1, sep - 1) + local cid = string.sub(member, sep + 1) + if not users[uid] then + users[uid] = true + user_count = user_count + 1 + end + if not clients[cid] then + clients[cid] = true + client_count = client_count + 1 + end + if target_user ~= '' and uid == target_user then + per_user_count = per_user_count + 1 + end + end +end + +return {total, user_count, client_count, per_user_count} diff --git a/backend/internal/repository/lua/update_affinity.lua b/backend/internal/repository/lua/update_affinity.lua new file mode 100644 index 0000000000..6e9b38fcb8 --- /dev/null +++ b/backend/internal/repository/lua/update_affinity.lua @@ -0,0 +1,15 @@ +-- 原子双写正向+反向索引 +-- KEYS[1] = affinity:{groupID}:{userID}/{clientID} (正向: member → accounts) +-- KEYS[2] = affinity_rev:{groupID}:{accountID} (反向: account → members) +-- ARGV[1] = 当前时间戳 (score) +-- ARGV[2] = TTL 秒数 +-- ARGV[3] = accountID (正向索引的成员) +-- ARGV[4] = 过期阈值时间戳 (now - ttl) +-- ARGV[5] = {userID}/{clientID} (反向索引的成员) +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[4]) +redis.call('ZADD', KEYS[1], ARGV[1], ARGV[3]) +redis.call('EXPIRE', KEYS[1], ARGV[2]) +redis.call('ZREMRANGEBYSCORE', KEYS[2], '-inf', ARGV[4]) +redis.call('ZADD', KEYS[2], ARGV[1], ARGV[5]) +redis.call('EXPIRE', KEYS[2], ARGV[2]) +return 1 diff --git a/backend/internal/repository/sora_generation_repo.go b/backend/internal/repository/sora_generation_repo.go index aaf3cb2f54..7894cdb5e1 100644 --- a/backend/internal/repository/sora_generation_repo.go +++ b/backend/internal/repository/sora_generation_repo.go @@ -417,3 +417,22 @@ func (r *soraGenerationRepository) CountByUserAndStatus(ctx context.Context, use err := r.sql.QueryRowContext(ctx, query, args...).Scan(&count) return count, err } + +// CountByStorageType 按存储类型和状态统计生成记录数。 +func (r *soraGenerationRepository) CountByStorageType(ctx context.Context, storageType string, statuses []string) (int64, error) { + if len(statuses) == 0 { + return 0, nil + } + + placeholders := make([]string, len(statuses)) + args := []any{storageType} + for i, s := range statuses { + placeholders[i] = fmt.Sprintf("$%d", i+2) + args = append(args, s) + } + + var count int64 + query := fmt.Sprintf("SELECT COUNT(*) FROM sora_generations WHERE storage_type = $1 AND status IN (%s)", strings.Join(placeholders, ",")) + err := r.sql.QueryRowContext(ctx, query, args...).Scan(&count) + return count, err +} diff --git a/backend/internal/repository/sora_task_repo.go b/backend/internal/repository/sora_task_repo.go new file mode 100644 index 0000000000..d46d50fe3c --- /dev/null +++ b/backend/internal/repository/sora_task_repo.go @@ -0,0 +1,136 @@ +package repository + +import ( + "context" + "database/sql" + "encoding/json" + "unicode/utf8" + + "github.com/Wei-Shaw/sub2api/internal/service" +) + +type SoraTaskRepository struct { + db *sql.DB +} + +func NewSoraTaskRepository(sqlDB *sql.DB) service.SoraTaskRepository { + return &SoraTaskRepository{db: sqlDB} +} + +func (r *SoraTaskRepository) Create(ctx context.Context, task *service.SoraTask) error { + var charStr, reqStr *string + if task.CharacterInfo != nil { + b, _ := json.Marshal(task.CharacterInfo) + s := string(b) + charStr = &s + } + if len(task.RequestBody) > 0 && utf8.Valid(task.RequestBody) { + s := string(task.RequestBody) + reqStr = &s + } + + _, err := r.db.ExecContext(ctx, ` + INSERT INTO sora_tasks ( + id, account_id, api_key_id, upstream_task_id, object_type, + model, prompt, status, progress, video_url, stored_key, storage_type, + share_id, character_info, error_message, error_type, + request_body, seconds, size, created_at, completed_at + ) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21)`, + task.ID, task.AccountID, task.APIKeyID, task.UpstreamTaskID, task.ObjectType, + task.Model, task.Prompt, task.Status, task.Progress, task.VideoURL, + task.StoredKey, task.StorageType, + task.ShareID, charStr, task.ErrorMessage, task.ErrorType, + reqStr, task.Seconds, task.Size, task.CreatedAt, task.CompletedAt, + ) + return err +} + +func (r *SoraTaskRepository) GetByID(ctx context.Context, id string) (*service.SoraTask, error) { + row := r.db.QueryRowContext(ctx, `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1`, id) + return scanTask(row) +} + +func (r *SoraTaskRepository) GetByIDAndAPIKey(ctx context.Context, id string, apiKeyID int64) (*service.SoraTask, error) { + row := r.db.QueryRowContext(ctx, + `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1 AND api_key_id = $2`, + id, apiKeyID, + ) + return scanTask(row) +} + +func (r *SoraTaskRepository) Update(ctx context.Context, task *service.SoraTask) error { + var charStr *string + if task.CharacterInfo != nil { + b, _ := json.Marshal(task.CharacterInfo) + s := string(b) + charStr = &s + } + + _, err := r.db.ExecContext(ctx, ` + UPDATE sora_tasks SET + upstream_task_id = $2, status = $3, progress = $4, + video_url = $5, stored_key = $6, storage_type = $7, + share_id = $8, character_info = $9, + error_message = $10, error_type = $11, + completed_at = $12, seconds = $13, size = $14, object_type = $15 + WHERE id = $1`, + task.ID, task.UpstreamTaskID, task.Status, task.Progress, + task.VideoURL, task.StoredKey, task.StorageType, + task.ShareID, charStr, + task.ErrorMessage, task.ErrorType, + task.CompletedAt, task.Seconds, task.Size, task.ObjectType, + ) + return err +} + +func (r *SoraTaskRepository) ListPending(ctx context.Context) ([]*service.SoraTask, error) { + rows, err := r.db.QueryContext(ctx, + `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE status IN ('queued', 'in_progress') ORDER BY created_at ASC LIMIT 200`, + ) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + var tasks []*service.SoraTask + for rows.Next() { + t, err := scanTaskFromRow(rows) + if err != nil { + return nil, err + } + tasks = append(tasks, t) + } + return tasks, rows.Err() +} + +const soraTaskColumns = `id, account_id, api_key_id, upstream_task_id, object_type, + model, prompt, status, progress, video_url, stored_key, storage_type, + share_id, character_info, error_message, error_type, + request_body, seconds, size, created_at, completed_at` + +func scanTask(s scannable) (*service.SoraTask, error) { + var t service.SoraTask + var charJSON, reqBody []byte + err := s.Scan( + &t.ID, &t.AccountID, &t.APIKeyID, &t.UpstreamTaskID, &t.ObjectType, + &t.Model, &t.Prompt, &t.Status, &t.Progress, &t.VideoURL, + &t.StoredKey, &t.StorageType, + &t.ShareID, &charJSON, &t.ErrorMessage, &t.ErrorType, + &reqBody, &t.Seconds, &t.Size, &t.CreatedAt, &t.CompletedAt, + ) + if err != nil { + return nil, err + } + if len(charJSON) > 0 { + var ch service.SoraCharacter + if json.Unmarshal(charJSON, &ch) == nil && ch.Username != "" { + t.CharacterInfo = &ch + } + } + t.RequestBody = reqBody + return &t, nil +} + +func scanTaskFromRow(rows *sql.Rows) (*service.SoraTask, error) { + return scanTask(rows) +} diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index ca45460651..94875eb577 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -2956,7 +2956,7 @@ func (r *usageLogRepository) GetGroupStatsWithFilters(ctx context.Context, start query := ` SELECT COALESCE(ul.group_id, 0) as group_id, - COALESCE(g.name, '') as group_name, + COALESCE(g.name, '(无分组)') as group_name, COUNT(*) as requests, COALESCE(SUM(ul.input_tokens + ul.output_tokens + ul.cache_creation_tokens + ul.cache_read_tokens), 0) as total_tokens, COALESCE(SUM(ul.total_cost), 0) as cost, diff --git a/backend/internal/repository/user_group_rate_repo.go b/backend/internal/repository/user_group_rate_repo.go index e2471ae5b5..eca5313f6d 100644 --- a/backend/internal/repository/user_group_rate_repo.go +++ b/backend/internal/repository/user_group_rate_repo.go @@ -100,7 +100,7 @@ func (r *userGroupRateRepository) GetByGroupID(ctx context.Context, groupID int6 query := ` SELECT ugr.user_id, u.username, u.email, COALESCE(u.notes, ''), u.status, ugr.rate_multiplier FROM user_group_rate_multipliers ugr - JOIN users u ON u.id = ugr.user_id + JOIN users u ON u.id = ugr.user_id AND u.deleted_at IS NULL WHERE ugr.group_id = $1 ORDER BY ugr.user_id ` diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 4ae5c27260..66a0717b2e 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -210,10 +210,9 @@ func TestAPIContracts(t *testing.T) { "sora_video_price_per_request": null, "sora_video_price_per_request_hd": null, "claude_code_only": false, - "allow_messages_dispatch": false, - "fallback_group_id": null, - "fallback_group_id_on_invalid_request": null, - "allow_messages_dispatch": false, + "allow_messages_dispatch": false, + "fallback_group_id": null, + "fallback_group_id_on_invalid_request": null, "created_at": "2025-01-02T03:04:05Z", "updated_at": "2025-01-02T03:04:05Z" } @@ -651,8 +650,8 @@ func newContractDeps(t *testing.T) *contractDeps { authHandler := handler.NewAuthHandler(cfg, nil, userService, settingService, nil, redeemService, nil) apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService) usageHandler := handler.NewUsageHandler(usageService, apiKeyService) - adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil) - adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil, nil, nil) + adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) jwtAuth := func(c *gin.Context) { c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{ diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index c80cca54a8..6f89e10854 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -264,6 +264,8 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) { accounts.POST("/today-stats/batch", h.Admin.Account.GetBatchTodayStats) accounts.POST("/:id/clear-rate-limit", h.Admin.Account.ClearRateLimit) accounts.POST("/:id/reset-quota", h.Admin.Account.ResetQuota) + accounts.GET("/:id/affinity-clients", h.Admin.Account.GetAffinityClients) + accounts.GET("/:id/affinity-details", h.Admin.Account.GetAffinityDetails) accounts.GET("/:id/temp-unschedulable", h.Admin.Account.GetTempUnschedulable) accounts.DELETE("/:id/temp-unschedulable", h.Admin.Account.ClearTempUnschedulable) accounts.POST("/:id/schedulable", h.Admin.Account.SetSchedulable) @@ -414,7 +416,7 @@ func registerSettingsRoutes(admin *gin.RouterGroup, h *handler.Handlers) { // Beta 策略配置 adminSettings.GET("/beta-policy", h.Admin.Setting.GetBetaPolicySettings) adminSettings.PUT("/beta-policy", h.Admin.Setting.UpdateBetaPolicySettings) - // Sora S3 存储配置 + // Sora S3 存储配置(旧路由,保留兼容) adminSettings.GET("/sora-s3", h.Admin.Setting.GetSoraS3Settings) adminSettings.PUT("/sora-s3", h.Admin.Setting.UpdateSoraS3Settings) adminSettings.POST("/sora-s3/test", h.Admin.Setting.TestSoraS3Connection) @@ -423,6 +425,22 @@ func registerSettingsRoutes(admin *gin.RouterGroup, h *handler.Handlers) { adminSettings.PUT("/sora-s3/profiles/:profile_id", h.Admin.Setting.UpdateSoraS3Profile) adminSettings.DELETE("/sora-s3/profiles/:profile_id", h.Admin.Setting.DeleteSoraS3Profile) adminSettings.POST("/sora-s3/profiles/:profile_id/activate", h.Admin.Setting.SetActiveSoraS3Profile) + // Sora 统一存储配置(新路由,指向相同 handler) + adminSettings.GET("/sora-storage", h.Admin.Setting.GetSoraS3Settings) + adminSettings.PUT("/sora-storage", h.Admin.Setting.UpdateSoraS3Settings) + adminSettings.POST("/sora-storage/test", h.Admin.Setting.TestSoraS3Connection) + adminSettings.GET("/sora-storage/profiles", h.Admin.Setting.ListSoraS3Profiles) + adminSettings.POST("/sora-storage/profiles", h.Admin.Setting.CreateSoraS3Profile) + adminSettings.PUT("/sora-storage/profiles/:profile_id", h.Admin.Setting.UpdateSoraS3Profile) + adminSettings.DELETE("/sora-storage/profiles/:profile_id", h.Admin.Setting.DeleteSoraS3Profile) + adminSettings.POST("/sora-storage/profiles/:profile_id/activate", h.Admin.Setting.SetActiveSoraS3Profile) + // Google Drive OAuth + adminSettings.POST("/sora-storage/gdrive-oauth/start", h.Admin.GDriveOAuth.StartOAuth) + adminSettings.POST("/sora-storage/gdrive-oauth/callback", h.Admin.GDriveOAuth.OAuthCallback) + adminSettings.POST("/sora-storage/gdrive-test", h.Admin.GDriveOAuth.TestGDriveStorage) + // Sora 存储统计 + adminSettings.GET("/sora-storage/gdrive-quota", h.Admin.Setting.GetGDriveQuota) + adminSettings.GET("/sora-storage/video-stats", h.Admin.Setting.GetStorageVideoStats) } } diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index fe82083096..f4bd810198 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -144,6 +144,16 @@ func RegisterGatewayRoutes( { soraV1.POST("/chat/completions", h.SoraGateway.ChatCompletions) soraV1.GET("/models", h.Gateway.Models) + + // Sora Videos/Images async task API + if h.SoraVideos != nil { + soraV1.POST("/videos", h.SoraVideos.CreateVideo) + soraV1.GET("/videos/:id", h.SoraVideos.GetVideo) + soraV1.POST("/videos/:id/remix", h.SoraVideos.RemixVideo) + soraV1.GET("/videos/:id/content", h.SoraVideos.GetVideoContent) + soraV1.POST("/images/generations", h.SoraVideos.CreateImage) + soraV1.POST("/images/edits", h.SoraVideos.EditImage) + } } // Sora 媒体代理(可选 API Key 验证) diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index b6408f5f7a..976d940e21 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -1204,6 +1204,244 @@ func (a *Account) IsSessionIDMaskingEnabled() bool { return false } +// IsAffinityEnabled 检查是否启用亲和调度(统一入口) +// 仅适用于 Anthropic 账号,同时检查新字段 affinity_enabled 和旧字段 client_affinity_enabled(向后兼容) +func (a *Account) IsAffinityEnabled() bool { + if a.Platform != PlatformAnthropic { + return false + } + if a.Extra == nil { + return false + } + if v, ok := a.Extra["affinity_enabled"]; ok { + if enabled, ok := v.(bool); ok { + return enabled + } + return false + } + if v, ok := a.Extra["client_affinity_enabled"]; ok { + if enabled, ok := v.(bool); ok { + return enabled + } + } + return false +} + +// IsClientAffinityEnabled 向后兼容别名,内部调用 IsAffinityEnabled +func (a *Account) IsClientAffinityEnabled() bool { + return a.IsAffinityEnabled() +} + +// AffinityZone 表示账号的客户端亲和分区 +type AffinityZone int + +const ( + AffinityZoneGreen AffinityZone = iota // 绿区:允许绑定新客户端,优先调度 + AffinityZoneYellow // 黄区:允许绑定,仅在无绿区账号时降级调度 + AffinityZoneRed // 红区:禁止调度 +) + +// GetAffinityBase 获取亲和基础限制(绿区上限),0 表示未配置 +func (a *Account) GetAffinityBase() int { + if a.Extra == nil { + return 0 + } + if v, ok := a.Extra["affinity_base"]; ok { + return parseExtraInt(v) + } + return 0 +} + +// GetAffinityBuffer 获取亲和缓冲区大小(黄区范围) +// 返回 (value, configured): +// - (0, false): 未配置 → 无限黄区(超过 base 永远黄区,永不红区) +// - (0, true): 显式设为 0 → 无黄区,超过 base 直接红区 +// - (n, true): n > 0 → 黄区范围为 base+1 到 base+n +func (a *Account) GetAffinityBuffer() (int, bool) { + if a.Extra == nil { + return 0, false + } + v, ok := a.Extra["affinity_buffer"] + if !ok { + return 0, false + } + // 显式设为 null/nil → 视为未配置 + if v == nil { + return 0, false + } + return parseExtraInt(v), true +} + +// GetAffinityZone 根据当前绑定的客户端数量计算账号的亲和分区。 +// 未开启亲和或未配置 base 的账号永远返回绿区。 +func (a *Account) GetAffinityZone(clientCount int64) AffinityZone { + if !a.IsClientAffinityEnabled() { + return AffinityZoneGreen + } + base := a.GetAffinityBase() + if base <= 0 { + return AffinityZoneGreen + } + if clientCount <= int64(base) { + return AffinityZoneGreen + } + buffer, configured := a.GetAffinityBuffer() + if !configured { + return AffinityZoneYellow // 未配置 buffer → 无限黄区 + } + if buffer == 0 { + return AffinityZoneRed // buffer=0 → 无黄区,直接红区 + } + if clientCount <= int64(base+buffer) { + return AffinityZoneYellow + } + return AffinityZoneRed +} + +// IsAffinityAllowSwitch 检查亲和账号是否允许在无绿区时切换到其他账号 +// 默认 true(允许切换),设为 false 时即使全部红区也不切换 +func (a *Account) IsAffinityAllowSwitch() bool { + if a.Extra == nil { + return true + } + if v, ok := a.Extra["affinity_allow_switch"]; ok { + if allow, ok := v.(bool); ok { + return allow + } + } + return true +} + +// GetPinnedUsers 获取指定亲和用户列表 +func (a *Account) GetPinnedUsers() []int64 { + if a.Extra == nil { + return nil + } + v, ok := a.Extra["pinned_users"] + if !ok || v == nil { + return nil + } + arr, ok := v.([]any) + if !ok || len(arr) == 0 { + return nil + } + result := make([]int64, 0, len(arr)) + for _, item := range arr { + switch id := item.(type) { + case float64: + result = append(result, int64(id)) + case int64: + result = append(result, id) + case int: + result = append(result, int64(id)) + case json.Number: + if i, err := id.Int64(); err == nil { + result = append(result, i) + } + } + } + return result +} + +// IsPinnedUser 检查 userID 是否在指定亲和用户列表中 +func (a *Account) IsPinnedUser(userID int64) bool { + for _, id := range a.GetPinnedUsers() { + if id == userID { + return true + } + } + return false +} + +// GetAffinityUserBase 获取用户维度亲和基础限制(绿区上限),0 表示未配置 +func (a *Account) GetAffinityUserBase() int { + if a.Extra == nil { + return 0 + } + if v, ok := a.Extra["affinity_user_base"]; ok { + return parseExtraInt(v) + } + return 0 +} + +// GetAffinityUserBuffer 获取用户维度亲和缓冲区大小(黄区范围) +// 返回 (value, configured),语义与 GetAffinityBuffer 相同 +func (a *Account) GetAffinityUserBuffer() (int, bool) { + if a.Extra == nil { + return 0, false + } + v, ok := a.Extra["affinity_user_buffer"] + if !ok { + return 0, false + } + if v == nil { + return 0, false + } + return parseExtraInt(v), true +} + +// GetPerUserClientLimit 获取每用户客户端限制,0 表示不限制 +func (a *Account) GetPerUserClientLimit() int { + if a.Extra == nil { + return 0 + } + if v, ok := a.Extra["per_user_client_limit"]; ok { + return parseExtraInt(v) + } + return 0 +} + +// getAffinityZoneForDim 根据单一维度的计数和三区参数计算亲和分区 +func getAffinityZoneForDim(count int64, base int, buffer int, bufferConfigured bool) AffinityZone { + if base <= 0 { + return AffinityZoneGreen + } + if count <= int64(base) { + return AffinityZoneGreen + } + if !bufferConfigured { + return AffinityZoneYellow + } + if buffer == 0 { + return AffinityZoneRed + } + if count <= int64(base+buffer) { + return AffinityZoneYellow + } + return AffinityZoneRed +} + +// GetMultiDimAffinityZone 根据多维度计数计算亲和分区,取所有维度中最严格的区域。 +// 未开启亲和的账号永远返回绿区。 +func (a *Account) GetMultiDimAffinityZone(userCount, clientCount, perUserCount int64) AffinityZone { + if !a.IsAffinityEnabled() { + return AffinityZoneGreen + } + worst := AffinityZoneGreen + + // 客户端维度 + clientBase := a.GetAffinityBase() + clientBuffer, clientBufCfg := a.GetAffinityBuffer() + if z := getAffinityZoneForDim(clientCount, clientBase, clientBuffer, clientBufCfg); z > worst { + worst = z + } + + // 用户维度 + userBase := a.GetAffinityUserBase() + userBuffer, userBufCfg := a.GetAffinityUserBuffer() + if z := getAffinityZoneForDim(userCount, userBase, userBuffer, userBufCfg); z > worst { + worst = z + } + + // 每用户客户端维度 + perUserLimit := a.GetPerUserClientLimit() + if perUserLimit > 0 && perUserCount > int64(perUserLimit) { + worst = AffinityZoneRed + } + + return worst +} + // IsCacheTTLOverrideEnabled 检查是否启用缓存 TTL 强制替换 // 仅适用于 Anthropic OAuth/SetupToken 类型账号 // 启用后将所有 cache creation tokens 归入指定的 TTL 类型(5m 或 1h) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 5eeac18329..0c4c3072fc 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -143,9 +143,10 @@ type CreateGroupInput struct { // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) - ModelRouting map[string][]int64 - ModelRoutingEnabled bool // 是否启用模型路由 - MCPXMLInject *bool + ModelRouting map[string][]int64 + ModelRoutingEnabled bool // 是否启用模型路由 + MCPXMLInject *bool + SimulateClaudeMaxEnabled *bool // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes []string // Sora 存储配额 @@ -182,9 +183,10 @@ type UpdateGroupInput struct { // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) - ModelRouting map[string][]int64 - ModelRoutingEnabled *bool // 是否启用模型路由 - MCPXMLInject *bool + ModelRouting map[string][]int64 + ModelRoutingEnabled *bool // 是否启用模型路由 + MCPXMLInject *bool + SimulateClaudeMaxEnabled *bool // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes *[]string // Sora 存储配额 @@ -868,6 +870,13 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn if input.MCPXMLInject != nil { mcpXMLInject = *input.MCPXMLInject } + simulateClaudeMaxEnabled := false + if input.SimulateClaudeMaxEnabled != nil { + if platform != PlatformAnthropic && *input.SimulateClaudeMaxEnabled { + return nil, fmt.Errorf("simulate_claude_max_enabled only supported for anthropic groups") + } + simulateClaudeMaxEnabled = *input.SimulateClaudeMaxEnabled + } // 如果指定了复制账号的源分组,先获取账号 ID 列表 var accountIDsToCopy []int64 @@ -924,6 +933,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn FallbackGroupIDOnInvalidRequest: fallbackOnInvalidRequest, ModelRouting: input.ModelRouting, MCPXMLInject: mcpXMLInject, + SimulateClaudeMaxEnabled: simulateClaudeMaxEnabled, SupportedModelScopes: input.SupportedModelScopes, SoraStorageQuotaBytes: input.SoraStorageQuotaBytes, AllowMessagesDispatch: input.AllowMessagesDispatch, @@ -1130,6 +1140,15 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.MCPXMLInject != nil { group.MCPXMLInject = *input.MCPXMLInject } + if input.SimulateClaudeMaxEnabled != nil { + if group.Platform != PlatformAnthropic && *input.SimulateClaudeMaxEnabled { + return nil, fmt.Errorf("simulate_claude_max_enabled only supported for anthropic groups") + } + group.SimulateClaudeMaxEnabled = *input.SimulateClaudeMaxEnabled + } + if group.Platform != PlatformAnthropic { + group.SimulateClaudeMaxEnabled = false + } // 支持的模型系列(仅 antigravity 平台使用) if input.SupportedModelScopes != nil { diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index 4845d87c10..e90ec93aaf 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -43,6 +43,16 @@ func (s *accountRepoStubForBulkUpdate) BindGroups(_ context.Context, accountID i return nil } +func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID int64) ([]Account, error) { + if err, ok := s.listByGroupErr[groupID]; ok { + return nil, err + } + if rows, ok := s.listByGroupData[groupID]; ok { + return rows, nil + } + return nil, nil +} + func (s *accountRepoStubForBulkUpdate) GetByIDs(_ context.Context, ids []int64) ([]*Account, error) { s.getByIDsCalled = true s.getByIDsIDs = append([]int64{}, ids...) @@ -63,16 +73,6 @@ func (s *accountRepoStubForBulkUpdate) GetByID(_ context.Context, id int64) (*Ac return nil, errors.New("account not found") } -func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID int64) ([]Account, error) { - if err, ok := s.listByGroupErr[groupID]; ok { - return nil, err - } - if rows, ok := s.listByGroupData[groupID]; ok { - return rows, nil - } - return nil, nil -} - // TestAdminService_BulkUpdateAccounts_AllSuccessIDs 验证批量更新成功时返回 success_ids/failed_ids。 func TestAdminService_BulkUpdateAccounts_AllSuccessIDs(t *testing.T) { repo := &accountRepoStubForBulkUpdate{} diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index 536be0b583..51d9b3d2fe 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -785,3 +785,57 @@ func TestAdminService_UpdateGroup_InvalidRequestFallbackAllowsAntigravity(t *tes require.NotNil(t, repo.updated) require.Equal(t, fallbackID, *repo.updated.FallbackGroupIDOnInvalidRequest) } + +func TestAdminService_CreateGroup_SimulateClaudeMaxRequiresAnthropic(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + + enabled := true + _, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "openai-group", + Platform: PlatformOpenAI, + SimulateClaudeMaxEnabled: &enabled, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "simulate_claude_max_enabled only supported for anthropic groups") + require.Nil(t, repo.created) +} + +func TestAdminService_UpdateGroup_SimulateClaudeMaxRequiresAnthropic(t *testing.T) { + existingGroup := &Group{ + ID: 1, + Name: "openai-group", + Platform: PlatformOpenAI, + Status: StatusActive, + } + repo := &groupRepoStubForAdmin{getByID: existingGroup} + svc := &adminServiceImpl{groupRepo: repo} + + enabled := true + _, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{ + SimulateClaudeMaxEnabled: &enabled, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "simulate_claude_max_enabled only supported for anthropic groups") + require.Nil(t, repo.updated) +} + +func TestAdminService_UpdateGroup_ClearsSimulateClaudeMaxWhenPlatformChanges(t *testing.T) { + existingGroup := &Group{ + ID: 1, + Name: "anthropic-group", + Platform: PlatformAnthropic, + Status: StatusActive, + SimulateClaudeMaxEnabled: true, + } + repo := &groupRepoStubForAdmin{getByID: existingGroup} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{ + Platform: PlatformOpenAI, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.updated) + require.False(t, repo.updated.SimulateClaudeMaxEnabled) +} diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go index f321ca89c2..9fc0641669 100644 --- a/backend/internal/service/antigravity_gateway_service.go +++ b/backend/internal/service/antigravity_gateway_service.go @@ -1717,7 +1717,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context, var clientDisconnect bool if claudeReq.Stream { // 客户端要求流式,直接透传转换 - streamRes, err := s.handleClaudeStreamingResponse(c, resp, startTime, originalModel) + streamRes, err := s.handleClaudeStreamingResponse(c, resp, startTime, originalModel, account.ID) if err != nil { logger.LegacyPrintf("service.antigravity_gateway", "%s status=stream_error error=%v", prefix, err) return nil, err @@ -1727,7 +1727,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context, clientDisconnect = streamRes.clientDisconnect } else { // 客户端要求非流式,收集流式响应后转换返回 - streamRes, err := s.handleClaudeStreamToNonStreaming(c, resp, startTime, originalModel) + streamRes, err := s.handleClaudeStreamToNonStreaming(c, resp, startTime, originalModel, account.ID) if err != nil { logger.LegacyPrintf("service.antigravity_gateway", "%s status=stream_collect_error error=%v", prefix, err) return nil, err @@ -1736,6 +1736,9 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context, firstTokenMs = streamRes.firstTokenMs } + // Claude Max cache billing: 同步 ForwardResult.Usage 与客户端响应体一致 + applyClaudeMaxCacheBillingPolicyToUsage(usage, parsedRequestFromGinContext(c), claudeMaxGroupFromGinContext(c), originalModel, account.ID) + return &ForwardResult{ RequestID: requestID, Usage: *usage, @@ -3670,7 +3673,7 @@ func (s *AntigravityGatewayService) writeGoogleError(c *gin.Context, status int, // handleClaudeStreamToNonStreaming 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 -func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { +func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string, accountID int64) (*antigravityStreamResult, error) { scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { @@ -3828,6 +3831,9 @@ returnResponse: return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") } + // Claude Max cache billing simulation (non-streaming) + claudeResp = applyClaudeMaxNonStreamingRewrite(c, claudeResp, agUsage, originalModel, accountID) + c.Data(http.StatusOK, "application/json", claudeResp) // 转换为 service.ClaudeUsage @@ -3842,7 +3848,7 @@ returnResponse: } // handleClaudeStreamingResponse 处理 Claude 流式响应(Gemini SSE → Claude SSE 转换) -func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { +func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string, accountID int64) (*antigravityStreamResult, error) { c.Header("Content-Type", "text/event-stream") c.Header("Cache-Control", "no-cache") c.Header("Connection", "keep-alive") @@ -3855,6 +3861,8 @@ func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context } processor := antigravity.NewStreamingProcessor(originalModel) + setupClaudeMaxStreamingHook(c, processor, originalModel, accountID) + var firstTokenMs *int // 使用 Scanner 并限制单行大小,避免 ReadString 无上限导致 OOM scanner := bufio.NewScanner(resp.Body) diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go index 6e0a730544..b2e2fc38ac 100644 --- a/backend/internal/service/antigravity_gateway_service_test.go +++ b/backend/internal/service/antigravity_gateway_service_test.go @@ -922,7 +922,7 @@ func TestHandleClaudeStreamingResponse_NormalComplete(t *testing.T) { fmt.Fprintln(pw, "") }() - result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5") + result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0) _ = pr.Close() require.NoError(t, err) @@ -999,7 +999,7 @@ func TestHandleClaudeStreamingResponse_ThoughtsTokenCount(t *testing.T) { fmt.Fprintln(pw, "") }() - result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "gemini-2.5-pro") + result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "gemini-2.5-pro", 0) _ = pr.Close() require.NoError(t, err) @@ -1202,7 +1202,7 @@ func TestHandleClaudeStreamingResponse_ClientDisconnect(t *testing.T) { fmt.Fprintln(pw, "") }() - result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5") + result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0) _ = pr.Close() require.NoError(t, err) @@ -1234,7 +1234,7 @@ func TestHandleClaudeStreamingResponse_EmptyStream(t *testing.T) { fmt.Fprintln(pw, "") }() - _, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5") + _, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0) _ = pr.Close() // 应当返回 UpstreamFailoverError 而非 nil,以便上层触发 failover @@ -1266,7 +1266,7 @@ func TestHandleClaudeStreamingResponse_ContextCanceled(t *testing.T) { resp := &http.Response{StatusCode: http.StatusOK, Body: cancelReadCloser{}, Header: http.Header{}} - result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5") + result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0) require.NoError(t, err) require.NotNil(t, result) diff --git a/backend/internal/service/antigravity_smart_retry_test.go b/backend/internal/service/antigravity_smart_retry_test.go index 218a128808..0f1b0215fd 100644 --- a/backend/internal/service/antigravity_smart_retry_test.go +++ b/backend/internal/service/antigravity_smart_retry_test.go @@ -9,6 +9,7 @@ import ( "net/http" "strings" "testing" + "time" "github.com/stretchr/testify/require" ) @@ -29,6 +30,27 @@ func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID c.deleteCalls = append(c.deleteCalls, deleteSessionCall{groupID: groupID, sessionHash: sessionHash}) return nil } +func (c *stubSmartRetryCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (c *stubSmartRetryCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (c *stubSmartRetryCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (c *stubSmartRetryCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (c *stubSmartRetryCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (c *stubSmartRetryCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (c *stubSmartRetryCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} // mockSmartRetryUpstream 用于 handleSmartRetry 测试的 mock upstream type mockSmartRetryUpstream struct { @@ -80,17 +102,12 @@ func (m *mockSmartRetryUpstream) Do(req *http.Request, proxyURL string, accountI m.responseBodies[respIdx] = bodyBytes } - // 用缓存的 body 字节重建新的 reader - var body io.ReadCloser + // 用缓存的 body 重建 reader(支持重试场景多次读取) + cloned := *resp if m.responseBodies[respIdx] != nil { - body = io.NopCloser(bytes.NewReader(m.responseBodies[respIdx])) + cloned.Body = io.NopCloser(bytes.NewReader(m.responseBodies[respIdx])) } - - return &http.Response{ - StatusCode: resp.StatusCode, - Header: resp.Header.Clone(), - Body: body, - }, respErr + return &cloned, respErr } func (m *mockSmartRetryUpstream) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, enableTLSFingerprint bool) (*http.Response, error) { diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index e8ad5c9c32..258b842bca 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -59,9 +59,10 @@ type APIKeyAuthGroupSnapshot struct { // Model routing is used by gateway account selection, so it must be part of auth cache snapshot. // Only anthropic groups use these fields; others may leave them empty. - ModelRouting map[string][]int64 `json:"model_routing,omitempty"` - ModelRoutingEnabled bool `json:"model_routing_enabled"` - MCPXMLInject bool `json:"mcp_xml_inject"` + ModelRouting map[string][]int64 `json:"model_routing,omitempty"` + ModelRoutingEnabled bool `json:"model_routing_enabled"` + MCPXMLInject bool `json:"mcp_xml_inject"` + SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled"` // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes []string `json:"supported_model_scopes,omitempty"` diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index f727ab10f3..d874ccf2ff 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -244,6 +244,7 @@ func (s *APIKeyService) snapshotFromAPIKey(apiKey *APIKey) *APIKeyAuthSnapshot { ModelRouting: apiKey.Group.ModelRouting, ModelRoutingEnabled: apiKey.Group.ModelRoutingEnabled, MCPXMLInject: apiKey.Group.MCPXMLInject, + SimulateClaudeMaxEnabled: apiKey.Group.SimulateClaudeMaxEnabled, SupportedModelScopes: apiKey.Group.SupportedModelScopes, AllowMessagesDispatch: apiKey.Group.AllowMessagesDispatch, DefaultMappedModel: apiKey.Group.DefaultMappedModel, @@ -303,6 +304,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho ModelRouting: snapshot.Group.ModelRouting, ModelRoutingEnabled: snapshot.Group.ModelRoutingEnabled, MCPXMLInject: snapshot.Group.MCPXMLInject, + SimulateClaudeMaxEnabled: snapshot.Group.SimulateClaudeMaxEnabled, SupportedModelScopes: snapshot.Group.SupportedModelScopes, AllowMessagesDispatch: snapshot.Group.AllowMessagesDispatch, DefaultMappedModel: snapshot.Group.DefaultMappedModel, diff --git a/backend/internal/service/claude_max_cache_billing_policy.go b/backend/internal/service/claude_max_cache_billing_policy.go new file mode 100644 index 0000000000..2381915ef9 --- /dev/null +++ b/backend/internal/service/claude_max_cache_billing_policy.go @@ -0,0 +1,450 @@ +package service + +import ( + "encoding/json" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/tidwall/gjson" +) + +type claudeMaxCacheBillingOutcome struct { + Simulated bool +} + +func applyClaudeMaxCacheBillingPolicyToUsage(usage *ClaudeUsage, parsed *ParsedRequest, group *Group, model string, accountID int64) claudeMaxCacheBillingOutcome { + var out claudeMaxCacheBillingOutcome + if usage == nil || !shouldApplyClaudeMaxBillingRulesForUsage(group, model, parsed) { + return out + } + + resolvedModel := strings.TrimSpace(model) + if resolvedModel == "" && parsed != nil { + resolvedModel = strings.TrimSpace(parsed.Model) + } + + if hasCacheCreationTokens(*usage) { + // Upstream already returned cache creation usage; keep original usage. + return out + } + + if !shouldSimulateClaudeMaxUsageForUsage(*usage, parsed) { + return out + } + beforeInputTokens := usage.InputTokens + out.Simulated = safelyProjectUsageToClaudeMax1H(usage, parsed) + if out.Simulated { + logger.LegacyPrintf("service.gateway", "simulate_claude_max_usage: model=%s account=%d input_tokens:%d->%d cache_creation_1h=%d", + resolvedModel, + accountID, + beforeInputTokens, + usage.InputTokens, + usage.CacheCreation1hTokens, + ) + } + return out +} + +func isClaudeFamilyModel(model string) bool { + normalized := strings.ToLower(strings.TrimSpace(claude.NormalizeModelID(model))) + if normalized == "" { + return false + } + return strings.Contains(normalized, "claude-") +} + +func shouldApplyClaudeMaxBillingRules(input *RecordUsageInput) bool { + if input == nil || input.Result == nil || input.APIKey == nil || input.APIKey.Group == nil { + return false + } + return shouldApplyClaudeMaxBillingRulesForUsage(input.APIKey.Group, input.Result.Model, input.ParsedRequest) +} + +func shouldApplyClaudeMaxBillingRulesForUsage(group *Group, model string, parsed *ParsedRequest) bool { + if group == nil { + return false + } + if !group.SimulateClaudeMaxEnabled || group.Platform != PlatformAnthropic { + return false + } + + resolvedModel := model + if resolvedModel == "" && parsed != nil { + resolvedModel = parsed.Model + } + if !isClaudeFamilyModel(resolvedModel) { + return false + } + return true +} + +func hasCacheCreationTokens(usage ClaudeUsage) bool { + return usage.CacheCreationInputTokens > 0 || usage.CacheCreation5mTokens > 0 || usage.CacheCreation1hTokens > 0 +} + +func shouldSimulateClaudeMaxUsage(input *RecordUsageInput) bool { + if input == nil || input.Result == nil { + return false + } + if !shouldApplyClaudeMaxBillingRules(input) { + return false + } + return shouldSimulateClaudeMaxUsageForUsage(input.Result.Usage, input.ParsedRequest) +} + +func shouldSimulateClaudeMaxUsageForUsage(usage ClaudeUsage, parsed *ParsedRequest) bool { + if usage.InputTokens <= 0 { + return false + } + if hasCacheCreationTokens(usage) { + return false + } + if !hasClaudeCacheSignals(parsed) { + return false + } + return true +} + +func safelyProjectUsageToClaudeMax1H(usage *ClaudeUsage, parsed *ParsedRequest) (changed bool) { + defer func() { + if r := recover(); r != nil { + logger.LegacyPrintf("service.gateway", "simulate_claude_max_usage skipped: panic=%v", r) + changed = false + } + }() + return projectUsageToClaudeMax1H(usage, parsed) +} + +func projectUsageToClaudeMax1H(usage *ClaudeUsage, parsed *ParsedRequest) bool { + if usage == nil { + return false + } + totalWindowTokens := usage.InputTokens + usage.CacheCreation5mTokens + usage.CacheCreation1hTokens + if totalWindowTokens <= 1 { + return false + } + + simulatedInputTokens := computeClaudeMaxProjectedInputTokens(totalWindowTokens, parsed) + if simulatedInputTokens <= 0 { + simulatedInputTokens = 1 + } + if simulatedInputTokens >= totalWindowTokens { + simulatedInputTokens = totalWindowTokens - 1 + } + + cacheCreation1hTokens := totalWindowTokens - simulatedInputTokens + if usage.InputTokens == simulatedInputTokens && + usage.CacheCreation5mTokens == 0 && + usage.CacheCreation1hTokens == cacheCreation1hTokens && + usage.CacheCreationInputTokens == cacheCreation1hTokens { + return false + } + + usage.InputTokens = simulatedInputTokens + usage.CacheCreation5mTokens = 0 + usage.CacheCreation1hTokens = cacheCreation1hTokens + usage.CacheCreationInputTokens = cacheCreation1hTokens + return true +} + +type claudeCacheProjection struct { + HasBreakpoint bool + BreakpointCount int + TotalEstimatedTokens int + TailEstimatedTokens int +} + +func computeClaudeMaxProjectedInputTokens(totalWindowTokens int, parsed *ParsedRequest) int { + if totalWindowTokens <= 1 { + return totalWindowTokens + } + + projection := analyzeClaudeCacheProjection(parsed) + if !projection.HasBreakpoint || projection.TotalEstimatedTokens <= 0 || projection.TailEstimatedTokens <= 0 { + return totalWindowTokens + } + + totalEstimate := int64(projection.TotalEstimatedTokens) + tailEstimate := int64(projection.TailEstimatedTokens) + if tailEstimate > totalEstimate { + tailEstimate = totalEstimate + } + + scaled := (int64(totalWindowTokens)*tailEstimate + totalEstimate/2) / totalEstimate + if scaled <= 0 { + scaled = 1 + } + if scaled >= int64(totalWindowTokens) { + scaled = int64(totalWindowTokens - 1) + } + return int(scaled) +} + +func hasClaudeCacheSignals(parsed *ParsedRequest) bool { + if parsed == nil { + return false + } + if hasTopLevelEphemeralCacheControl(parsed) { + return true + } + return countExplicitCacheBreakpoints(parsed) > 0 +} + +func hasTopLevelEphemeralCacheControl(parsed *ParsedRequest) bool { + if parsed == nil || len(parsed.Body) == 0 { + return false + } + cacheType := strings.TrimSpace(gjson.GetBytes(parsed.Body, "cache_control.type").String()) + return strings.EqualFold(cacheType, "ephemeral") +} + +func analyzeClaudeCacheProjection(parsed *ParsedRequest) claudeCacheProjection { + var projection claudeCacheProjection + if parsed == nil { + return projection + } + + total := 0 + lastBreakpointAt := -1 + + switch system := parsed.System.(type) { + case string: + total += claudeMaxMessageOverheadTokens + estimateClaudeTextTokens(system) + case []any: + for _, raw := range system { + block, ok := raw.(map[string]any) + if !ok { + total += claudeMaxUnknownContentTokens + continue + } + total += estimateClaudeBlockTokens(block) + if hasEphemeralCacheControl(block) { + lastBreakpointAt = total + projection.BreakpointCount++ + projection.HasBreakpoint = true + } + } + } + + for _, rawMsg := range parsed.Messages { + total += claudeMaxMessageOverheadTokens + msg, ok := rawMsg.(map[string]any) + if !ok { + total += claudeMaxUnknownContentTokens + continue + } + content, exists := msg["content"] + if !exists { + continue + } + msgTokens, msgLastBreak, msgBreakCount := estimateClaudeContentTokens(content) + total += msgTokens + if msgBreakCount > 0 { + lastBreakpointAt = total - msgTokens + msgLastBreak + projection.BreakpointCount += msgBreakCount + projection.HasBreakpoint = true + } + } + + if total <= 0 { + total = 1 + } + projection.TotalEstimatedTokens = total + + if projection.HasBreakpoint && lastBreakpointAt >= 0 { + tail := total - lastBreakpointAt + if tail <= 0 { + tail = 1 + } + projection.TailEstimatedTokens = tail + return projection + } + + if hasTopLevelEphemeralCacheControl(parsed) { + tail := estimateLastUserMessageTokens(parsed) + if tail <= 0 { + tail = 1 + } + projection.HasBreakpoint = true + projection.BreakpointCount = 1 + projection.TailEstimatedTokens = tail + } + return projection +} + +func countExplicitCacheBreakpoints(parsed *ParsedRequest) int { + if parsed == nil { + return 0 + } + total := 0 + if system, ok := parsed.System.([]any); ok { + for _, raw := range system { + if block, ok := raw.(map[string]any); ok && hasEphemeralCacheControl(block) { + total++ + } + } + } + for _, rawMsg := range parsed.Messages { + msg, ok := rawMsg.(map[string]any) + if !ok { + continue + } + content, ok := msg["content"].([]any) + if !ok { + continue + } + for _, raw := range content { + if block, ok := raw.(map[string]any); ok && hasEphemeralCacheControl(block) { + total++ + } + } + } + return total +} + +func hasEphemeralCacheControl(block map[string]any) bool { + if block == nil { + return false + } + raw, ok := block["cache_control"] + if !ok || raw == nil { + return false + } + switch cc := raw.(type) { + case map[string]any: + cacheType, _ := cc["type"].(string) + return strings.EqualFold(strings.TrimSpace(cacheType), "ephemeral") + case map[string]string: + return strings.EqualFold(strings.TrimSpace(cc["type"]), "ephemeral") + default: + return false + } +} + +func estimateClaudeContentTokens(content any) (tokens int, lastBreakAt int, breakpointCount int) { + switch value := content.(type) { + case string: + return estimateClaudeTextTokens(value), -1, 0 + case []any: + total := 0 + lastBreak := -1 + breaks := 0 + for _, raw := range value { + block, ok := raw.(map[string]any) + if !ok { + total += claudeMaxUnknownContentTokens + continue + } + total += estimateClaudeBlockTokens(block) + if hasEphemeralCacheControl(block) { + lastBreak = total + breaks++ + } + } + return total, lastBreak, breaks + default: + return estimateStructuredTokens(value), -1, 0 + } +} + +func estimateClaudeBlockTokens(block map[string]any) int { + if block == nil { + return claudeMaxUnknownContentTokens + } + tokens := claudeMaxBlockOverheadTokens + blockType, _ := block["type"].(string) + switch blockType { + case "text": + if text, ok := block["text"].(string); ok { + tokens += estimateClaudeTextTokens(text) + } + case "tool_result": + if content, ok := block["content"]; ok { + nested, _, _ := estimateClaudeContentTokens(content) + tokens += nested + } + case "tool_use": + if name, ok := block["name"].(string); ok { + tokens += estimateClaudeTextTokens(name) + } + if input, ok := block["input"]; ok { + tokens += estimateStructuredTokens(input) + } + default: + if text, ok := block["text"].(string); ok { + tokens += estimateClaudeTextTokens(text) + } else if content, ok := block["content"]; ok { + nested, _, _ := estimateClaudeContentTokens(content) + tokens += nested + } + } + if tokens <= claudeMaxBlockOverheadTokens { + tokens += claudeMaxUnknownContentTokens + } + return tokens +} + +func estimateLastUserMessageTokens(parsed *ParsedRequest) int { + if parsed == nil || len(parsed.Messages) == 0 { + return 0 + } + for i := len(parsed.Messages) - 1; i >= 0; i-- { + msg, ok := parsed.Messages[i].(map[string]any) + if !ok { + continue + } + role, _ := msg["role"].(string) + if !strings.EqualFold(strings.TrimSpace(role), "user") { + continue + } + tokens, _, _ := estimateClaudeContentTokens(msg["content"]) + return claudeMaxMessageOverheadTokens + tokens + } + return 0 +} + +func estimateStructuredTokens(v any) int { + if v == nil { + return 0 + } + raw, err := json.Marshal(v) + if err != nil { + return claudeMaxUnknownContentTokens + } + return estimateClaudeTextTokens(string(raw)) +} + +func estimateClaudeTextTokens(text string) int { + if tokens, ok := estimateTokensByThirdPartyTokenizer(text); ok { + return tokens + } + return estimateClaudeTextTokensHeuristic(text) +} + +func estimateClaudeTextTokensHeuristic(text string) int { + normalized := strings.Join(strings.Fields(strings.TrimSpace(text)), " ") + if normalized == "" { + return 0 + } + asciiChars := 0 + nonASCIIChars := 0 + for _, r := range normalized { + if r <= 127 { + asciiChars++ + } else { + nonASCIIChars++ + } + } + tokens := nonASCIIChars + if asciiChars > 0 { + tokens += (asciiChars + 3) / 4 + } + if words := len(strings.Fields(normalized)); words > tokens { + tokens = words + } + if tokens <= 0 { + return 1 + } + return tokens +} diff --git a/backend/internal/service/claude_max_simulation_test.go b/backend/internal/service/claude_max_simulation_test.go new file mode 100644 index 0000000000..3d2ae2e674 --- /dev/null +++ b/backend/internal/service/claude_max_simulation_test.go @@ -0,0 +1,156 @@ +package service + +import ( + "strings" + "testing" +) + +func TestProjectUsageToClaudeMax1H_Conservation(t *testing.T) { + usage := &ClaudeUsage{ + InputTokens: 1200, + CacheCreationInputTokens: 0, + CacheCreation5mTokens: 0, + CacheCreation1hTokens: 0, + } + parsed := &ParsedRequest{ + Model: "claude-sonnet-4-5", + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": strings.Repeat("cached context ", 200), + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "summarize quickly", + }, + }, + }, + }, + } + + changed := projectUsageToClaudeMax1H(usage, parsed) + if !changed { + t.Fatalf("expected usage to be projected") + } + + total := usage.InputTokens + usage.CacheCreation5mTokens + usage.CacheCreation1hTokens + if total != 1200 { + t.Fatalf("total tokens changed: got=%d want=%d", total, 1200) + } + if usage.CacheCreation5mTokens != 0 { + t.Fatalf("cache_creation_5m should be 0, got=%d", usage.CacheCreation5mTokens) + } + if usage.InputTokens <= 0 || usage.InputTokens >= 1200 { + t.Fatalf("simulated input out of range, got=%d", usage.InputTokens) + } + if usage.InputTokens > 100 { + t.Fatalf("simulated input should stay near cache breakpoint tail, got=%d", usage.InputTokens) + } + if usage.CacheCreation1hTokens <= 0 { + t.Fatalf("cache_creation_1h should be > 0, got=%d", usage.CacheCreation1hTokens) + } + if usage.CacheCreationInputTokens != usage.CacheCreation1hTokens { + t.Fatalf("cache_creation_input_tokens mismatch: got=%d want=%d", usage.CacheCreationInputTokens, usage.CacheCreation1hTokens) + } +} + +func TestComputeClaudeMaxProjectedInputTokens_Deterministic(t *testing.T) { + parsed := &ParsedRequest{ + Model: "claude-opus-4-5", + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "build context", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "what is failing now", + }, + }, + }, + }, + } + + got1 := computeClaudeMaxProjectedInputTokens(4096, parsed) + got2 := computeClaudeMaxProjectedInputTokens(4096, parsed) + if got1 != got2 { + t.Fatalf("non-deterministic input tokens: %d != %d", got1, got2) + } +} + +func TestShouldSimulateClaudeMaxUsage(t *testing.T) { + group := &Group{ + Platform: PlatformAnthropic, + SimulateClaudeMaxEnabled: true, + } + input := &RecordUsageInput{ + Result: &ForwardResult{ + Model: "claude-sonnet-4-5", + Usage: ClaudeUsage{ + InputTokens: 3000, + CacheCreationInputTokens: 0, + CacheCreation5mTokens: 0, + CacheCreation1hTokens: 0, + }, + }, + ParsedRequest: &ParsedRequest{ + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "cached", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "tail", + }, + }, + }, + }, + }, + APIKey: &APIKey{Group: group}, + } + + if !shouldSimulateClaudeMaxUsage(input) { + t.Fatalf("expected simulate=true for claude group with cache signal") + } + + input.ParsedRequest = &ParsedRequest{ + Messages: []any{ + map[string]any{"role": "user", "content": "no cache signal"}, + }, + } + if shouldSimulateClaudeMaxUsage(input) { + t.Fatalf("expected simulate=false when request has no cache signal") + } + + input.ParsedRequest = &ParsedRequest{ + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "cached", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + }, + }, + }, + } + input.Result.Usage.CacheCreationInputTokens = 100 + if shouldSimulateClaudeMaxUsage(input) { + t.Fatalf("expected simulate=false when cache creation already exists") + } +} diff --git a/backend/internal/service/claude_tokenizer.go b/backend/internal/service/claude_tokenizer.go new file mode 100644 index 0000000000..61f5e9616e --- /dev/null +++ b/backend/internal/service/claude_tokenizer.go @@ -0,0 +1,41 @@ +package service + +import ( + "sync" + + tiktoken "github.com/pkoukk/tiktoken-go" + tiktokenloader "github.com/pkoukk/tiktoken-go-loader" +) + +var ( + claudeTokenizerOnce sync.Once + claudeTokenizer *tiktoken.Tiktoken +) + +func getClaudeTokenizer() *tiktoken.Tiktoken { + claudeTokenizerOnce.Do(func() { + // Use offline loader to avoid runtime dictionary download. + tiktoken.SetBpeLoader(tiktokenloader.NewOfflineLoader()) + // Use a high-capacity tokenizer as the default approximation for Claude payloads. + enc, err := tiktoken.GetEncoding(tiktoken.MODEL_O200K_BASE) + if err != nil { + enc, err = tiktoken.GetEncoding(tiktoken.MODEL_CL100K_BASE) + } + if err == nil { + claudeTokenizer = enc + } + }) + return claudeTokenizer +} + +func estimateTokensByThirdPartyTokenizer(text string) (int, bool) { + enc := getClaudeTokenizer() + if enc == nil { + return 0, false + } + tokens := len(enc.EncodeOrdinary(text)) + if tokens <= 0 { + return 0, false + } + return tokens, true +} diff --git a/backend/internal/service/concurrency_service.go b/backend/internal/service/concurrency_service.go index 217b83d6c8..386d5ed05d 100644 --- a/backend/internal/service/concurrency_service.go +++ b/backend/internal/service/concurrency_service.go @@ -343,8 +343,9 @@ func (s *ConcurrencyService) StartSlotCleanupWorker(accountRepo AccountRepositor }() } -// GetAccountConcurrencyBatch gets current concurrency counts for multiple accounts -// Returns a map of accountID -> current concurrency count +// GetAccountConcurrencyBatch gets current concurrency counts for multiple accounts. +// Uses a detached context with timeout to prevent HTTP request cancellation from +// causing the entire batch to fail (which would show all concurrency as 0). func (s *ConcurrencyService) GetAccountConcurrencyBatch(ctx context.Context, accountIDs []int64) (map[int64]int, error) { if len(accountIDs) == 0 { return map[int64]int{}, nil @@ -356,5 +357,11 @@ func (s *ConcurrencyService) GetAccountConcurrencyBatch(ctx context.Context, acc } return result, nil } - return s.cache.GetAccountConcurrencyBatch(ctx, accountIDs) + + // Use a detached context so that a cancelled HTTP request doesn't cause + // the Redis pipeline to fail and return all-zero concurrency counts. + redisCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + return s.cache.GetAccountConcurrencyBatch(redisCtx, accountIDs) } diff --git a/backend/internal/service/concurrency_service_test.go b/backend/internal/service/concurrency_service_test.go index 078ba0dc17..65f1c7c5f3 100644 --- a/backend/internal/service/concurrency_service_test.go +++ b/backend/internal/service/concurrency_service_test.go @@ -87,6 +87,7 @@ func (c *stubConcurrencyCacheForTest) GetAccountsLoadBatch(_ context.Context, _ func (c *stubConcurrencyCacheForTest) GetUsersLoadBatch(_ context.Context, _ []UserWithConcurrency) (map[int64]*UserLoadInfo, error) { return c.usersLoadBatch, c.usersLoadErr } + func (c *stubConcurrencyCacheForTest) CleanupExpiredAccountSlots(_ context.Context, _ int64) error { return c.cleanupErr } diff --git a/backend/internal/service/error_passthrough_runtime_test.go b/backend/internal/service/error_passthrough_runtime_test.go index 7032d15b95..2b7bbf602c 100644 --- a/backend/internal/service/error_passthrough_runtime_test.go +++ b/backend/internal/service/error_passthrough_runtime_test.go @@ -220,7 +220,7 @@ func TestApplyErrorPassthroughRule_SkipMonitoringSetsContextKey(t *testing.T) { v, exists := c.Get(OpsSkipPassthroughKey) assert.True(t, exists, "OpsSkipPassthroughKey should be set when skip_monitoring=true") boolVal, ok := v.(bool) - assert.True(t, ok, "value should be bool") + assert.True(t, ok, "value should be a bool") assert.True(t, boolVal) } diff --git a/backend/internal/service/error_policy_test.go b/backend/internal/service/error_policy_test.go index 297a954c34..4099399b7c 100644 --- a/backend/internal/service/error_policy_test.go +++ b/backend/internal/service/error_policy_test.go @@ -110,13 +110,12 @@ func TestCheckErrorPolicy(t *testing.T) { expected: ErrorPolicyTempUnscheduled, }, { - // Antigravity 401 不走升级逻辑(由 applyErrorPolicy 的 temp_unschedulable_rules 自行控制), - // second hit 仍然返回 TempUnscheduled。 - name: "temp_unschedulable_401_second_hit_antigravity_stays_temp", + // Gemini OAuth 401 second hit 会升级为 error(返回 None,交由默认错误逻辑处理)。 + name: "temp_unschedulable_401_second_hit_gemini_escalates", account: &Account{ ID: 15, Type: AccountTypeOAuth, - Platform: PlatformAntigravity, + Platform: PlatformGemini, // 非 Antigravity 平台 401 second hit 升级 TempUnschedulableReason: `{"status_code":401,"until_unix":1735689600}`, Credentials: map[string]any{ "temp_unschedulable_enabled": true, @@ -131,7 +130,29 @@ func TestCheckErrorPolicy(t *testing.T) { }, statusCode: 401, body: []byte(`unauthorized`), - expected: ErrorPolicyTempUnscheduled, + expected: ErrorPolicyNone, // Gemini 401 second hit 升级为 error + }, + { + name: "temp_unschedulable_401_antigravity_no_escalation", + account: &Account{ + ID: 16, + Type: AccountTypeOAuth, + Platform: PlatformAntigravity, // Antigravity 跳过 401 升级,由 rules 正常处理 + TempUnschedulableReason: `{"status_code":401,"until_unix":1735689600}`, + Credentials: map[string]any{ + "temp_unschedulable_enabled": true, + "temp_unschedulable_rules": []any{ + map[string]any{ + "error_code": float64(401), + "keywords": []any{"unauthorized"}, + "duration_minutes": float64(10), + }, + }, + }, + }, + statusCode: 401, + body: []byte(`unauthorized`), + expected: ErrorPolicyTempUnscheduled, // Antigravity 不升级,继续走规则匹配 }, { name: "temp_unschedulable_body_miss_returns_none", diff --git a/backend/internal/service/gateway_affinity_flow.go b/backend/internal/service/gateway_affinity_flow.go new file mode 100644 index 0000000000..cf6e7e6e26 --- /dev/null +++ b/backend/internal/service/gateway_affinity_flow.go @@ -0,0 +1,239 @@ +package service + +import "context" + +// gatewayAffinityFlow encapsulates affinity-specific scheduling steps so the +// main account selection flow can stay focused on generic scheduling. +type gatewayAffinityFlow struct { + svc *GatewayService + ctx context.Context + groupID *int64 + sessionHash string + requestedModel string + affinityClientID string + affinityUserID int64 + platform string + useMixed bool + accountByID map[int64]*Account + isExcluded func(int64) bool +} + +type affinityWaitCandidate struct { + account *Account +} + +func newGatewayAffinityFlow( + svc *GatewayService, + ctx context.Context, + groupID *int64, + sessionHash string, + requestedModel string, + affinityClientID string, + affinityUserID int64, + platform string, + useMixed bool, + accountByID map[int64]*Account, + isExcluded func(int64) bool, +) *gatewayAffinityFlow { + return &gatewayAffinityFlow{ + svc: svc, + ctx: ctx, + groupID: groupID, + sessionHash: sessionHash, + requestedModel: requestedModel, + affinityClientID: affinityClientID, + affinityUserID: affinityUserID, + platform: platform, + useMixed: useMixed, + accountByID: accountByID, + isExcluded: isExcluded, + } +} + +// shouldFilterAccountWithoutClientID excludes affinity-enabled Anthropic accounts +// when metadata.user_id does not provide a usable client_id. +func shouldFilterAccountWithoutClientID(account *Account, affinityClientID string) bool { + if account == nil || affinityClientID != "" { + return false + } + if account.Platform != PlatformAnthropic { + return false + } + return account.IsAffinityEnabled() +} + +func filterAccountsWithoutClientID(accounts []Account, affinityClientID string) []Account { + if affinityClientID != "" { + return accounts + } + filtered := make([]Account, 0, len(accounts)) + for _, acc := range accounts { + if shouldFilterAccountWithoutClientID(&acc, affinityClientID) { + continue + } + filtered = append(filtered, acc) + } + return filtered +} + +func (f *gatewayAffinityFlow) preprocessPinnedUsers(accounts []Account) { + if f.affinityUserID <= 0 || f.affinityClientID == "" || f.svc.cache == nil { + return + } + for i := range accounts { + if accounts[i].IsPinnedUser(f.affinityUserID) && accounts[i].IsAffinityEnabled() { + _ = f.svc.cache.UpdateAffinity( + f.ctx, + derefGroupID(f.groupID), + f.affinityUserID, + f.affinityClientID, + accounts[i].ID, + ClientAffinityTTL, + ) + } + } +} + +// trySelectAffinityAccount runs Layer 1.4 and returns: +// - result != nil: affinity path selected an account or wait plan +// - affinityHit == true: an effective affinity-enabled record was considered and should suppress sticky fallback +func (f *gatewayAffinityFlow) trySelectAffinityAccount() (*AccountSelectionResult, bool, error) { + if f.affinityClientID == "" || f.affinityUserID <= 0 || f.svc.cache == nil { + return nil, false, nil + } + + gid := derefGroupID(f.groupID) + affinityAccountIDs, err := f.svc.cache.GetAffinityAccounts(f.ctx, gid, f.affinityUserID, f.affinityClientID, ClientAffinityTTL) + if err != nil || len(affinityAccountIDs) == 0 { + return nil, false, nil + } + + noSwitchBlocked := false + anyAllowSwitch := false + effectiveAffinityHit := false + waitCandidates := make(map[int64]*affinityWaitCandidate) + + for _, affinityAccID := range affinityAccountIDs { + account, ok := f.accountByID[affinityAccID] + + if f.isExcluded != nil && f.isExcluded(affinityAccID) { + checkAcc := account + if !ok && f.svc.accountRepo != nil { + if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil { + checkAcc = acc + } + } + if checkAcc != nil && checkAcc.IsAffinityEnabled() { + effectiveAffinityHit = true + if !checkAcc.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } else { + anyAllowSwitch = true + } + } + continue + } + + if !ok || !f.svc.isAccountSchedulableForSelection(account) { + checkAcc := account + if !ok && f.svc.accountRepo != nil { + if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil { + checkAcc = acc + } + } + if checkAcc != nil && checkAcc.IsAffinityEnabled() { + effectiveAffinityHit = true + if !checkAcc.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } else { + anyAllowSwitch = true + } + } + continue + } + + if !account.IsAffinityEnabled() { + continue + } + effectiveAffinityHit = true + if account.IsAffinityAllowSwitch() { + anyAllowSwitch = true + } + + if !f.svc.isAccountAllowedForPlatform(account, f.platform, f.useMixed) || + (f.requestedModel != "" && !f.svc.isModelSupportedByAccountWithContext(f.ctx, account, f.requestedModel)) || + !f.svc.isAccountSchedulableForModelSelection(f.ctx, account, f.requestedModel) || + !f.svc.isAccountSchedulableForQuota(account) || + !f.svc.isAccountSchedulableForWindowCost(f.ctx, account, false) || + !f.svc.isAccountSchedulableForRPM(f.ctx, account, false) { + if !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } + continue + } + + userCount, clientCount, perUserCount, multiErr := f.svc.cache.GetAffinityMultiCount( + f.ctx, gid, affinityAccID, f.affinityUserID, ClientAffinityTTL, + ) + if multiErr == nil { + zone := account.GetMultiDimAffinityZone(userCount, clientCount, perUserCount) + if zone == AffinityZoneRed { + if !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } + continue + } + } + + result, acquireErr := f.svc.tryAcquireAccountSlot(f.ctx, affinityAccID, account.Concurrency) + if acquireErr == nil && result.Acquired { + if !f.svc.checkAndRegisterSession(f.ctx, account, f.sessionHash) { + result.ReleaseFunc() + continue + } + _ = f.svc.cache.UpdateAffinity(f.ctx, gid, f.affinityUserID, f.affinityClientID, affinityAccID, ClientAffinityTTL) + if f.sessionHash != "" { + _ = f.svc.cache.SetSessionAccountID(f.ctx, gid, f.sessionHash, affinityAccID, stickySessionTTL) + } + return &AccountSelectionResult{ + Account: account, + Acquired: true, + ReleaseFunc: result.ReleaseFunc, + }, true, nil + } + if acquireErr == nil && !result.Acquired && !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + waitCandidates[affinityAccID] = &affinityWaitCandidate{account: account} + } + } + + if noSwitchBlocked && !anyAllowSwitch && f.svc.concurrencyService != nil { + for _, waitAccID := range affinityAccountIDs { + candidate, ok := waitCandidates[waitAccID] + if !ok || candidate == nil || candidate.account == nil { + continue + } + acc := candidate.account + waitingCount, _ := f.svc.concurrencyService.GetAccountWaitingCount(f.ctx, waitAccID) + if waitingCount >= f.svc.schedulingConfig().StickySessionMaxWaiting { + continue + } + if !f.svc.checkAndRegisterSession(f.ctx, acc, f.sessionHash) { + continue + } + cfg := f.svc.schedulingConfig() + return &AccountSelectionResult{ + Account: acc, + WaitPlan: &AccountWaitPlan{ + AccountID: waitAccID, + MaxConcurrency: acc.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }, + }, true, nil + } + return nil, effectiveAffinityHit, ErrAffinityNoSwitch + } + + return nil, effectiveAffinityHit, nil +} diff --git a/backend/internal/service/gateway_affinity_scheduling_test.go b/backend/internal/service/gateway_affinity_scheduling_test.go new file mode 100644 index 0000000000..f5bf974c0a --- /dev/null +++ b/backend/internal/service/gateway_affinity_scheduling_test.go @@ -0,0 +1,1759 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "sort" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// --------------------------------------------------------------------------- +// Mock: GatewayCache for affinity tests +// --------------------------------------------------------------------------- + +// mockAffinityCache 为亲和调度测试提供可控的 GatewayCache mock。 +// 通过 getCountBatchFunc 可以自定义 GetAccountAffinityCountBatch 的行为。 +type mockAffinityCache struct { + getCountBatchFunc func(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) + getCountBatchCalls int // 记录 GetAccountAffinityCountBatch 被调用次数 +} + +func (m *mockAffinityCache) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) { + return 0, errors.New("not found") +} +func (m *mockAffinityCache) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (m *mockAffinityCache) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error { + return nil +} +func (m *mockAffinityCache) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { + return nil +} +func (m *mockAffinityCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (m *mockAffinityCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (m *mockAffinityCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (m *mockAffinityCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) { + m.getCountBatchCalls++ + if m.getCountBatchFunc != nil { + return m.getCountBatchFunc(ctx, groupID, accountIDs, ttl) + } + return map[int64]int64{}, nil +} +func (m *mockAffinityCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (m *mockAffinityCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (m *mockAffinityCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} + +// --------------------------------------------------------------------------- +// Helper: 构造启用了客户端亲和的 Anthropic 账号 +// --------------------------------------------------------------------------- + +func newAffinityAccount(id int64, priority int, affinityEnabled bool) *Account { + acc := &Account{ + ID: id, + Platform: PlatformAnthropic, + Priority: priority, + Status: StatusActive, + } + if affinityEnabled { + acc.Extra = map[string]any{"client_affinity_enabled": true} + } + return acc +} + +func newAffinityAccountWithLoad(id int64, priority int, loadRate int, affinityCount int64, lastUsedAt *time.Time) accountWithLoad { + return accountWithLoad{ + account: newAffinityAccount(id, priority, true), + loadInfo: &AccountLoadInfo{AccountID: id, LoadRate: loadRate}, + affinityCount: affinityCount, + } +} + +// =========================================================================== +// 1. filterByMinAffinityCount 测试 +// =========================================================================== + +func TestAffinityFilterByMinAffinityCount(t *testing.T) { + t.Run("empty slice returns empty", func(t *testing.T) { + result := filterByMinAffinityCount(nil) + require.Empty(t, result) + }) + + t.Run("single element returned as-is", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 5}, + } + result := filterByMinAffinityCount(accounts) + require.Len(t, result, 1) + require.Equal(t, int64(1), result[0].account.ID) + }) + + t.Run("all same affinityCount returns all", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 3}, + {account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 3}, + {account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 3}, + } + result := filterByMinAffinityCount(accounts) + require.Len(t, result, 3) + }) + + t.Run("filters to min affinityCount only", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 10}, + {account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 2}, + {account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 5}, + {account: &Account{ID: 4}, loadInfo: &AccountLoadInfo{}, affinityCount: 2}, + } + result := filterByMinAffinityCount(accounts) + require.Len(t, result, 2) + require.Equal(t, int64(2), result[0].account.ID) + require.Equal(t, int64(4), result[1].account.ID) + }) + + t.Run("zero affinityCount is smallest", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 5}, + {account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 0}, + {account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 3}, + {account: &Account{ID: 4}, loadInfo: &AccountLoadInfo{}, affinityCount: 0}, + } + result := filterByMinAffinityCount(accounts) + require.Len(t, result, 2) + require.Equal(t, int64(2), result[0].account.ID) + require.Equal(t, int64(4), result[1].account.ID) + }) + + t.Run("preserves order within same affinityCount", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 5}, loadInfo: &AccountLoadInfo{}, affinityCount: 1}, + {account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 1}, + {account: &Account{ID: 7}, loadInfo: &AccountLoadInfo{}, affinityCount: 2}, + {account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 1}, + } + result := filterByMinAffinityCount(accounts) + require.Len(t, result, 3) + // 验证保持原始顺序 + require.Equal(t, int64(5), result[0].account.ID) + require.Equal(t, int64(3), result[1].account.ID) + require.Equal(t, int64(1), result[2].account.ID) + }) +} + +// =========================================================================== +// 2. populateAffinityCounts 测试 +// =========================================================================== + +func TestAffinityPopulateAffinityCounts(t *testing.T) { + ctx := context.Background() + + t.Run("nil cache does not panic", func(t *testing.T) { + svc := &GatewayService{cache: nil} + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + } + // 不应 panic + svc.populateAffinityCounts(ctx, accounts, 0) + // affinityCount 保持零值 + require.Equal(t, int64(0), accounts[0].affinityCount) + }) + + t.Run("empty accounts returns immediately", func(t *testing.T) { + cache := &mockAffinityCache{} + svc := &GatewayService{cache: cache} + svc.populateAffinityCounts(ctx, nil, 0) + require.Equal(t, 0, cache.getCountBatchCalls, "should not call Redis for empty accounts") + }) + + t.Run("no affinity-enabled accounts skips Redis call", func(t *testing.T) { + cache := &mockAffinityCache{} + svc := &GatewayService{cache: cache} + accounts := []accountWithLoad{ + // Anthropic 但未启用亲和 + {account: newAffinityAccount(1, 1, false), loadInfo: &AccountLoadInfo{}}, + // 非 Anthropic 平台 + {account: &Account{ID: 2, Platform: PlatformOpenAI}, loadInfo: &AccountLoadInfo{}}, + } + svc.populateAffinityCounts(ctx, accounts, 0) + require.Equal(t, 0, cache.getCountBatchCalls, "should skip Redis when no affinity-enabled accounts") + }) + + t.Run("correctly populates affinityCount from Redis", func(t *testing.T) { + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, accountIDs []int64, _ time.Duration) (map[int64]int64, error) { + result := map[int64]int64{} + for _, id := range accountIDs { + switch id { + case 1: + result[1] = 5 + case 2: + result[2] = 0 + case 3: + result[3] = 12 + } + } + return result, nil + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, false), loadInfo: &AccountLoadInfo{}}, // 未启用,但仍在列表中 + {account: newAffinityAccount(3, 1, true), loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 100) + + require.Equal(t, 1, cache.getCountBatchCalls, "should call Redis exactly once") + require.Equal(t, int64(5), accounts[0].affinityCount, "account 1 should have count 5") + require.Equal(t, int64(0), accounts[1].affinityCount, "account 2 should have count 0") + require.Equal(t, int64(12), accounts[2].affinityCount, "account 3 should have count 12") + }) + + t.Run("Redis error degrades gracefully with counts at 0", func(t *testing.T) { + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return nil, errors.New("redis connection refused") + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 0) + + require.Equal(t, 1, cache.getCountBatchCalls) + require.Equal(t, int64(0), accounts[0].affinityCount, "should remain 0 on error") + require.Equal(t, int64(0), accounts[1].affinityCount, "should remain 0 on error") + }) + + t.Run("partial Redis result fills only known accounts", func(t *testing.T) { + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + // 只返回部分账号的计数 + return map[int64]int64{1: 7}, nil + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 0) + + require.Equal(t, int64(7), accounts[0].affinityCount) + require.Equal(t, int64(0), accounts[1].affinityCount, "missing account should default to 0") + }) + + t.Run("queries all account IDs regardless of affinity status", func(t *testing.T) { + // 验证:只要有至少一个 affinity-enabled 账号,就查询 ALL 账号的计数 + var queriedIDs []int64 + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, accountIDs []int64, _ time.Duration) (map[int64]int64, error) { + queriedIDs = accountIDs + return map[int64]int64{}, nil + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, false), loadInfo: &AccountLoadInfo{}}, + {account: &Account{ID: 3, Platform: PlatformOpenAI}, loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 0) + + require.Equal(t, 1, cache.getCountBatchCalls) + require.Equal(t, []int64{1, 2, 3}, queriedIDs, "should query ALL account IDs, not just affinity-enabled ones") + }) +} + +// =========================================================================== +// 3. Layer 1 排序链测试(sort.SliceStable 中 affinityCount 的正确性) +// =========================================================================== + +func TestAffinityLayer1SortChain(t *testing.T) { + now := time.Now() + earlier := now.Add(-1 * time.Hour) + muchEarlier := now.Add(-2 * time.Hour) + + // 复现 Layer 1 的排序逻辑 + sortByLayer1 := func(accounts []accountWithLoad) { + sort.SliceStable(accounts, func(i, j int) bool { + a, b := accounts[i], accounts[j] + if a.account.Priority != b.account.Priority { + return a.account.Priority < b.account.Priority + } + if a.loadInfo.LoadRate != b.loadInfo.LoadRate { + return a.loadInfo.LoadRate < b.loadInfo.LoadRate + } + if a.affinityCount != b.affinityCount { + return a.affinityCount < b.affinityCount + } + switch { + case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: + return true + case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil: + return false + case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil: + return false + default: + return a.account.LastUsedAt.Before(*b.account.LastUsedAt) + } + }) + } + + t.Run("same priority same loadRate sorts by affinityCount asc", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 10}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 2}, + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 5}, + } + sortByLayer1(accounts) + require.Equal(t, int64(2), accounts[0].account.ID, "lowest affinityCount first") + require.Equal(t, int64(3), accounts[1].account.ID) + require.Equal(t, int64(1), accounts[2].account.ID, "highest affinityCount last") + }) + + t.Run("priority takes precedence over affinityCount", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 2, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 0}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100}, + } + sortByLayer1(accounts) + require.Equal(t, int64(2), accounts[0].account.ID, "lower priority wins despite higher affinityCount") + }) + + t.Run("loadRate takes precedence over affinityCount", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 100}, + } + sortByLayer1(accounts) + require.Equal(t, int64(2), accounts[0].account.ID, "lower loadRate wins despite higher affinityCount") + }) + + t.Run("affinityCount takes precedence over LRU", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 5}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 1}, + } + sortByLayer1(accounts) + require.Equal(t, int64(2), accounts[0].account.ID, "lower affinityCount wins despite older LRU") + }) + + t.Run("same affinityCount falls through to LRU", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3}, + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3}, + } + sortByLayer1(accounts) + require.Equal(t, int64(3), accounts[0].account.ID, "LRU: oldest used first") + require.Equal(t, int64(2), accounts[1].account.ID) + require.Equal(t, int64(1), accounts[2].account.ID, "LRU: most recently used last") + }) + + t.Run("full chain: priority > loadRate > affinityCount > LRU", func(t *testing.T) { + accounts := []accountWithLoad{ + // 优先级 2 - 不管其他维度如何,排在后面 + {account: &Account{ID: 10, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0}, + // 优先级 1,负载 80% - 负载高 + {account: &Account{ID: 20, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0}, + // 优先级 1,负载 20%,亲和 5 + {account: &Account{ID: 30, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 5}, + // 优先级 1,负载 20%,亲和 1,最近使用 + {account: &Account{ID: 40, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 1}, + // 优先级 1,负载 20%,亲和 1,更早使用(应排最前) + {account: &Account{ID: 50, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 1}, + } + sortByLayer1(accounts) + // 期望排序:50 → 40 → 30 → 20 → 10 + require.Equal(t, int64(50), accounts[0].account.ID, "best: p1, lr20, ac1, LRU earlier") + require.Equal(t, int64(40), accounts[1].account.ID, "second: p1, lr20, ac1, LRU now") + require.Equal(t, int64(30), accounts[2].account.ID, "third: p1, lr20, ac5") + require.Equal(t, int64(20), accounts[3].account.ID, "fourth: p1, lr80") + require.Equal(t, int64(10), accounts[4].account.ID, "last: p2") + }) +} + +// =========================================================================== +// 4. Layer 2 分层过滤链完整性测试 +// =========================================================================== + +func TestAffinityLayer2FilterChain(t *testing.T) { + now := time.Now() + earlier := now.Add(-1 * time.Hour) + muchEarlier := now.Add(-2 * time.Hour) + + // 模拟 Layer 2 的完整过滤链:Priority → LoadRate → AffinityCount → LRU + applyLayer2 := func(accounts []accountWithLoad) *accountWithLoad { + candidates := filterByMinPriority(accounts) + candidates = filterByMinLoadRate(candidates) + candidates = filterByMinAffinityCount(candidates) + return selectByLRU(candidates, false) + } + + t.Run("priority different - affinityCount does not matter", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(2), selected.account.ID, "higher priority dimension overrides affinityCount") + }) + + t.Run("same priority same loadRate different affinityCount", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 10}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 2}, + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(2), selected.account.ID, "lowest affinityCount wins") + }) + + t.Run("same priority same loadRate same affinityCount falls through to LRU", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5}, + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(3), selected.account.ID, "LRU selects oldest") + }) + + t.Run("loadRate different overrides affinityCount", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 100}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(2), selected.account.ID, "lower loadRate wins over lower affinityCount") + }) + + t.Run("full chain integration: p → lr → ac → lru", func(t *testing.T) { + accounts := []accountWithLoad{ + // p=2 淘汰 + {account: &Account{ID: 1, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0}, + // p=1, lr=50 淘汰 + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 0}, + // p=1, lr=10, ac=8 淘汰 + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 8}, + // p=1, lr=10, ac=2, lru=now 淘汰 + {account: &Account{ID: 4, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 2}, + // p=1, lr=10, ac=2, lru=muchEarlier → 胜出 + {account: &Account{ID: 5, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 2}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(5), selected.account.ID, "full chain selects ID=5") + }) + + t.Run("empty input returns nil", func(t *testing.T) { + selected := applyLayer2(nil) + require.Nil(t, selected) + }) + + t.Run("single account always selected", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 42, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(42), selected.account.ID) + }) + + t.Run("affinityCount zero preferred among same p and lr", func(t *testing.T) { + accounts := []accountWithLoad{ + {account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 5}, + {account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0}, + {account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 3}, + } + selected := applyLayer2(accounts) + require.NotNil(t, selected) + require.Equal(t, int64(2), selected.account.ID, "zero affinityCount preferred") + }) +} + +// =========================================================================== +// 5. populateAffinityCounts + filterByMinAffinityCount 联合测试 +// =========================================================================== + +func TestAffinityPopulateAndFilterIntegration(t *testing.T) { + ctx := context.Background() + + t.Run("populate then filter selects least-loaded accounts", func(t *testing.T) { + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{ + 1: 10, + 2: 3, + 3: 3, + 4: 7, + }, nil + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(3, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(4, 1, true), loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 0) + result := filterByMinAffinityCount(accounts) + + require.Len(t, result, 2) + require.Equal(t, int64(2), result[0].account.ID) + require.Equal(t, int64(3), result[1].account.ID) + }) + + t.Run("Redis failure results in all accounts having 0 affinityCount", func(t *testing.T) { + cache := &mockAffinityCache{ + getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return nil, errors.New("timeout") + }, + } + svc := &GatewayService{cache: cache} + + accounts := []accountWithLoad{ + {account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}}, + {account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}}, + } + + svc.populateAffinityCounts(ctx, accounts, 0) + result := filterByMinAffinityCount(accounts) + + // 全部为 0,全部返回 + require.Len(t, result, 2, "all accounts should pass filter when Redis fails (all have count 0)") + }) +} + +// =========================================================================== +// 6. IsClientAffinityEnabled 边界测试 +// =========================================================================== + +func TestAffinityIsClientAffinityEnabled(t *testing.T) { + t.Run("Anthropic with enabled flag", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"client_affinity_enabled": true}, + } + assert.True(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("Anthropic with disabled flag", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"client_affinity_enabled": false}, + } + assert.False(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("Anthropic with nil Extra", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: nil, + } + assert.False(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("Anthropic without the key", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"other_key": true}, + } + assert.False(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("non-Anthropic platform always false", func(t *testing.T) { + platforms := []string{PlatformOpenAI, PlatformGemini, PlatformAntigravity} + for _, p := range platforms { + acc := &Account{ + Platform: p, + Extra: map[string]any{"client_affinity_enabled": true}, + } + assert.False(t, acc.IsClientAffinityEnabled(), "platform=%s should not support affinity", p) + } + }) + + t.Run("wrong type for enabled value", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"client_affinity_enabled": "true"}, // string 而非 bool + } + assert.False(t, acc.IsClientAffinityEnabled(), "string 'true' should not enable affinity") + }) + + t.Run("new flag false overrides legacy flag true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": false, + "client_affinity_enabled": true, + }, + } + assert.False(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("new flag true overrides legacy flag false", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": true, + "client_affinity_enabled": false, + }, + } + assert.True(t, acc.IsClientAffinityEnabled()) + }) +} + +// =========================================================================== +// GetAffinityZone 测试 +// =========================================================================== + +func TestGetAffinityZone(t *testing.T) { + makeAccount := func(enabled bool, base int, buffer any) *Account { + extra := map[string]any{"client_affinity_enabled": enabled} + if base > 0 { + extra["affinity_base"] = base + } + if buffer != nil { + extra["affinity_buffer"] = buffer + } + return &Account{ + Platform: PlatformAnthropic, + Extra: extra, + } + } + + t.Run("affinity disabled always green", func(t *testing.T) { + acc := makeAccount(false, 5, 3) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(100)) + }) + + t.Run("no base configured always green", func(t *testing.T) { + acc := makeAccount(true, 0, nil) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(100)) + }) + + t.Run("within base limit is green", func(t *testing.T) { + acc := makeAccount(true, 5, 3) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(0)) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(3)) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(5)) + }) + + t.Run("no buffer configured infinite yellow", func(t *testing.T) { + acc := makeAccount(true, 5, nil) // buffer not set → infinite yellow + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(6)) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(100)) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(9999)) + }) + + t.Run("buffer zero no yellow zone", func(t *testing.T) { + acc := makeAccount(true, 5, 0) // buffer=0 → no yellow, direct red + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(5)) + assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(6)) + assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(100)) + }) + + t.Run("within buffer is yellow", func(t *testing.T) { + acc := makeAccount(true, 5, 3) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(6)) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(7)) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(8)) // base(5)+buffer(3)=8 + }) + + t.Run("beyond buffer is red", func(t *testing.T) { + acc := makeAccount(true, 5, 3) + assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(9)) + assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(100)) + }) + + t.Run("boundary exactly at base", func(t *testing.T) { + acc := makeAccount(true, 10, 5) + assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(10)) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(11)) + }) + + t.Run("boundary exactly at base plus buffer", func(t *testing.T) { + acc := makeAccount(true, 10, 5) + assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(15)) + assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(16)) + }) +} + +// =========================================================================== +// classifyByAffinityZone 测试 +// =========================================================================== + +func TestClassifyByAffinityZone(t *testing.T) { + makeAWL := func(id int64, base int, buffer any, count int64) accountWithLoad { + extra := map[string]any{"client_affinity_enabled": true} + if base > 0 { + extra["affinity_base"] = base + } + if buffer != nil { + extra["affinity_buffer"] = buffer + } + return accountWithLoad{ + account: &Account{ID: id, Platform: PlatformAnthropic, Extra: extra}, + loadInfo: &AccountLoadInfo{AccountID: id}, + affinityCount: count, + } + } + + t.Run("empty input returns empty", func(t *testing.T) { + result := classifyByAffinityZone(nil) + require.Empty(t, result) + }) + + t.Run("no zone config returns all", func(t *testing.T) { + // 没有账号配置 affinity_base → 原样返回 + accs := []accountWithLoad{ + {account: newAffinityAccount(1, 50, true), loadInfo: &AccountLoadInfo{AccountID: 1}}, + {account: newAffinityAccount(2, 50, true), loadInfo: &AccountLoadInfo{AccountID: 2}}, + } + result := classifyByAffinityZone(accs) + require.Len(t, result, 2) + }) + + t.Run("greens preferred over yellows", func(t *testing.T) { + accs := []accountWithLoad{ + makeAWL(1, 5, 3, 3), // green (3 ≤ 5) + makeAWL(2, 5, 3, 7), // yellow (5 < 7 ≤ 8) + makeAWL(3, 5, 3, 2), // green (2 ≤ 5) + } + result := classifyByAffinityZone(accs) + require.Len(t, result, 2) + + ids := []int64{result[0].account.ID, result[1].account.ID} + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + assert.Equal(t, []int64{1, 3}, ids) + }) + + t.Run("reds excluded", func(t *testing.T) { + accs := []accountWithLoad{ + makeAWL(1, 5, 3, 10), // red (10 > 8) + makeAWL(2, 5, 3, 6), // yellow (5 < 6 ≤ 8) + makeAWL(3, 5, 3, 9), // red (9 > 8) + } + result := classifyByAffinityZone(accs) + require.Len(t, result, 1) + assert.Equal(t, int64(2), result[0].account.ID) + }) + + t.Run("all red returns empty", func(t *testing.T) { + accs := []accountWithLoad{ + makeAWL(1, 5, 0, 6), // buffer=0 → red (6 > 5) + makeAWL(2, 5, 0, 10), // buffer=0 → red (10 > 5) + } + result := classifyByAffinityZone(accs) + require.Empty(t, result) + }) + + t.Run("mixed with unconfigured accounts", func(t *testing.T) { + // 账号 1: 配置了 base=5,buffer=3 → green(3≤5) + // 账号 2: 未配置 base → 视为 green + // 账号 3: 配置了 base=5,buffer=3 → red(10>8) + accs := []accountWithLoad{ + makeAWL(1, 5, 3, 3), + {account: newAffinityAccount(2, 50, true), loadInfo: &AccountLoadInfo{AccountID: 2}, affinityCount: 20}, + makeAWL(3, 5, 3, 10), + } + result := classifyByAffinityZone(accs) + require.Len(t, result, 2) + + ids := []int64{result[0].account.ID, result[1].account.ID} + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + assert.Equal(t, []int64{1, 2}, ids) + }) + + t.Run("infinite yellow never reds", func(t *testing.T) { + // buffer 未配置 → 无限黄区,永不红区 + accs := []accountWithLoad{ + makeAWL(1, 5, nil, 100), // yellow (100 > 5, no buffer → infinite yellow) + makeAWL(2, 5, nil, 3), // green (3 ≤ 5) + } + result := classifyByAffinityZone(accs) + // green 优先 + require.Len(t, result, 1) + assert.Equal(t, int64(2), result[0].account.ID) + }) + + t.Run("only yellows when no greens", func(t *testing.T) { + accs := []accountWithLoad{ + makeAWL(1, 5, nil, 10), // yellow + makeAWL(2, 5, nil, 20), // yellow + } + result := classifyByAffinityZone(accs) + require.Len(t, result, 2) + }) +} + +// =========================================================================== +// GetMultiDimAffinityZone 测试 +// =========================================================================== + +func TestGetMultiDimAffinityZone(t *testing.T) { + makeAccount := func(opts map[string]any) *Account { + extra := map[string]any{"affinity_enabled": true} + for k, v := range opts { + extra[k] = v + } + return &Account{ + Platform: PlatformAnthropic, + Extra: extra, + } + } + + t.Run("user green + client green = green", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(2, 5, 0)) + }) + + t.Run("user green + client yellow = yellow", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(2, 12, 0)) + }) + + t.Run("user yellow + client green = yellow", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(4, 5, 0)) + }) + + t.Run("user red + client green = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 0, // no yellow, direct red + }) + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(4, 5, 0)) + }) + + t.Run("user green + client red = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 0, // no yellow, direct red + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(2, 11, 0)) + }) + + t.Run("user_base=0 only checks client dimension", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + // no affinity_user_base → 0 → user dimension always green + }) + // client in yellow range + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(100, 12, 0)) + }) + + t.Run("perUserCount exceeds limit = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "per_user_client_limit": 3, + }) + // client green, user green, but perUser > limit + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(1, 5, 4)) + }) + + t.Run("perUserCount within limit = green", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "per_user_client_limit": 3, + }) + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(1, 5, 3)) + }) + + t.Run("affinity disabled always green", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": false, + "affinity_base": 1, + "affinity_buffer": 0, + }, + } + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(100, 100, 100)) + }) +} + +// =========================================================================== +// IsAffinityAllowSwitch 测试 +// =========================================================================== + +func TestIsAffinityAllowSwitch(t *testing.T) { + t.Run("default true when no field in Extra", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_enabled": true}, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("explicit false", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_allow_switch": false}, + } + assert.False(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("Extra nil defaults to true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: nil, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("explicit true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_allow_switch": true}, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) +} + +// =========================================================================== +// GetPinnedUsers / IsPinnedUser 测试 +// =========================================================================== + +func TestGetPinnedUsers(t *testing.T) { + t.Run("normal list", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{float64(10), float64(20), float64(30)}, + }, + } + result := acc.GetPinnedUsers() + assert.Equal(t, []int64{10, 20, 30}, result) + }) + + t.Run("empty list", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{}, + }, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) + + t.Run("Extra nil", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: nil, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) + + t.Run("pinned_users not set", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{}, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) +} + +func TestIsPinnedUser(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{float64(10), float64(20), float64(30)}, + }, + } + + t.Run("hit", func(t *testing.T) { + assert.True(t, acc.IsPinnedUser(20)) + }) + + t.Run("miss", func(t *testing.T) { + assert.False(t, acc.IsPinnedUser(99)) + }) + + t.Run("nil Extra", func(t *testing.T) { + nilAcc := &Account{Platform: PlatformAnthropic, Extra: nil} + assert.False(t, nilAcc.IsPinnedUser(10)) + }) +} + +// =========================================================================== +// Enhanced mock for integration tests (supports affinity lookups + tracking) +// =========================================================================== + +type affinityIntegrationCache struct { + mockAffinityCache + getAffinityAccountsFunc func(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) + updateAffinityCalls []affinityUpdateCall + getMultiCountFunc func(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (int64, int64, int64, error) + sessionBindings map[string]int64 + getAffinityAccountsCalled bool // 标记 GetAffinityAccounts 是否已被调用,用于区分预处理/Layer 2 阶段 +} + +type affinityUpdateCall struct { + groupID int64 + userID int64 + clientID string + accountID int64 + beforeAffinityLookup bool // true 表示该调用发生在 GetAffinityAccounts 之前(预处理阶段) +} + +func (c *affinityIntegrationCache) GetSessionAccountID(_ context.Context, _ int64, hash string) (int64, error) { + if c.sessionBindings != nil { + if id, ok := c.sessionBindings[hash]; ok { + return id, nil + } + } + return 0, errors.New("not found") +} + +func (c *affinityIntegrationCache) GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) { + c.getAffinityAccountsCalled = true + if c.getAffinityAccountsFunc != nil { + return c.getAffinityAccountsFunc(ctx, groupID, userID, clientID, ttl) + } + return nil, nil +} + +func (c *affinityIntegrationCache) UpdateAffinity(_ context.Context, groupID int64, userID int64, clientID string, accountID int64, _ time.Duration) error { + c.updateAffinityCalls = append(c.updateAffinityCalls, affinityUpdateCall{ + groupID: groupID, userID: userID, clientID: clientID, accountID: accountID, + beforeAffinityLookup: !c.getAffinityAccountsCalled, + }) + return nil +} + +func (c *affinityIntegrationCache) GetAffinityMultiCount(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (int64, int64, int64, error) { + if c.getMultiCountFunc != nil { + return c.getMultiCountFunc(ctx, groupID, accountID, targetUserID, ttl) + } + return 0, 0, 0, nil +} + +// =========================================================================== +// TestAffinityPreprocessPinnedUsers — 验证 pinned_users 预处理逻辑 +// =========================================================================== + +func TestAffinityPreprocessPinnedUsers(t *testing.T) { + ctx := context.Background() + testUserID := int64(42) + testClientID := "user_" + "a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2" + "_account_test" + + t.Run("pinned user triggers UpdateAffinity", func(t *testing.T) { + cache := &affinityIntegrationCache{} + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "pinned_users": []any{float64(42)}, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{"affinity_enabled": true}, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true, 2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + + // 验证 UpdateAffinity 至少被调用了一次(pinned preprocess for account 1) + found := false + for _, call := range cache.updateAffinityCalls { + if call.accountID == 1 && call.userID == testUserID { + found = true + break + } + } + assert.True(t, found, "UpdateAffinity should be called for pinned account 1 with userID %d", testUserID) + }) + + t.Run("non-pinned user does not trigger UpdateAffinity preprocess", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil // 无亲和记录,直接进入 Layer 2 + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "pinned_users": []any{float64(99)}, // different user + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + + // 验证预处理阶段(beforeAffinityLookup=true)没有为 userID=42 在 account 1 上调用 UpdateAffinity + for _, call := range cache.updateAffinityCalls { + if call.beforeAffinityLookup && call.accountID == 1 && call.userID == testUserID { + t.Errorf("UpdateAffinity should NOT be called in preprocess for non-pinned user %d on account %d", testUserID, call.accountID) + } + } + }) +} + +// =========================================================================== +// TestAffinityNoSwitchError — 验证 ErrAffinityNoSwitch 行为 +// =========================================================================== + +func TestAffinityNoSwitchError(t *testing.T) { + ctx := context.Background() + testUserID := int64(42) + testClientID := "user_" + "a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2" + "_account_test" + futureTime := time.Now().Add(1 * time.Hour) + + t.Run("affinity hit but unschedulable + allow_switch=false returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited = unschedulable for selection + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("affinity hit but unschedulable + allow_switch=true continues to Layer 2", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited = unschedulable + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": true, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(2), result.Account.ID, "should fall through to Layer 2 and select account 2") + }) + + t.Run("no affinity records - not affected by allow_switch", func(t *testing.T) { + cache := &affinityIntegrationCache{} // returns nil from GetAffinityAccounts + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(1), result.Account.ID, "should select normally via Layer 2") + }) + + t.Run("affinity disabled account ignores allow_switch=false", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 2, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.NotEqual(t, int64(0), result.Account.ID) + }) + + t.Run("allow_switch=false + model_unsupported returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + // model_mapping 仅包含 claude-3-haiku → 请求 claude-3-5-sonnet 将被拒绝 + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "claude-3-haiku-20240307": "claude-3-haiku-20240307", + }, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("allow_switch=false + red_zone returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + getMultiCountFunc: func(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + // userCount=5, clientCount=5, perUserCount=0 + // 账号 affinity_base=1, affinity_buffer=0 → clientCount(5) > base(1) + buffer(0) → 红区 + return 5, 5, 0, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + "affinity_base": 1, + "affinity_buffer": 0, // buffer=0 → 超过 base 直接红区 + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("one_allow_switch=true overrides another allow_switch=false (一票放行)", func(t *testing.T) { + // 账号 1: allow_switch=false + rate limited(不可调度) + // 账号 3: allow_switch=true + rate limited(不可调度) + // 一票放行:账号 3 允许切换 → 不阻断 → 降级到 Layer 2 选中账号 2 + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1, 3}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + { + ID: 3, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": true, + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(2), result.Account.ID, "一票放行: should fall through to Layer 2") + }) + + t.Run("excluded affinity account with allow_switch=false returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + excluded := map[int64]struct{}{1: {}} + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excluded, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("disabled affinity record does not suppress sticky fallback", func(t *testing.T) { + sessionHash := "sticky-disabled-affinity" + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + sessionBindings: map[string]int64{sessionHash: 2}, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.Equal(t, int64(2), result.Account.ID) + }) + + t.Run("historical affinity record with disabled account still falls back to sticky account", func(t *testing.T) { + sessionHash := "sticky-disabled-status-affinity" + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + sessionBindings: map[string]int64{sessionHash: 2}, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusDisabled, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.Equal(t, int64(2), result.Account.ID) + }) +} + +// =========================================================================== +// TestAffinityMultiDimClassify — 验证 classifyByAffinityZone 与多维度计数 +// =========================================================================== + +func TestAffinityMultiDimClassify(t *testing.T) { + t.Run("multi-dim zone classification with user and client dimensions", func(t *testing.T) { + makeAWL := func(id int64, extra map[string]any, count int64) accountWithLoad { + return accountWithLoad{ + account: &Account{ID: id, Platform: PlatformAnthropic, Extra: extra}, + loadInfo: &AccountLoadInfo{AccountID: id}, + affinityCount: count, + } + } + + // Account 1: client green (3 <= 5), classified as green by old classifyByAffinityZone + // Account 2: client red (10 > 8), classified as red + accs := []accountWithLoad{ + makeAWL(1, map[string]any{ + "affinity_enabled": true, + "affinity_base": 5, + "affinity_buffer": 3, + }, 3), + makeAWL(2, map[string]any{ + "affinity_enabled": true, + "affinity_base": 5, + "affinity_buffer": 3, + }, 10), + } + + result := classifyByAffinityZone(accs) + require.Len(t, result, 1) + assert.Equal(t, int64(1), result[0].account.ID, "only green account should remain") + }) + + t.Run("all green returns all", func(t *testing.T) { + makeAWL := func(id int64, count int64) accountWithLoad { + return accountWithLoad{ + account: &Account{ + ID: id, + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_base": 10, + "affinity_buffer": 5, + }, + }, + loadInfo: &AccountLoadInfo{AccountID: id}, + affinityCount: count, + } + } + + accs := []accountWithLoad{ + makeAWL(1, 3), + makeAWL(2, 5), + makeAWL(3, 10), // exactly at base + } + + result := classifyByAffinityZone(accs) + require.Len(t, result, 3) + }) +} diff --git a/backend/internal/service/gateway_claude_max_response_helpers.go b/backend/internal/service/gateway_claude_max_response_helpers.go new file mode 100644 index 0000000000..a5f5f3d2d2 --- /dev/null +++ b/backend/internal/service/gateway_claude_max_response_helpers.go @@ -0,0 +1,196 @@ +package service + +import ( + "context" + "encoding/json" + + "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" + "github.com/gin-gonic/gin" + "github.com/tidwall/sjson" +) + +type claudeMaxResponseRewriteContext struct { + Parsed *ParsedRequest + Group *Group +} + +type claudeMaxResponseRewriteContextKeyType struct{} + +var claudeMaxResponseRewriteContextKey = claudeMaxResponseRewriteContextKeyType{} + +func withClaudeMaxResponseRewriteContext(ctx context.Context, c *gin.Context, parsed *ParsedRequest) context.Context { + if ctx == nil { + ctx = context.Background() + } + value := claudeMaxResponseRewriteContext{ + Parsed: parsed, + Group: claudeMaxGroupFromGinContext(c), + } + return context.WithValue(ctx, claudeMaxResponseRewriteContextKey, value) +} + +func claudeMaxResponseRewriteContextFromContext(ctx context.Context) claudeMaxResponseRewriteContext { + if ctx == nil { + return claudeMaxResponseRewriteContext{} + } + value, _ := ctx.Value(claudeMaxResponseRewriteContextKey).(claudeMaxResponseRewriteContext) + return value +} + +func claudeMaxGroupFromGinContext(c *gin.Context) *Group { + if c == nil { + return nil + } + raw, exists := c.Get("api_key") + if !exists { + return nil + } + apiKey, ok := raw.(*APIKey) + if !ok || apiKey == nil { + return nil + } + return apiKey.Group +} + +func parsedRequestFromGinContext(c *gin.Context) *ParsedRequest { + if c == nil { + return nil + } + raw, exists := c.Get("parsed_request") + if !exists { + return nil + } + parsed, _ := raw.(*ParsedRequest) + return parsed +} + +func applyClaudeMaxSimulationToUsage(ctx context.Context, usage *ClaudeUsage, model string, accountID int64) claudeMaxCacheBillingOutcome { + var out claudeMaxCacheBillingOutcome + if usage == nil { + return out + } + rewriteCtx := claudeMaxResponseRewriteContextFromContext(ctx) + return applyClaudeMaxCacheBillingPolicyToUsage(usage, rewriteCtx.Parsed, rewriteCtx.Group, model, accountID) +} + +func applyClaudeMaxSimulationToUsageJSONMap(ctx context.Context, usageObj map[string]any, model string, accountID int64) claudeMaxCacheBillingOutcome { + var out claudeMaxCacheBillingOutcome + if usageObj == nil { + return out + } + usage := claudeUsageFromJSONMap(usageObj) + out = applyClaudeMaxSimulationToUsage(ctx, &usage, model, accountID) + if out.Simulated { + rewriteClaudeUsageJSONMap(usageObj, usage) + } + return out +} + +func rewriteClaudeUsageJSONBytes(body []byte, usage ClaudeUsage) []byte { + updated := body + var err error + + updated, err = sjson.SetBytes(updated, "usage.input_tokens", usage.InputTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.cache_creation_input_tokens", usage.CacheCreationInputTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.cache_creation.ephemeral_5m_input_tokens", usage.CacheCreation5mTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.cache_creation.ephemeral_1h_input_tokens", usage.CacheCreation1hTokens) + if err != nil { + return body + } + return updated +} + +func claudeUsageFromJSONMap(usageObj map[string]any) ClaudeUsage { + var usage ClaudeUsage + if usageObj == nil { + return usage + } + + usage.InputTokens = usageIntFromAny(usageObj["input_tokens"]) + usage.OutputTokens = usageIntFromAny(usageObj["output_tokens"]) + usage.CacheCreationInputTokens = usageIntFromAny(usageObj["cache_creation_input_tokens"]) + usage.CacheReadInputTokens = usageIntFromAny(usageObj["cache_read_input_tokens"]) + + if ccObj, ok := usageObj["cache_creation"].(map[string]any); ok { + usage.CacheCreation5mTokens = usageIntFromAny(ccObj["ephemeral_5m_input_tokens"]) + usage.CacheCreation1hTokens = usageIntFromAny(ccObj["ephemeral_1h_input_tokens"]) + } + return usage +} + +func rewriteClaudeUsageJSONMap(usageObj map[string]any, usage ClaudeUsage) { + if usageObj == nil { + return + } + usageObj["input_tokens"] = usage.InputTokens + usageObj["cache_creation_input_tokens"] = usage.CacheCreationInputTokens + + ccObj, _ := usageObj["cache_creation"].(map[string]any) + if ccObj == nil { + ccObj = make(map[string]any, 2) + usageObj["cache_creation"] = ccObj + } + ccObj["ephemeral_5m_input_tokens"] = usage.CacheCreation5mTokens + ccObj["ephemeral_1h_input_tokens"] = usage.CacheCreation1hTokens +} + +func usageIntFromAny(v any) int { + switch value := v.(type) { + case int: + return value + case int64: + return int(value) + case float64: + return int(value) + case json.Number: + if n, err := value.Int64(); err == nil { + return int(n) + } + } + return 0 +} + +// setupClaudeMaxStreamingHook 为 Antigravity 流式路径设置 SSE usage 改写 hook。 +func setupClaudeMaxStreamingHook(c *gin.Context, processor *antigravity.StreamingProcessor, originalModel string, accountID int64) { + group := claudeMaxGroupFromGinContext(c) + parsed := parsedRequestFromGinContext(c) + if !shouldApplyClaudeMaxBillingRulesForUsage(group, originalModel, parsed) { + return + } + processor.SetUsageMapHook(func(usageMap map[string]any) { + svcUsage := claudeUsageFromJSONMap(usageMap) + outcome := applyClaudeMaxCacheBillingPolicyToUsage(&svcUsage, parsed, group, originalModel, accountID) + if outcome.Simulated { + rewriteClaudeUsageJSONMap(usageMap, svcUsage) + } + }) +} + +// applyClaudeMaxNonStreamingRewrite 为 Antigravity 非流式路径改写响应体中的 usage。 +func applyClaudeMaxNonStreamingRewrite(c *gin.Context, claudeResp []byte, agUsage *antigravity.ClaudeUsage, originalModel string, accountID int64) []byte { + group := claudeMaxGroupFromGinContext(c) + parsed := parsedRequestFromGinContext(c) + if !shouldApplyClaudeMaxBillingRulesForUsage(group, originalModel, parsed) { + return claudeResp + } + svcUsage := &ClaudeUsage{ + InputTokens: agUsage.InputTokens, + OutputTokens: agUsage.OutputTokens, + CacheCreationInputTokens: agUsage.CacheCreationInputTokens, + CacheReadInputTokens: agUsage.CacheReadInputTokens, + } + outcome := applyClaudeMaxCacheBillingPolicyToUsage(svcUsage, parsed, group, originalModel, accountID) + if outcome.Simulated { + return rewriteClaudeUsageJSONBytes(claudeResp, *svcUsage) + } + return claudeResp +} diff --git a/backend/internal/service/gateway_hotpath_optimization_test.go b/backend/internal/service/gateway_hotpath_optimization_test.go index 161c4ba4b1..af108c6fc2 100644 --- a/backend/internal/service/gateway_hotpath_optimization_test.go +++ b/backend/internal/service/gateway_hotpath_optimization_test.go @@ -143,6 +143,27 @@ func (s *stickyGatewayCacheHotpathStub) RefreshSessionTTL(ctx context.Context, g func (s *stickyGatewayCacheHotpathStub) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error { return nil } +func (s *stickyGatewayCacheHotpathStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (s *stickyGatewayCacheHotpathStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (s *stickyGatewayCacheHotpathStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (s *stickyGatewayCacheHotpathStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} func (s *modelsListAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]Account, error) { s.listByGroupCalls.Add(1) @@ -732,7 +753,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { modelsListCacheTTL: time.Minute, } - result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -754,7 +775,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, account.ID) ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(0)) - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -776,7 +797,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, int64(999)) ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(77)) - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index 718cd42ade..c258d1a5fb 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -235,6 +235,28 @@ func (m *mockGatewayCacheForPlatform) DeleteSessionAccountID(ctx context.Context return nil } +func (m *mockGatewayCacheForPlatform) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (m *mockGatewayCacheForPlatform) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (m *mockGatewayCacheForPlatform) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (m *mockGatewayCacheForPlatform) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (m *mockGatewayCacheForPlatform) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (m *mockGatewayCacheForPlatform) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (m *mockGatewayCacheForPlatform) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} + type mockGroupRepoForGateway struct { groups map[int64]*Group getByIDCalls int @@ -2031,7 +2053,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, // No concurrency service } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2084,7 +2106,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, // legacy path } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2116,7 +2138,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2148,13 +2170,314 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { } excludedIDs := map[int64]struct{}{1: {}} - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) require.Equal(t, int64(2), result.Account.ID, "不应选择被排除的账号") }) + t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-load-aware路径", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(2), result.Account.ID, "无客户端ID时应过滤亲和OAuth账号") + }) + + t.Run("无客户端ID且亲和关闭时-load-aware路径恢复旧逻辑", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "关闭亲和后不应再被无clientID过滤") + }) + + t.Run("有客户端ID时不过滤Anthropic OAuth亲和账号", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + metadataUserID := "user_0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef_account_test_session_test" + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, metadataUserID, 123) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "有客户端ID时应允许亲和OAuth账号参与调度") + }) + + t.Run("过滤对Anthropic全类型生效-SetupToken也受影响", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeSetupToken, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 3, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 3, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(3), result.Account.ID, "无客户端ID时,Anthropic SetupToken/OAuth(开启亲和)都应被过滤") + }) + + t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-legacy路径", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = false + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: nil, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(2), result.Account.ID, "legacy路径也应过滤亲和OAuth账号") + }) + + t.Run("无客户端ID且亲和关闭时-legacy路径恢复旧逻辑", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": false, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = false + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: nil, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "关闭亲和后 legacy 路径不应再被无clientID过滤") + }) + t.Run("粘性命中-不调用GetByID", func(t *testing.T) { repo := &mockAccountRepoForPlatform{ accounts: []Account{ @@ -2182,7 +2505,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2218,7 +2541,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2259,7 +2582,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2287,7 +2610,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.Error(t, err) require.Nil(t, result) require.ErrorIs(t, err, ErrNoAvailableAccounts) @@ -2319,7 +2642,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2352,7 +2675,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2390,7 +2713,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2426,7 +2749,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2485,7 +2808,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2539,7 +2862,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2593,7 +2916,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2651,7 +2974,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2709,7 +3032,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2767,7 +3090,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2804,7 +3127,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2856,7 +3179,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2934,7 +3257,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { } excluded := map[int64]struct{}{1: {}} - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2988,7 +3311,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -3021,7 +3344,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.Error(t, err) require.Nil(t, result) require.ErrorIs(t, err, ErrClaudeCodeOnly) @@ -3059,7 +3382,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -3097,7 +3420,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) diff --git a/backend/internal/service/gateway_record_usage_claude_max_test.go b/backend/internal/service/gateway_record_usage_claude_max_test.go new file mode 100644 index 0000000000..3cd8693848 --- /dev/null +++ b/backend/internal/service/gateway_record_usage_claude_max_test.go @@ -0,0 +1,199 @@ +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type usageLogRepoRecordUsageStub struct { + UsageLogRepository + + last *UsageLog + inserted bool + err error +} + +func (s *usageLogRepoRecordUsageStub) Create(_ context.Context, log *UsageLog) (bool, error) { + copied := *log + s.last = &copied + return s.inserted, s.err +} + +func newGatewayServiceForRecordUsageTest(repo UsageLogRepository) *GatewayService { + return &GatewayService{ + usageLogRepo: repo, + billingService: NewBillingService(&config.Config{}, nil), + cfg: &config.Config{RunMode: config.RunModeSimple}, + deferredService: &DeferredService{}, + } +} + +func TestRecordUsage_SimulateClaudeMaxEnabled_ProjectsUsageAndSkipsTTLOverride(t *testing.T) { + repo := &usageLogRepoRecordUsageStub{inserted: true} + svc := newGatewayServiceForRecordUsageTest(repo) + + groupID := int64(11) + input := &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "req-sim-1", + Model: "claude-sonnet-4", + Duration: time.Second, + Usage: ClaudeUsage{ + InputTokens: 160, + }, + }, + ParsedRequest: &ParsedRequest{ + Model: "claude-sonnet-4", + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "long cached context for prior turns", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "please summarize the logs and provide root cause analysis", + }, + }, + }, + }, + }, + APIKey: &APIKey{ + ID: 1, + GroupID: &groupID, + Group: &Group{ + ID: groupID, + Platform: PlatformAnthropic, + RateMultiplier: 1, + SimulateClaudeMaxEnabled: true, + }, + }, + User: &User{ID: 2}, + Account: &Account{ + ID: 3, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "cache_ttl_override_enabled": true, + "cache_ttl_override_target": "5m", + }, + }, + } + + err := svc.RecordUsage(context.Background(), input) + require.NoError(t, err) + require.NotNil(t, repo.last) + + log := repo.last + require.Equal(t, 80, log.InputTokens) + require.Equal(t, 80, log.CacheCreationTokens) + require.Equal(t, 0, log.CacheCreation5mTokens) + require.Equal(t, 80, log.CacheCreation1hTokens) + require.False(t, log.CacheTTLOverridden, "simulate outcome should skip account ttl override") +} + +func TestRecordUsage_SimulateClaudeMaxDisabled_AppliesTTLOverride(t *testing.T) { + repo := &usageLogRepoRecordUsageStub{inserted: true} + svc := newGatewayServiceForRecordUsageTest(repo) + + groupID := int64(12) + input := &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "req-sim-2", + Model: "claude-sonnet-4", + Duration: time.Second, + Usage: ClaudeUsage{ + InputTokens: 40, + CacheCreationInputTokens: 120, + CacheCreation1hTokens: 120, + }, + }, + APIKey: &APIKey{ + ID: 2, + GroupID: &groupID, + Group: &Group{ + ID: groupID, + Platform: PlatformAnthropic, + RateMultiplier: 1, + SimulateClaudeMaxEnabled: false, + }, + }, + User: &User{ID: 3}, + Account: &Account{ + ID: 4, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "cache_ttl_override_enabled": true, + "cache_ttl_override_target": "5m", + }, + }, + } + + err := svc.RecordUsage(context.Background(), input) + require.NoError(t, err) + require.NotNil(t, repo.last) + + log := repo.last + require.Equal(t, 120, log.CacheCreationTokens) + require.Equal(t, 120, log.CacheCreation5mTokens) + require.Equal(t, 0, log.CacheCreation1hTokens) + require.True(t, log.CacheTTLOverridden) +} + +func TestRecordUsage_SimulateClaudeMaxEnabled_ExistingCacheCreationBypassesSimulation(t *testing.T) { + repo := &usageLogRepoRecordUsageStub{inserted: true} + svc := newGatewayServiceForRecordUsageTest(repo) + + groupID := int64(13) + input := &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "req-sim-3", + Model: "claude-sonnet-4", + Duration: time.Second, + Usage: ClaudeUsage{ + InputTokens: 20, + CacheCreationInputTokens: 120, + CacheCreation5mTokens: 120, + }, + }, + APIKey: &APIKey{ + ID: 3, + GroupID: &groupID, + Group: &Group{ + ID: groupID, + Platform: PlatformAnthropic, + RateMultiplier: 1, + SimulateClaudeMaxEnabled: true, + }, + }, + User: &User{ID: 4}, + Account: &Account{ + ID: 5, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "cache_ttl_override_enabled": true, + "cache_ttl_override_target": "5m", + }, + }, + } + + err := svc.RecordUsage(context.Background(), input) + require.NoError(t, err) + require.NotNil(t, repo.last) + + log := repo.last + require.Equal(t, 20, log.InputTokens) + require.Equal(t, 120, log.CacheCreation5mTokens) + require.Equal(t, 0, log.CacheCreation1hTokens) + require.Equal(t, 120, log.CacheCreationTokens) + require.False(t, log.CacheTTLOverridden, "existing cache_creation with SimulateClaudeMax enabled should skip account ttl override") +} diff --git a/backend/internal/service/gateway_response_usage_sync_test.go b/backend/internal/service/gateway_response_usage_sync_test.go new file mode 100644 index 0000000000..445ee8ad53 --- /dev/null +++ b/backend/internal/service/gateway_response_usage_sync_test.go @@ -0,0 +1,170 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestHandleNonStreamingResponse_UsageAlignedWithClaudeMaxSimulation(t *testing.T) { + gin.SetMode(gin.TestMode) + + svc := &GatewayService{ + cfg: &config.Config{}, + rateLimitService: &RateLimitService{}, + } + + account := &Account{ + ID: 11, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "cache_ttl_override_enabled": true, + "cache_ttl_override_target": "5m", + }, + } + group := &Group{ + ID: 99, + Platform: PlatformAnthropic, + SimulateClaudeMaxEnabled: true, + } + parsed := &ParsedRequest{ + Model: "claude-sonnet-4", + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "long cached context", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "new user question", + }, + }, + }, + }, + } + + upstreamBody := []byte(`{"id":"msg_1","model":"claude-sonnet-4","usage":{"input_tokens":120,"output_tokens":8}}`) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserBytes(upstreamBody), + } + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(nil)) + c.Set("api_key", &APIKey{Group: group}) + requestCtx := withClaudeMaxResponseRewriteContext(context.Background(), c, parsed) + + usage, err := svc.handleNonStreamingResponse(requestCtx, resp, c, account, "claude-sonnet-4", "claude-sonnet-4") + require.NoError(t, err) + require.NotNil(t, usage) + + var rendered struct { + Usage ClaudeUsage `json:"usage"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &rendered)) + rendered.Usage.CacheCreation5mTokens = int(gjson.GetBytes(rec.Body.Bytes(), "usage.cache_creation.ephemeral_5m_input_tokens").Int()) + rendered.Usage.CacheCreation1hTokens = int(gjson.GetBytes(rec.Body.Bytes(), "usage.cache_creation.ephemeral_1h_input_tokens").Int()) + + require.Equal(t, rendered.Usage.InputTokens, usage.InputTokens) + require.Equal(t, rendered.Usage.OutputTokens, usage.OutputTokens) + require.Equal(t, rendered.Usage.CacheCreationInputTokens, usage.CacheCreationInputTokens) + require.Equal(t, rendered.Usage.CacheCreation5mTokens, usage.CacheCreation5mTokens) + require.Equal(t, rendered.Usage.CacheCreation1hTokens, usage.CacheCreation1hTokens) + require.Equal(t, rendered.Usage.CacheReadInputTokens, usage.CacheReadInputTokens) + + require.Greater(t, usage.CacheCreation1hTokens, 0) + require.Equal(t, 0, usage.CacheCreation5mTokens) + require.Less(t, usage.InputTokens, 120) +} + +func TestHandleNonStreamingResponse_ClaudeMaxDisabled_NoSimulationIntercept(t *testing.T) { + gin.SetMode(gin.TestMode) + + svc := &GatewayService{ + cfg: &config.Config{}, + rateLimitService: &RateLimitService{}, + } + + account := &Account{ + ID: 12, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "cache_ttl_override_enabled": true, + "cache_ttl_override_target": "5m", + }, + } + group := &Group{ + ID: 100, + Platform: PlatformAnthropic, + SimulateClaudeMaxEnabled: false, + } + parsed := &ParsedRequest{ + Model: "claude-sonnet-4", + Messages: []any{ + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{ + "type": "text", + "text": "long cached context", + "cache_control": map[string]any{"type": "ephemeral"}, + }, + map[string]any{ + "type": "text", + "text": "new user question", + }, + }, + }, + }, + } + + upstreamBody := []byte(`{"id":"msg_2","model":"claude-sonnet-4","usage":{"input_tokens":120,"output_tokens":8}}`) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserBytes(upstreamBody), + } + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(nil)) + c.Set("api_key", &APIKey{Group: group}) + requestCtx := withClaudeMaxResponseRewriteContext(context.Background(), c, parsed) + + usage, err := svc.handleNonStreamingResponse(requestCtx, resp, c, account, "claude-sonnet-4", "claude-sonnet-4") + require.NoError(t, err) + require.NotNil(t, usage) + + require.Equal(t, 120, usage.InputTokens) + require.Equal(t, 0, usage.CacheCreationInputTokens) + require.Equal(t, 0, usage.CacheCreation5mTokens) + require.Equal(t, 0, usage.CacheCreation1hTokens) +} + +func ioNopCloserBytes(b []byte) *readCloserFromBytes { + return &readCloserFromBytes{Reader: bytes.NewReader(b)} +} + +type readCloserFromBytes struct { + *bytes.Reader +} + +func (r *readCloserFromBytes) Close() error { + return nil +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 7e962f7f3f..6c2f7e8e9c 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -40,7 +40,8 @@ import ( const ( claudeAPIURL = "https://api.anthropic.com/v1/messages?beta=true" claudeAPICountTokensURL = "https://api.anthropic.com/v1/messages/count_tokens?beta=true" - stickySessionTTL = time.Hour // 粘性会话TTL + stickySessionTTL = time.Hour // 粘性会话TTL + ClientAffinityTTL = 24 * time.Hour // 客户端亲和TTL defaultMaxLineSize = 500 * 1024 * 1024 // Canonical Claude Code banner. Keep it EXACT (no trailing whitespace/newlines) // to match real Claude CLI traffic as closely as possible. When we need a visual @@ -57,14 +58,21 @@ const ( claudeMimicDebugInfoKey = "claude_mimic_debug_info" ) +const ( + claudeMaxMessageOverheadTokens = 3 + claudeMaxBlockOverheadTokens = 1 + claudeMaxUnknownContentTokens = 4 +) + // ForceCacheBillingContextKey 强制缓存计费上下文键 // 用于粘性会话切换时,将 input_tokens 转为 cache_read_input_tokens 计费 type forceCacheBillingKeyType struct{} // accountWithLoad 账号与负载信息的组合,用于负载感知调度 type accountWithLoad struct { - account *Account - loadInfo *AccountLoadInfo + account *Account + loadInfo *AccountLoadInfo + affinityCount int64 // 亲和客户端数量(反向索引),越少越优先 } var ForceCacheBillingContextKey = forceCacheBillingKeyType{} @@ -328,6 +336,10 @@ var ( sseDataRe = regexp.MustCompile(`^data:\s*`) claudeCliUserAgentRe = regexp.MustCompile(`^claude-cli/\d+\.\d+\.\d+`) + // clientIDFromMetadataRegex 从 metadata.user_id 中提取客户端 ID(64位 hex) + // 格式: user_{64位hex}_account_... + clientIDFromMetadataRegex = regexp.MustCompile(`^user_([a-f0-9]{64})_account_`) + // claudeCodePromptPrefixes 用于检测 Claude Code 系统提示词的前缀列表 // 支持多种变体:标准版、Agent SDK 版、Explore Agent 版、Compact 版等 // 注意:前缀之间不应存在包含关系,否则会导致冗余匹配 @@ -351,6 +363,12 @@ var ErrNoAvailableAccounts = errors.New("no available accounts") // ErrClaudeCodeOnly 表示分组仅允许 Claude Code 客户端访问 var ErrClaudeCodeOnly = errors.New("this group only allows Claude Code clients") +// ErrAffinityNoSwitch 表示亲和账号不可用且不允许切换到其他账号 +var ErrAffinityNoSwitch = errors.New("affinity account unavailable and switching is disabled") + +// ErrAffinityLimitExceeded 表示亲和客户端限制已达上限 +var ErrAffinityLimitExceeded = errors.New("affinity client limit exceeded") + // allowedHeaders 白名单headers(参考CRS项目) var allowedHeaders = map[string]bool{ "accept": true, @@ -391,6 +409,39 @@ type GatewayCache interface { // DeleteSessionAccountID 删除粘性会话绑定,用于账号不可用时主动清理 // Delete sticky session binding, used to proactively clean up when account becomes unavailable DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error + + // GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员 + GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) + // UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL) + UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error + // GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员) + GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) + // GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重) + // accountGroups: map[accountID][]groupID + // 返回值成员格式为 {userID}/{clientID} + GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error) + // GetAccountAffinityClientsWithScores 获取单个账号跨所有分组的亲和客户端列表(含最后活跃时间) + GetAccountAffinityClientsWithScores(ctx context.Context, accountID int64, groupIDs []int64, ttl time.Duration) ([]AffinityClient, error) + // ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引) + // 用于账号关闭亲和时立即清理旧绑定 + ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error + // GetAffinityMultiCount 获取账号的多维度亲和计数 + // 返回: uniqueUsers, uniqueClients, perUserClients + GetAffinityMultiCount(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (users, clients, perUser int64, err error) +} + +// AffinityClient 亲和客户端信息(含用户 ID 和最后活跃时间) +type AffinityClient struct { + UserID int64 `json:"user_id"` + ClientID string `json:"client_id"` + LastActive time.Time `json:"last_active"` +} + +// SortAffinityClients 按最后活跃时间降序排序 +func SortAffinityClients(clients []AffinityClient) { + sort.Slice(clients, func(i, j int) bool { + return clients[i].LastActive.After(clients[j].LastActive) + }) } // derefGroupID safely dereferences *int64 to int64, returning 0 if nil @@ -461,6 +512,20 @@ func shouldClearStickySession(account *Account, requestedModel string) bool { return false } +// extractClientIDFromMetadata 从 metadata.user_id 中提取客户端 ID(64位 hex)。 +// 格式: user_{64位hex}_account_..._session_... +// 返回空字符串表示无法提取(非 Claude Code/Console 客户端)。 +func extractClientIDFromMetadata(metadataUserID string) string { + if metadataUserID == "" { + return "" + } + matches := clientIDFromMetadataRegex.FindStringSubmatch(metadataUserID) + if matches == nil { + return "" + } + return matches[1] +} + type AccountWaitPlan struct { AccountID int64 MaxConcurrency int @@ -1094,8 +1159,10 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context } // SelectAccountWithLoadAwareness selects account with load-awareness and wait plan. -// metadataUserID: 已废弃参数,会话限制现在统一使用 sessionHash -func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string) (*AccountSelectionResult, error) { +// 调度流程文档见 docs/ACCOUNT_SCHEDULING_FLOW.md 。 +// metadataUserID: 用于客户端亲和调度,从中提取客户端 ID +// sub2apiUserID: 系统用户 ID,用于二维亲和调度 +func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string, sub2apiUserID int64) (*AccountSelectionResult, error) { // 调试日志:记录调度入口参数 excludedIDsList := make([]int64, 0, len(excludedIDs)) for id := range excludedIDs { @@ -1125,6 +1192,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } + // 提取客户端 ID(用于客户端亲和调度) + affinityClientID := extractClientIDFromMetadata(metadataUserID) + affinityUserID := sub2apiUserID + if s.debugModelRoutingEnabled() && requestedModel != "" { groupPlatform := "" if group != nil { @@ -1146,6 +1217,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if err != nil { return nil, err } + if shouldFilterAccountWithoutClientID(account, affinityClientID) { + localExcluded[account.ID] = struct{}{} + continue + } result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) if err == nil && result.Acquired { @@ -1207,12 +1282,18 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if err != nil { return nil, err } + accounts = filterAccountsWithoutClientID(accounts, affinityClientID) if len(accounts) == 0 { return nil, ErrNoAvailableAccounts } ctx = s.withWindowCostPrefetch(ctx, accounts) ctx = s.withRPMPrefetch(ctx, accounts) + // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) + accountByID := make(map[int64]*Account, len(accounts)) + for i := range accounts { + accountByID[accounts[i].ID] = &accounts[i] + } isExcluded := func(accountID int64) bool { if excludedIDs == nil { return false @@ -1220,12 +1301,19 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro _, excluded := excludedIDs[accountID] return excluded } - - // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) - accountByID := make(map[int64]*Account, len(accounts)) - for i := range accounts { - accountByID[accounts[i].ID] = &accounts[i] - } + affinityFlow := newGatewayAffinityFlow( + s, + ctx, + groupID, + sessionHash, + requestedModel, + affinityClientID, + affinityUserID, + platform, + useMixed, + accountByID, + isExcluded, + ) // 获取模型路由配置(仅 anthropic 平台) var routingAccountIDs []int64 @@ -1388,7 +1476,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } if len(routingAvailable) > 0 { - // 排序:优先级 > 负载率 > 最后使用时间 + // 批量获取亲和客户端数量 + s.populateAffinityCounts(ctx, routingAvailable, derefGroupID(groupID)) + + // 排序:优先级 > 负载率 > 亲和客户端数 > 最后使用时间 sort.SliceStable(routingAvailable, func(i, j int) bool { a, b := routingAvailable[i], routingAvailable[j] if a.account.Priority != b.account.Priority { @@ -1397,6 +1488,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if a.loadInfo.LoadRate != b.loadInfo.LoadRate { return a.loadInfo.LoadRate < b.loadInfo.LoadRate } + if a.affinityCount != b.affinityCount { + return a.affinityCount < b.affinityCount + } switch { case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil: return true @@ -1422,6 +1516,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL) } + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && item.account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, item.account.ID, ClientAffinityTTL) + } if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) } @@ -1459,14 +1556,27 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } - // ============ Layer 1.5: 粘性会话(仅在无模型路由配置时生效) ============ - if len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { + // ============ Layer 1.3: 用户亲和预处理(pinned_users 自动注入) ============ + affinityFlow.preprocessPinnedUsers(accounts) + + // ============ Layer 1.4: 客户端亲和调度(优先于粘性会话) ============ + affinityHit := false + if affinityResult, hit, err := affinityFlow.trySelectAffinityAccount(); err != nil { + return nil, err + } else { + affinityHit = hit + if affinityResult != nil { + return affinityResult, nil + } + } + + // ============ Layer 1.5: 粘性会话(仅在无模型路由配置 且 亲和未命中时生效) ============ + if !affinityHit && len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { accountID := stickyAccountID if accountID > 0 && !isExcluded(accountID) { account, ok := accountByID[accountID] if ok { // 检查账户是否需要清理粘性会话绑定 - // Check if the account needs sticky session cleanup clearSticky := shouldClearStickySession(account, requestedModel) if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) @@ -1482,7 +1592,6 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) if err == nil && result.Acquired { // 会话数量限制检查 - // Session count limit check if !s.checkAndRegisterSession(ctx, account, sessionHash) { result.ReleaseFunc() // 释放槽位,继续到 Layer 2 } else { @@ -1497,10 +1606,8 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) if waitingCount < cfg.StickySessionMaxWaiting { // 会话数量限制检查(等待计划也需要占用会话配额) - // Session count limit check (wait plan also requires session quota) if !s.checkAndRegisterSession(ctx, account, sessionHash) { // 会话限制已满,继续到 Layer 2 - // Session limit full, continue to Layer 2 } else { return &AccountSelectionResult{ Account: account, @@ -1570,6 +1677,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) if err != nil { if result, ok := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); ok { + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && result.Account != nil && result.Account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, result.Account.ID, ClientAffinityTTL) + } return result, nil } } else { @@ -1587,13 +1697,37 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } - // 分层过滤选择:优先级 → 负载率 → LRU + // 批量获取亲和客户端数量(用于均衡分配新客户端) + s.populateAffinityCounts(ctx, available, derefGroupID(groupID)) + + // 分层过滤选择:优先级 → 亲和三区 → 负载率 → 亲和客户端数 → LRU for len(available) > 0 { // 1. 取优先级最小的集合 candidates := filterByMinPriority(available) - // 2. 取负载率最低的集合 + // 2. 按亲和三区过滤:绿区优先 → 黄区降级 → 红区移除(在同优先级内) + candidates = classifyByAffinityZone(candidates) + if len(candidates) == 0 { + // 当前优先级组全部在红区,移除后回退到下一优先级组 + minPri := available[0].account.Priority + for _, a := range available[1:] { + if a.account.Priority < minPri { + minPri = a.account.Priority + } + } + newAvailable := make([]accountWithLoad, 0, len(available)) + for _, a := range available { + if a.account.Priority != minPri { + newAvailable = append(newAvailable, a) + } + } + available = newAvailable + continue + } + // 3. 取负载率最低的集合 candidates = filterByMinLoadRate(candidates) - // 3. LRU 选择最久未用的账号 + // 3. 取亲和客户端数最少的集合 + candidates = filterByMinAffinityCount(candidates) + // 4. LRU 选择最久未用的账号 selected := selectByLRU(candidates, preferOAuth) if selected == nil { break @@ -1608,6 +1742,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) } + // 更新亲和关系 + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && selected.account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, selected.account.ID, ClientAffinityTTL) + } return &AccountSelectionResult{ Account: selected.account, Acquired: true, @@ -2365,6 +2503,36 @@ func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID in return s.accountRepo.GetByID(ctx, accountID) } +// populateAffinityCounts 批量获取账号的亲和客户端数量并填入 accountWithLoad 切片。 +// 仅当存在开启了客户端亲和的账号时才查询 Redis,否则跳过。 +func (s *GatewayService) populateAffinityCounts(ctx context.Context, accounts []accountWithLoad, groupID int64) { + if s.cache == nil || len(accounts) == 0 { + return + } + // 快速检查:是否有任何账号开启了亲和 + hasAffinity := false + for _, acc := range accounts { + if acc.account.IsAffinityEnabled() { + hasAffinity = true + break + } + } + if !hasAffinity { + return + } + accountIDs := make([]int64, len(accounts)) + for i, acc := range accounts { + accountIDs[i] = acc.account.ID + } + countMap, err := s.cache.GetAccountAffinityCountBatch(ctx, groupID, accountIDs, ClientAffinityTTL) + if err != nil { + return // 查询失败不影响调度,affinityCount 保持 0 + } + for i := range accounts { + accounts[i].affinityCount = countMap[accounts[i].account.ID] + } +} + // filterByMinPriority 过滤出优先级最小的账号集合 func filterByMinPriority(accounts []accountWithLoad) []accountWithLoad { if len(accounts) == 0 { @@ -2405,6 +2573,64 @@ func filterByMinLoadRate(accounts []accountWithLoad) []accountWithLoad { return result } +// filterByMinAffinityCount 过滤出亲和客户端数最少的账号集合 +func filterByMinAffinityCount(accounts []accountWithLoad) []accountWithLoad { + if len(accounts) == 0 { + return accounts + } + minCount := accounts[0].affinityCount + for _, acc := range accounts[1:] { + if acc.affinityCount < minCount { + minCount = acc.affinityCount + } + } + result := make([]accountWithLoad, 0, len(accounts)) + for _, acc := range accounts { + if acc.affinityCount == minCount { + result = append(result, acc) + } + } + return result +} + +// classifyByAffinityZone 按亲和分区对候选账号进行分类。 +// 返回值:仅绿区账号(有绿区时),否则返回黄区账号。红区账号被移除。 +// 如果没有任何账号开启了亲和三区配置(即 affinity_base <= 0),则原样返回所有账号。 +func classifyByAffinityZone(accounts []accountWithLoad) []accountWithLoad { + if len(accounts) == 0 { + return accounts + } + // 快速检查:是否有任何账号配置了 affinity_base + hasZoneConfig := false + for _, acc := range accounts { + if acc.account.IsAffinityEnabled() && acc.account.GetAffinityBase() > 0 { + hasZoneConfig = true + break + } + } + if !hasZoneConfig { + return accounts + } + + greens := make([]accountWithLoad, 0, len(accounts)) + yellows := make([]accountWithLoad, 0, len(accounts)) + for _, acc := range accounts { + zone := acc.account.GetAffinityZone(acc.affinityCount) + switch zone { + case AffinityZoneGreen: + greens = append(greens, acc) + case AffinityZoneYellow: + yellows = append(yellows, acc) + case AffinityZoneRed: + // 红区:移除,不参与调度 + } + } + if len(greens) > 0 { + return greens + } + return yellows +} + // selectByLRU 从集合中选择最久未用的账号 // 如果有多个账号具有相同的最小 LastUsedAt,则随机选择一个 func selectByLRU(accounts []accountWithLoad, preferOAuth bool) *accountWithLoad { @@ -3378,6 +3604,10 @@ func (s *GatewayService) isModelSupportedByAccount(account *Account, requestedMo _, ok := ResolveBedrockModelID(account, requestedModel) return ok } + // OpenAI 透传模式:仅替换认证,允许所有模型 + if account.Platform == PlatformOpenAI && account.IsOpenAIPassthroughEnabled() { + return true + } // OAuth/SetupToken 账号使用 Anthropic 标准映射(短ID → 长ID) if account.Platform == PlatformAnthropic && account.Type != AccountTypeAPIKey { requestedModel = claude.NormalizeModelID(requestedModel) @@ -4486,6 +4716,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A } // 处理正常响应 + ctx = withClaudeMaxResponseRewriteContext(ctx, c, parsed) // 触发上游接受回调(提前释放串行锁,不等流完成) if parsed.OnUpstreamAccepted != nil { @@ -6591,6 +6822,7 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http needModelReplace := originalModel != mappedModel clientDisconnected := false // 客户端断开标志,断开后继续读取上游以获取完整usage sawTerminalEvent := false + skipAccountTTLOverride := false pendingEventLines := make([]string, 0, 4) @@ -6652,17 +6884,25 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http if msg, ok := event["message"].(map[string]any); ok { if u, ok := msg["usage"].(map[string]any); ok { eventChanged = reconcileCachedTokens(u) || eventChanged + claudeMaxOutcome := applyClaudeMaxSimulationToUsageJSONMap(ctx, u, originalModel, account.ID) + if claudeMaxOutcome.Simulated { + skipAccountTTLOverride = true + } } } } if eventType == "message_delta" { if u, ok := event["usage"].(map[string]any); ok { eventChanged = reconcileCachedTokens(u) || eventChanged + claudeMaxOutcome := applyClaudeMaxSimulationToUsageJSONMap(ctx, u, originalModel, account.ID) + if claudeMaxOutcome.Simulated { + skipAccountTTLOverride = true + } } } // Cache TTL Override: 重写 SSE 事件中的 cache_creation 分类 - if account.IsCacheTTLOverrideEnabled() { + if account.IsCacheTTLOverrideEnabled() && !skipAccountTTLOverride { overrideTarget := account.GetCacheTTLOverrideTarget() if eventType == "message_start" { if msg, ok := event["message"].(map[string]any); ok { @@ -7094,8 +7334,13 @@ func (s *GatewayService) handleNonStreamingResponse(ctx context.Context, resp *h } } + claudeMaxOutcome := applyClaudeMaxSimulationToUsage(ctx, &response.Usage, originalModel, account.ID) + if claudeMaxOutcome.Simulated { + body = rewriteClaudeUsageJSONBytes(body, response.Usage) + } + // Cache TTL Override: 重写 non-streaming 响应中的 cache_creation 分类 - if account.IsCacheTTLOverrideEnabled() { + if account.IsCacheTTLOverrideEnabled() && !claudeMaxOutcome.Simulated { overrideTarget := account.GetCacheTTLOverrideTarget() if applyCacheTTLOverride(&response.Usage, overrideTarget) { // 同步更新 body JSON 中的嵌套 cache_creation 对象 @@ -7161,6 +7406,7 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID, // RecordUsageInput 记录使用量的输入参数 type RecordUsageInput struct { Result *ForwardResult + ParsedRequest *ParsedRequest APIKey *APIKey User *User Account *Account @@ -7473,9 +7719,19 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu result.Usage.InputTokens = 0 } + // Claude Max cache billing policy (group-level): + // - GatewayService 路径: Forward 已改写 usage(含 cache tokens)→ apply 见到 cache tokens 跳过 → simulatedClaudeMax=true(通过第二条件) + // - Antigravity 路径: Forward 中 hook 改写了客户端 SSE,但 ForwardResult.Usage 是原始值 → apply 实际执行模拟 → simulatedClaudeMax=true + var apiKeyGroup *Group + if apiKey != nil { + apiKeyGroup = apiKey.Group + } + claudeMaxOutcome := applyClaudeMaxCacheBillingPolicyToUsage(&result.Usage, input.ParsedRequest, apiKeyGroup, result.Model, account.ID) + simulatedClaudeMax := claudeMaxOutcome.Simulated || + (shouldApplyClaudeMaxBillingRulesForUsage(apiKeyGroup, result.Model, input.ParsedRequest) && hasCacheCreationTokens(result.Usage)) // Cache TTL Override: 确保计费时 token 分类与账号设置一致 cacheTTLOverridden := false - if account.IsCacheTTLOverrideEnabled() { + if account.IsCacheTTLOverrideEnabled() && !simulatedClaudeMax { applyCacheTTLOverride(&result.Usage, account.GetCacheTTLOverrideTarget()) cacheTTLOverridden = (result.Usage.CacheCreation5mTokens + result.Usage.CacheCreation1hTokens) > 0 } diff --git a/backend/internal/service/gateway_service_affinity_test.go b/backend/internal/service/gateway_service_affinity_test.go new file mode 100644 index 0000000000..d3dc6f4bb0 --- /dev/null +++ b/backend/internal/service/gateway_service_affinity_test.go @@ -0,0 +1,122 @@ +//go:build unit + +// Package service 提供 API 网关核心服务。 +// 本文件包含 SortAffinityClients 函数的单元测试, +// 验证 AffinityClient 切片排序逻辑在各种输入条件下的正确行为。 +// +// This file contains unit tests for the SortAffinityClients function, +// verifying correct sorting behavior for AffinityClient slices under +// various input conditions including empty, single, sorted, reverse, +// and duplicate-timestamp scenarios. +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestSortAffinityClients_Empty(t *testing.T) { + var clients []AffinityClient + SortAffinityClients(clients) + require.Empty(t, clients) + + clients = []AffinityClient{} + SortAffinityClients(clients) + require.Empty(t, clients) +} + +func TestSortAffinityClients_SingleElement(t *testing.T) { + now := time.Now() + clients := []AffinityClient{ + {ClientID: "client-1", LastActive: now}, + } + SortAffinityClients(clients) + require.Len(t, clients, 1) + require.Equal(t, "client-1", clients[0].ClientID) + require.Equal(t, now, clients[0].LastActive) +} + +func TestSortAffinityClients_AlreadySorted(t *testing.T) { + now := time.Now() + clients := []AffinityClient{ + {ClientID: "newest", LastActive: now}, + {ClientID: "middle", LastActive: now.Add(-1 * time.Hour)}, + {ClientID: "oldest", LastActive: now.Add(-2 * time.Hour)}, + } + SortAffinityClients(clients) + + require.Equal(t, "newest", clients[0].ClientID) + require.Equal(t, "middle", clients[1].ClientID) + require.Equal(t, "oldest", clients[2].ClientID) +} + +func TestSortAffinityClients_ReverseOrder(t *testing.T) { + now := time.Now() + clients := []AffinityClient{ + {ClientID: "oldest", LastActive: now.Add(-2 * time.Hour)}, + {ClientID: "middle", LastActive: now.Add(-1 * time.Hour)}, + {ClientID: "newest", LastActive: now}, + } + SortAffinityClients(clients) + + require.Equal(t, "newest", clients[0].ClientID) + require.Equal(t, "middle", clients[1].ClientID) + require.Equal(t, "oldest", clients[2].ClientID) +} + +func TestSortAffinityClients_SameTimestamps(t *testing.T) { + now := time.Now() + clients := []AffinityClient{ + {ClientID: "c1", LastActive: now}, + {ClientID: "c2", LastActive: now}, + {ClientID: "c3", LastActive: now}, + } + SortAffinityClients(clients) + + // 所有时间戳相同时,排序结果应保持稳定(sort.Slice 不保证稳定性, + // 但只要结果是某种确定的顺序即可)。 + // 验证所有元素仍然存在且时间相同。 + require.Len(t, clients, 3) + ids := map[string]bool{} + for _, c := range clients { + ids[c.ClientID] = true + require.Equal(t, now, c.LastActive) + } + require.True(t, ids["c1"]) + require.True(t, ids["c2"]) + require.True(t, ids["c3"]) +} + +func TestSortAffinityClients_MixedOrder(t *testing.T) { + now := time.Now() + clients := []AffinityClient{ + {ClientID: "c3", LastActive: now.Add(-30 * time.Minute)}, + {ClientID: "c1", LastActive: now}, + {ClientID: "c5", LastActive: now.Add(-2 * time.Hour)}, + {ClientID: "c2", LastActive: now.Add(-10 * time.Minute)}, + {ClientID: "c4", LastActive: now.Add(-1 * time.Hour)}, + } + SortAffinityClients(clients) + + // 按 LastActive 降序排列 + require.Equal(t, "c1", clients[0].ClientID) // now + require.Equal(t, "c2", clients[1].ClientID) // -10m + require.Equal(t, "c3", clients[2].ClientID) // -30m + require.Equal(t, "c4", clients[3].ClientID) // -1h + require.Equal(t, "c5", clients[4].ClientID) // -2h +} + +func TestSortAffinityClients_SubSecondDifferences(t *testing.T) { + base := time.Now() + clients := []AffinityClient{ + {ClientID: "early", LastActive: base}, + {ClientID: "late", LastActive: base.Add(500 * time.Millisecond)}, + } + SortAffinityClients(clients) + + // 500ms 差异也应正确排序(更晚的在前) + require.Equal(t, "late", clients[0].ClientID) + require.Equal(t, "early", clients[1].ClientID) +} diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index a78c56e768..f86e7656ef 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -288,6 +288,28 @@ func (m *mockGatewayCacheForGemini) DeleteSessionAccountID(ctx context.Context, return nil } +func (m *mockGatewayCacheForGemini) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (m *mockGatewayCacheForGemini) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (m *mockGatewayCacheForGemini) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (m *mockGatewayCacheForGemini) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (m *mockGatewayCacheForGemini) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (m *mockGatewayCacheForGemini) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (m *mockGatewayCacheForGemini) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} + // TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform 测试 Gemini 单平台选择 func TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform(t *testing.T) { ctx := context.Background() diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index e17032e000..b4d7845afd 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -50,6 +50,9 @@ type Group struct { // MCP XML 协议注入开关(仅 antigravity 平台使用) MCPXMLInject bool + // Claude usage 模拟开关:将无写缓存 usage 模拟为 claude-max 风格 + SimulateClaudeMaxEnabled bool + // 支持的模型系列(仅 antigravity 平台使用) // 可选值: claude, gemini_text, gemini_image SupportedModelScopes []string diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 789888cb36..3f465d458e 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -323,7 +323,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, nil } - if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) { + if req.RequestedModel != "" && !account.IsOpenAIPassthroughEnabled() && !account.IsModelSupported(req.RequestedModel) { return nil, nil } if !s.isAccountTransportCompatible(account, req.RequiredTransport) { @@ -582,7 +582,7 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( if !account.IsSchedulable() || !account.IsOpenAI() { continue } - if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) { + if req.RequestedModel != "" && !account.IsOpenAIPassthroughEnabled() && !account.IsModelSupported(req.RequestedModel) { continue } if !s.isAccountTransportCompatible(account, req.RequiredTransport) { diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index cf902c20df..53f7e099e6 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1393,7 +1393,7 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex if !acc.IsSchedulable() { continue } - if requestedModel != "" && !acc.IsModelSupported(requestedModel) { + if requestedModel != "" && !acc.IsOpenAIPassthroughEnabled() && !acc.IsModelSupported(requestedModel) { continue } candidates = append(candidates, acc) @@ -1554,7 +1554,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. if !fresh.IsSchedulable() || !fresh.IsOpenAI() { return nil } - if requestedModel != "" && !fresh.IsModelSupported(requestedModel) { + if requestedModel != "" && !fresh.IsOpenAIPassthroughEnabled() && !fresh.IsModelSupported(requestedModel) { return nil } return fresh diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 9e2f33f22a..2828802fe6 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -282,6 +282,28 @@ func (c *stubGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID i return nil } +func (c *stubGatewayCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (c *stubGatewayCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (c *stubGatewayCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (c *stubGatewayCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (c *stubGatewayCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (c *stubGatewayCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (c *stubGatewayCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} + func TestOpenAISelectAccountWithLoadAwareness_FiltersUnschedulable(t *testing.T) { now := time.Now() resetAt := now.Add(10 * time.Minute) diff --git a/backend/internal/service/openai_ws_state_store_test.go b/backend/internal/service/openai_ws_state_store_test.go index 235d42331d..9f35cd7ec6 100644 --- a/backend/internal/service/openai_ws_state_store_test.go +++ b/backend/internal/service/openai_ws_state_store_test.go @@ -193,6 +193,28 @@ func (c *openAIWSStateStoreTimeoutProbeCache) DeleteSessionAccountID(ctx context return nil } +func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { + return nil, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} + func TestOpenAIWSStateStore_RedisOpsUseShortTimeout(t *testing.T) { probe := &openAIWSStateStoreTimeoutProbeCache{} store := NewOpenAIWSStateStore(probe) diff --git a/backend/internal/service/ops_concurrency.go b/backend/internal/service/ops_concurrency.go index a571dd4df4..c03108c4b9 100644 --- a/backend/internal/service/ops_concurrency.go +++ b/backend/internal/service/ops_concurrency.go @@ -64,12 +64,9 @@ func (s *OpsService) getAccountsLoadMapBestEffort(ctx context.Context, accounts if acc.ID <= 0 { continue } - c := acc.Concurrency - if c <= 0 { - c = 1 - } - if prev, ok := unique[acc.ID]; !ok || c > prev { - unique[acc.ID] = c + lf := acc.EffectiveLoadFactor() + if prev, ok := unique[acc.ID]; !ok || lf > prev { + unique[acc.ID] = lf } } diff --git a/backend/internal/service/ops_metrics_collector.go b/backend/internal/service/ops_metrics_collector.go index f93481e7ff..6c33707138 100644 --- a/backend/internal/service/ops_metrics_collector.go +++ b/backend/internal/service/ops_metrics_collector.go @@ -391,7 +391,7 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con } batch = append(batch, AccountWithConcurrency{ ID: acc.ID, - MaxConcurrency: acc.Concurrency, + MaxConcurrency: acc.EffectiveLoadFactor(), }) } if len(batch) == 0 { diff --git a/backend/internal/service/ops_retry.go b/backend/internal/service/ops_retry.go index fdabbafde9..c0e814ab7b 100644 --- a/backend/internal/service/ops_retry.go +++ b/backend/internal/service/ops_retry.go @@ -519,7 +519,7 @@ func (s *OpsService) selectAccountForRetry(ctx context.Context, reqType opsRetry if s.gatewayService == nil { return nil, fmt.Errorf("gateway service not available") } - return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "") // 重试不使用会话限制 + return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "", int64(0)) // 重试不使用会话限制 default: return nil, fmt.Errorf("unsupported retry type: %s", reqType) } diff --git a/backend/internal/service/ops_system_log_sink_test.go b/backend/internal/service/ops_system_log_sink_test.go index 12a2ec0c7d..137ee33c72 100644 --- a/backend/internal/service/ops_system_log_sink_test.go +++ b/backend/internal/service/ops_system_log_sink_test.go @@ -183,6 +183,15 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) { if strings.TrimSpace(item.Message) == "" { t.Fatalf("message should not be empty") } + // writtenCount is incremented after BatchInsertSystemLogsFn returns, + // so poll briefly to avoid a race between the done signal and the atomic add. + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + if sink.Health().WrittenCount > 0 { + break + } + time.Sleep(time.Millisecond) + } health := sink.Health() if health.WrittenCount == 0 { t.Fatalf("written_count should be >0") diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go index 4a6e5d6cc3..c145cefbbf 100644 --- a/backend/internal/service/ratelimit_service_401_test.go +++ b/backend/internal/service/ratelimit_service_401_test.go @@ -94,6 +94,8 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t }) } +// TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError +// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testing.T) { repo := &rateLimitAccountRepoStub{} invalidator := &tokenCacheInvalidatorRecorder{err: errors.New("boom")} @@ -101,7 +103,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testin service.SetTokenCacheInvalidator(invalidator) account := &Account{ ID: 101, - Platform: PlatformGemini, + Platform: PlatformOpenAI, Type: AccountTypeOAuth, } diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index ece95c4e45..a360ebe670 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -210,6 +210,49 @@ func (s *SettingService) SetOnS3UpdateCallback(callback func()) { s.onS3Update = callback } +// SetOnStorageUpdateCallback 设置存储配置变更时的回调函数(用于刷新所有存储客户端缓存)。 +// 替代 SetOnS3UpdateCallback,支持 S3 + GDrive 统一刷新。 +func (s *SettingService) SetOnStorageUpdateCallback(callback func()) { + s.onS3Update = callback +} + +// --- 统一存储 Profile 方法别名 --- + +// ListSoraStorageProfiles 获取 Sora 存储多配置列表(统一方法名)。 +func (s *SettingService) ListSoraStorageProfiles(ctx context.Context) (*SoraS3ProfileList, error) { + return s.ListSoraS3Profiles(ctx) +} + +// CreateSoraStorageProfile 创建 Sora 存储配置(统一方法名)。 +func (s *SettingService) CreateSoraStorageProfile(ctx context.Context, profile *SoraS3Profile, setActive bool) (*SoraS3Profile, error) { + return s.CreateSoraS3Profile(ctx, profile, setActive) +} + +// UpdateSoraStorageProfile 更新 Sora 存储配置(统一方法名)。 +func (s *SettingService) UpdateSoraStorageProfile(ctx context.Context, profileID string, profile *SoraS3Profile) (*SoraS3Profile, error) { + return s.UpdateSoraS3Profile(ctx, profileID, profile) +} + +// DeleteSoraStorageProfile 删除 Sora 存储配置(统一方法名)。 +func (s *SettingService) DeleteSoraStorageProfile(ctx context.Context, profileID string) error { + return s.DeleteSoraS3Profile(ctx, profileID) +} + +// SetActiveSoraStorageProfile 设置激活的 Sora 存储配置(统一方法名)。 +func (s *SettingService) SetActiveSoraStorageProfile(ctx context.Context, profileID string) (*SoraS3Profile, error) { + return s.SetActiveSoraS3Profile(ctx, profileID) +} + +// GetActiveStorageProfile 获取当前激活的存储配置 profile。 +func (s *SettingService) GetActiveStorageProfile(ctx context.Context) (*SoraS3Profile, error) { + profiles, err := s.ListSoraS3Profiles(ctx) + if err != nil { + return nil, err + } + active := pickActiveSoraS3Profile(profiles.Items, profiles.ActiveProfileID) + return active, nil +} + // SetVersion sets the application version for injection into public settings func (s *SettingService) SetVersion(version string) { s.version = version @@ -1475,6 +1518,8 @@ type soraS3ProfilesStore struct { type soraS3ProfileStoreItem struct { ProfileID string `json:"profile_id"` Name string `json:"name"` + Provider string `json:"provider,omitempty"` // "s3" / "gdrive",空值视为 "s3" + AccessMode string `json:"access_mode,omitempty"` // "direct" / "proxy",空值视为 "direct" Enabled bool `json:"enabled"` Endpoint string `json:"endpoint"` Region string `json:"region"` @@ -1486,6 +1531,14 @@ type soraS3ProfileStoreItem struct { CDNURL string `json:"cdn_url"` DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` UpdatedAt string `json:"updated_at"` + + // --- Google Drive 专属 --- + AuthType string `json:"auth_type,omitempty"` + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` + RefreshToken string `json:"refresh_token,omitempty"` + ServiceAccountJSON string `json:"service_account_json,omitempty"` + FolderID string `json:"folder_id,omitempty"` } // GetSoraS3Settings 获取 Sora S3 存储配置(兼容旧单配置语义:返回当前激活配置) @@ -1597,6 +1650,8 @@ func (s *SettingService) CreateSoraS3Profile(ctx context.Context, profile *SoraS store.Items = append(store.Items, soraS3ProfileStoreItem{ ProfileID: profileID, Name: name, + Provider: profile.Provider, + AccessMode: profile.AccessMode, Enabled: profile.Enabled, Endpoint: strings.TrimSpace(profile.Endpoint), Region: strings.TrimSpace(profile.Region), @@ -1608,6 +1663,13 @@ func (s *SettingService) CreateSoraS3Profile(ctx context.Context, profile *SoraS CDNURL: strings.TrimSpace(profile.CDNURL), DefaultStorageQuotaBytes: maxInt64(profile.DefaultStorageQuotaBytes, 0), UpdatedAt: now, + // Google Drive 专属 + AuthType: profile.AuthType, + ClientID: strings.TrimSpace(profile.ClientID), + ClientSecret: profile.ClientSecret, + RefreshToken: profile.RefreshToken, + ServiceAccountJSON: profile.ServiceAccountJSON, + FolderID: strings.TrimSpace(profile.FolderID), }) if setActive || store.ActiveProfileID == "" { @@ -1653,6 +1715,8 @@ func (s *SettingService) UpdateSoraS3Profile(ctx context.Context, profileID stri return nil, infraerrors.BadRequest("SORA_S3_PROFILE_NAME_REQUIRED", "name is required") } target.Name = name + target.Provider = profile.Provider + target.AccessMode = profile.AccessMode target.Enabled = profile.Enabled target.Endpoint = strings.TrimSpace(profile.Endpoint) target.Region = strings.TrimSpace(profile.Region) @@ -1665,6 +1729,19 @@ func (s *SettingService) UpdateSoraS3Profile(ctx context.Context, profileID stri if profile.SecretAccessKey != "" { target.SecretAccessKey = profile.SecretAccessKey } + // Google Drive 专属 + target.AuthType = profile.AuthType + target.ClientID = strings.TrimSpace(profile.ClientID) + if profile.ClientSecret != "" { + target.ClientSecret = profile.ClientSecret + } + if profile.RefreshToken != "" { + target.RefreshToken = profile.RefreshToken + } + if profile.ServiceAccountJSON != "" { + target.ServiceAccountJSON = profile.ServiceAccountJSON + } + target.FolderID = strings.TrimSpace(profile.FolderID) target.UpdatedAt = time.Now().UTC().Format(time.RFC3339) store.Items[targetIndex] = target @@ -1967,6 +2044,8 @@ func convertSoraS3ProfilesStore(store *soraS3ProfilesStore) *SoraS3ProfileList { ProfileID: item.ProfileID, Name: item.Name, IsActive: item.ProfileID == store.ActiveProfileID, + Provider: item.Provider, + AccessMode: item.AccessMode, Enabled: item.Enabled, Endpoint: item.Endpoint, Region: item.Region, @@ -1979,6 +2058,16 @@ func convertSoraS3ProfilesStore(store *soraS3ProfilesStore) *SoraS3ProfileList { CDNURL: item.CDNURL, DefaultStorageQuotaBytes: item.DefaultStorageQuotaBytes, UpdatedAt: item.UpdatedAt, + // Google Drive 专属 + AuthType: item.AuthType, + ClientID: item.ClientID, + ClientSecret: item.ClientSecret, + ClientSecretConfigured: item.ClientSecret != "", + RefreshToken: item.RefreshToken, + RefreshTokenConfigured: item.RefreshToken != "", + ServiceAccountJSON: item.ServiceAccountJSON, + ServiceAccountConfigured: item.ServiceAccountJSON != "", + FolderID: item.FolderID, }) } return &SoraS3ProfileList{ diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 23188a09ac..182551c5d1 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -129,6 +129,8 @@ type SoraS3Profile struct { ProfileID string `json:"profile_id"` Name string `json:"name"` IsActive bool `json:"is_active"` + Provider string `json:"provider"` // "s3" / "gdrive",空值视为 "s3" + AccessMode string `json:"access_mode"` // "direct" / "proxy",空值视为 "direct" Enabled bool `json:"enabled"` Endpoint string `json:"endpoint"` Region string `json:"region"` @@ -141,6 +143,25 @@ type SoraS3Profile struct { CDNURL string `json:"cdn_url"` DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"` UpdatedAt string `json:"updated_at"` + + // --- Google Drive 专属 --- + AuthType string `json:"auth_type,omitempty"` // "oauth2" / "service_account" + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"-"` + ClientSecretConfigured bool `json:"client_secret_configured"` + RefreshToken string `json:"-"` + RefreshTokenConfigured bool `json:"refresh_token_configured"` + ServiceAccountJSON string `json:"-"` + ServiceAccountConfigured bool `json:"service_account_configured"` + FolderID string `json:"folder_id,omitempty"` +} + +// GetProvider 返回 Provider,空值视为 "s3"。 +func (p *SoraS3Profile) GetProvider() string { + if p.Provider == "" { + return SoraStorageTypeS3 + } + return p.Provider } // SoraS3ProfileList Sora S3 多配置列表 diff --git a/backend/internal/service/sora_gdrive_oauth.go b/backend/internal/service/sora_gdrive_oauth.go new file mode 100644 index 0000000000..61282053ca --- /dev/null +++ b/backend/internal/service/sora_gdrive_oauth.go @@ -0,0 +1,71 @@ +package service + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + + "golang.org/x/oauth2" + "golang.org/x/oauth2/google" + "google.golang.org/api/drive/v3" +) + +// SoraGDriveOAuthService 处理 Google Drive OAuth2 授权流程。 +type SoraGDriveOAuthService struct{} + +// NewSoraGDriveOAuthService 创建 GDrive OAuth 服务。 +func NewSoraGDriveOAuthService(_ *SettingService) *SoraGDriveOAuthService { + return &SoraGDriveOAuthService{} +} + +// GenerateAuthURL 生成 Google OAuth 授权 URL。 +func (s *SoraGDriveOAuthService) GenerateAuthURL(clientID, clientSecret, redirectURI string) (authURL, state string, err error) { + if clientID == "" || clientSecret == "" || redirectURI == "" { + return "", "", fmt.Errorf("client_id, client_secret, redirect_uri are required") + } + + config := &oauth2.Config{ + ClientID: clientID, + ClientSecret: clientSecret, + Endpoint: google.Endpoint, + Scopes: []string{drive.DriveFileScope}, + RedirectURL: redirectURI, + } + + // 生成随机 state + stateBytes := make([]byte, 16) + if _, err := rand.Read(stateBytes); err != nil { + return "", "", fmt.Errorf("generate state: %w", err) + } + state = hex.EncodeToString(stateBytes) + + authURL = config.AuthCodeURL(state, oauth2.AccessTypeOffline, oauth2.ApprovalForce) + return authURL, state, nil +} + +// ExchangeCode 用授权码换取 refresh_token。 +func (s *SoraGDriveOAuthService) ExchangeCode(ctx context.Context, clientID, clientSecret, redirectURI, code string) (string, error) { + if code == "" { + return "", fmt.Errorf("authorization code is required") + } + + config := &oauth2.Config{ + ClientID: clientID, + ClientSecret: clientSecret, + Endpoint: google.Endpoint, + Scopes: []string{drive.DriveFileScope}, + RedirectURL: redirectURI, + } + + token, err := config.Exchange(ctx, code) + if err != nil { + return "", fmt.Errorf("exchange code: %w", err) + } + + if token.RefreshToken == "" { + return "", fmt.Errorf("no refresh_token received, please revoke app access and try again") + } + + return token.RefreshToken, nil +} diff --git a/backend/internal/service/sora_gdrive_storage.go b/backend/internal/service/sora_gdrive_storage.go new file mode 100644 index 0000000000..afbc3ea5c9 --- /dev/null +++ b/backend/internal/service/sora_gdrive_storage.go @@ -0,0 +1,433 @@ +package service + +import ( + "context" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/google/uuid" + "golang.org/x/oauth2" + "golang.org/x/oauth2/google" + "google.golang.org/api/drive/v3" + "google.golang.org/api/option" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraGDriveStorage 负责 Sora 媒体文件的 Google Drive 存储操作。 +type SoraGDriveStorage struct { + settingService *SettingService + + mu sync.RWMutex + srv *drive.Service + cfg *SoraS3Profile // 缓存当前 GDrive 配置 + healthCheckedAt time.Time + healthErr error + healthTTL time.Duration +} + +const defaultGDriveHealthTTL = 30 * time.Second + +// NewSoraGDriveStorage 创建 Google Drive 存储服务实例。 +func NewSoraGDriveStorage(settingService *SettingService) *SoraGDriveStorage { + return &SoraGDriveStorage{ + settingService: settingService, + healthTTL: defaultGDriveHealthTTL, + } +} + +// StorageType 返回存储类型标识。 +func (s *SoraGDriveStorage) StorageType() string { + return SoraStorageTypeGDrive +} + +// Enabled 返回 Google Drive 存储是否已启用。 +func (s *SoraGDriveStorage) Enabled(ctx context.Context) bool { + profile := s.getActiveGDriveProfile(ctx) + if profile == nil { + return false + } + return profile.Enabled && s.hasValidCredentials(profile) +} + +// getActiveGDriveProfile 获取当前激活的 GDrive 配置。 +func (s *SoraGDriveStorage) getActiveGDriveProfile(ctx context.Context) *SoraS3Profile { + if s.settingService == nil { + return nil + } + profile, err := s.settingService.GetActiveStorageProfile(ctx) + if err != nil || profile == nil { + return nil + } + if profile.GetProvider() != SoraStorageTypeGDrive { + return nil + } + return profile +} + +// hasValidCredentials 检查 GDrive 配置是否有有效凭证。 +func (s *SoraGDriveStorage) hasValidCredentials(profile *SoraS3Profile) bool { + switch profile.AuthType { + case "oauth2": + return profile.ClientID != "" && profile.ClientSecret != "" && profile.RefreshToken != "" + case "service_account": + return profile.ServiceAccountJSON != "" + default: + return false + } +} + +// getService 获取或初始化 Drive 服务(带缓存)。 +func (s *SoraGDriveStorage) getService(ctx context.Context) (*drive.Service, *SoraS3Profile, error) { + s.mu.RLock() + if s.srv != nil && s.cfg != nil { + srv, cfg := s.srv, s.cfg + s.mu.RUnlock() + return srv, cfg, nil + } + s.mu.RUnlock() + + return s.initService(ctx) +} + +func (s *SoraGDriveStorage) initService(ctx context.Context) (*drive.Service, *SoraS3Profile, error) { + s.mu.Lock() + defer s.mu.Unlock() + + // 双重检查 + if s.srv != nil && s.cfg != nil { + return s.srv, s.cfg, nil + } + + profile := s.getActiveGDriveProfile(ctx) + if profile == nil { + return nil, nil, fmt.Errorf("no active gdrive profile found") + } + if !profile.Enabled { + return nil, nil, fmt.Errorf("gdrive storage is disabled") + } + + srv, err := s.buildDriveService(ctx, profile) + if err != nil { + return nil, nil, fmt.Errorf("build gdrive service: %w", err) + } + + s.srv = srv + s.cfg = profile + logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 客户端已初始化 auth_type=%s folder_id=%s", profile.AuthType, profile.FolderID) + return srv, profile, nil +} + +// buildDriveService 根据认证类型创建 Google Drive 服务。 +func (s *SoraGDriveStorage) buildDriveService(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) { + switch profile.AuthType { + case "oauth2": + return s.buildOAuth2Service(ctx, profile) + case "service_account": + return s.buildServiceAccountService(ctx, profile) + default: + return nil, fmt.Errorf("unsupported auth_type: %s", profile.AuthType) + } +} + +func (s *SoraGDriveStorage) buildOAuth2Service(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) { + config := &oauth2.Config{ + ClientID: profile.ClientID, + ClientSecret: profile.ClientSecret, + Endpoint: google.Endpoint, + Scopes: []string{drive.DriveFileScope}, + } + token := &oauth2.Token{ + RefreshToken: profile.RefreshToken, + } + tokenSource := config.TokenSource(ctx, token) + srv, err := drive.NewService(ctx, option.WithTokenSource(tokenSource)) + if err != nil { + return nil, fmt.Errorf("create gdrive oauth2 service: %w", err) + } + return srv, nil +} + +func (s *SoraGDriveStorage) buildServiceAccountService(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) { + srv, err := drive.NewService(ctx, option.WithCredentialsJSON([]byte(profile.ServiceAccountJSON))) //nolint:staticcheck // SA1019: admin-controlled service account JSON, safe to use + if err != nil { + return nil, fmt.Errorf("create gdrive service account service: %w", err) + } + return srv, nil +} + +// RefreshClient 清除缓存的 Drive 客户端。 +func (s *SoraGDriveStorage) RefreshClient() { + s.mu.Lock() + defer s.mu.Unlock() + s.srv = nil + s.cfg = nil + s.healthCheckedAt = time.Time{} + s.healthErr = nil + logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 客户端缓存已清除") +} + +// GDriveQuotaInfo 包含 Google Drive 配额信息。 +type GDriveQuotaInfo struct { + LimitBytes int64 `json:"limit_bytes"` + UsedBytes int64 `json:"used_bytes"` +} + +// TestConnection 测试 Google Drive 连接。 +func (s *SoraGDriveStorage) TestConnection(ctx context.Context) error { + srv, _, err := s.getService(ctx) + if err != nil { + return err + } + _, err = srv.About.Get().Fields("storageQuota").Context(ctx).Do() + if err != nil { + return fmt.Errorf("gdrive About.Get failed: %w", err) + } + return nil +} + +// GetQuotaInfo 获取 Google Drive 配额信息(总量和已用量)。 +func (s *SoraGDriveStorage) GetQuotaInfo(ctx context.Context) (*GDriveQuotaInfo, error) { + srv, _, err := s.getService(ctx) + if err != nil { + return nil, err + } + about, err := srv.About.Get().Fields("storageQuota").Context(ctx).Do() + if err != nil { + return nil, fmt.Errorf("gdrive About.Get failed: %w", err) + } + if about.StorageQuota == nil { + return nil, fmt.Errorf("storageQuota not available") + } + return &GDriveQuotaInfo{ + LimitBytes: about.StorageQuota.Limit, + UsedBytes: about.StorageQuota.Usage, + }, nil +} + +// TestFullCycle 执行完整的上传→获取链接→删除测试。 +func (s *SoraGDriveStorage) TestFullCycle(ctx context.Context) (map[string]any, error) { + srv, cfg, err := s.getService(ctx) + if err != nil { + return nil, fmt.Errorf("init client: %w", err) + } + + result := map[string]any{} + + // 1. 测试 API 连接 + about, err := srv.About.Get().Fields("storageQuota").Context(ctx).Do() + if err != nil { + return nil, fmt.Errorf("API connection failed: %w", err) + } + if about.StorageQuota != nil { + result["quota_limit_bytes"] = about.StorageQuota.Limit + result["quota_used_bytes"] = about.StorageQuota.Usage + } + + // 2. 上传测试文件 + testContent := "sub2api GDrive test file - " + time.Now().Format(time.RFC3339) + fileMeta := &drive.File{ + Name: "sub2api_test_" + uuid.NewString()[:8] + ".txt", + MimeType: "text/plain", + } + if cfg.FolderID != "" { + fileMeta.Parents = []string{cfg.FolderID} + } + uploaded, err := srv.Files.Create(fileMeta). + Media(strings.NewReader(testContent)). + Fields("id,name,size,webViewLink"). + Context(ctx).Do() + if err != nil { + return nil, fmt.Errorf("upload test file failed: %w", err) + } + result["uploaded_file_id"] = uploaded.Id + result["uploaded_file_name"] = uploaded.Name + result["uploaded_file_size"] = uploaded.Size + result["web_view_link"] = uploaded.WebViewLink + + // 3. 获取访问链接 + accessURL, err := s.GetAccessURL(ctx, uploaded.Id) + if err != nil { + // 即使获取链接失败,仍尝试清理 + _ = srv.Files.Delete(uploaded.Id).Context(ctx).Do() + return nil, fmt.Errorf("get access URL failed: %w", err) + } + result["access_url"] = accessURL + + // 4. 删除测试文件 + if err := srv.Files.Delete(uploaded.Id).Context(ctx).Do(); err != nil { + result["delete_warning"] = fmt.Sprintf("delete failed (manual cleanup needed): %v", err) + } else { + result["deleted"] = true + } + + result["status"] = "ok" + return result, nil +} + +// IsHealthy 返回 Google Drive 健康状态(带短缓存)。 +func (s *SoraGDriveStorage) IsHealthy(ctx context.Context) bool { + if s == nil { + return false + } + now := time.Now() + s.mu.RLock() + lastCheck := s.healthCheckedAt + lastErr := s.healthErr + ttl := s.healthTTL + s.mu.RUnlock() + + if ttl <= 0 { + ttl = defaultGDriveHealthTTL + } + if !lastCheck.IsZero() && now.Sub(lastCheck) < ttl { + return lastErr == nil + } + + err := s.TestConnection(ctx) + s.mu.Lock() + s.healthCheckedAt = time.Now() + s.healthErr = err + s.mu.Unlock() + return err == nil +} + +// UploadFromURL 从上游 URL 下载并上传到 Google Drive。 +// 返回 Google Drive 文件 ID 作为 objectKey、文件大小、存储类型。 +func (s *SoraGDriveStorage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) { + srv, cfg, err := s.getService(ctx) + if err != nil { + return "", 0, "", err + } + + // 下载源文件 + req, err := http.NewRequestWithContext(ctx, http.MethodGet, sourceURL, nil) + if err != nil { + return "", 0, "", fmt.Errorf("create download request: %w", err) + } + httpClient := &http.Client{Timeout: 5 * time.Minute} + resp, err := httpClient.Do(req) + if err != nil { + return "", 0, "", fmt.Errorf("download from upstream: %w", err) + } + defer func() { + _ = resp.Body.Close() + }() + + if resp.StatusCode != http.StatusOK { + return "", 0, "", &UpstreamDownloadError{StatusCode: resp.StatusCode} + } + + // 推断文件扩展名和 MIME + ext := fileExtFromURL(sourceURL) + if ext == "" { + ext = fileExtFromContentType(resp.Header.Get("Content-Type")) + } + if ext == "" { + ext = ".bin" + } + contentType := resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "application/octet-stream" + } + + // 生成文件名 + datePath := time.Now().Format("2006-01-02") + fileName := fmt.Sprintf("sora_%d_%s_%s%s", userID, datePath, uuid.NewString()[:8], ext) + + // 创建文件元数据 + fileMeta := &drive.File{ + Name: fileName, + MimeType: contentType, + } + if cfg.FolderID != "" { + fileMeta.Parents = []string{cfg.FolderID} + } + + // 使用 CountingReader 统计大小 + cr := &countingReader{Reader: resp.Body} + + // 上传到 Google Drive + created, err := srv.Files.Create(fileMeta). + Media(cr). + Fields("id, size"). + Context(ctx). + Do() + if err != nil { + return "", 0, "", fmt.Errorf("gdrive upload: %w", err) + } + + fileSize := cr.BytesRead + if created.Size > 0 { + fileSize = created.Size + } + + // 根据 access_mode 设置权限 + if cfg.AccessMode == "" || cfg.AccessMode == "direct" { + // 设为任何人可读 + _, permErr := srv.Permissions.Create(created.Id, &drive.Permission{ + Type: "anyone", + Role: "reader", + }).Context(ctx).Do() + if permErr != nil { + logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 设置公开权限失败 fileID=%s err=%v", created.Id, permErr) + } + } + + logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 上传完成 fileID=%s size=%d", created.Id, fileSize) + return created.Id, fileSize, SoraStorageTypeGDrive, nil +} + +// DeleteObjects 删除一组 Google Drive 文件。 +func (s *SoraGDriveStorage) DeleteObjects(ctx context.Context, objectKeys []string) error { + if len(objectKeys) == 0 { + return nil + } + + srv, _, err := s.getService(ctx) + if err != nil { + return err + } + + var lastErr error + for _, fileID := range objectKeys { + if err := srv.Files.Delete(fileID).Context(ctx).Do(); err != nil { + logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 删除失败 fileID=%s err=%v", fileID, err) + lastErr = err + } + } + return lastErr +} + +// GetAccessURL 获取 Google Drive 文件的访问 URL。 +func (s *SoraGDriveStorage) GetAccessURL(ctx context.Context, objectKey string) (string, error) { + _, cfg, err := s.getService(ctx) + if err != nil { + return "", err + } + + // CDN URL 优先 + if cfg.CDNURL != "" { + cdnBase := strings.TrimRight(cfg.CDNURL, "/") + return cdnBase + "/" + objectKey, nil + } + + // 默认使用 Google Drive 直链 + return fmt.Sprintf("https://drive.google.com/uc?export=download&id=%s", objectKey), nil +} + +// countingReader 包装 io.Reader 以统计读取的字节数。 +type countingReader struct { + Reader io.Reader + BytesRead int64 +} + +func (r *countingReader) Read(p []byte) (int, error) { + n, err := r.Reader.Read(p) + r.BytesRead += int64(n) + return n, err +} diff --git a/backend/internal/service/sora_gdrive_storage_test.go b/backend/internal/service/sora_gdrive_storage_test.go new file mode 100644 index 0000000000..e82f70c903 --- /dev/null +++ b/backend/internal/service/sora_gdrive_storage_test.go @@ -0,0 +1,124 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestSoraGDriveStorage_StorageType(t *testing.T) { + s := NewSoraGDriveStorage(nil) + assert.Equal(t, SoraStorageTypeGDrive, s.StorageType()) +} + +func TestSoraGDriveStorage_EnabledWithNilSettingService(t *testing.T) { + s := NewSoraGDriveStorage(nil) + assert.False(t, s.Enabled(context.Background())) +} + +func TestSoraGDriveStorage_IsHealthyWithNilReceiver(t *testing.T) { + var s *SoraGDriveStorage + assert.False(t, s.IsHealthy(context.Background())) +} + +func TestSoraGDriveStorage_GetServiceWithoutProfile(t *testing.T) { + s := NewSoraGDriveStorage(nil) + _, _, err := s.getService(context.Background()) + assert.Error(t, err) + assert.Contains(t, err.Error(), "no active gdrive profile") +} + +func TestSoraGDriveStorage_DeleteObjectsEmpty(t *testing.T) { + s := NewSoraGDriveStorage(nil) + err := s.DeleteObjects(context.Background(), []string{}) + assert.NoError(t, err) +} + +func TestSoraGDriveStorage_RefreshClient(t *testing.T) { + s := NewSoraGDriveStorage(nil) + // 不应 panic + s.RefreshClient() + assert.Nil(t, s.srv) + assert.Nil(t, s.cfg) +} + +func TestSoraGDriveStorage_HasValidCredentials(t *testing.T) { + s := NewSoraGDriveStorage(nil) + + tests := []struct { + name string + profile *SoraS3Profile + want bool + }{ + { + name: "oauth2 with all fields", + profile: &SoraS3Profile{ + AuthType: "oauth2", + ClientID: "id", + ClientSecret: "secret", + RefreshToken: "token", + }, + want: true, + }, + { + name: "oauth2 missing refresh token", + profile: &SoraS3Profile{ + AuthType: "oauth2", + ClientID: "id", + ClientSecret: "secret", + }, + want: false, + }, + { + name: "service_account with json", + profile: &SoraS3Profile{ + AuthType: "service_account", + ServiceAccountJSON: `{"type":"service_account"}`, + }, + want: true, + }, + { + name: "service_account without json", + profile: &SoraS3Profile{ + AuthType: "service_account", + }, + want: false, + }, + { + name: "unknown auth type", + profile: &SoraS3Profile{ + AuthType: "unknown", + }, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := s.hasValidCredentials(tt.profile) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestSoraStorageRouter_DefaultsToS3(t *testing.T) { + s3 := NewSoraS3Storage(nil) + router := NewSoraStorageRouter(nil, s3, nil) + // settingService 为 nil,应返回 s3Storage + backend := router.activeBackend(context.Background()) + assert.Equal(t, s3, backend) +} + +func TestSoraStorageRouter_StorageType(t *testing.T) { + router := NewSoraStorageRouter(nil, nil, nil) + assert.Equal(t, SoraStorageTypeS3, router.StorageType()) +} + +func TestSoraStorageRouter_RefreshAllNoPanic(t *testing.T) { + router := NewSoraStorageRouter(nil, nil, nil) + // 不应 panic + router.RefreshAll() +} diff --git a/backend/internal/service/sora_generation.go b/backend/internal/service/sora_generation.go index a704454b82..7323bb38ac 100644 --- a/backend/internal/service/sora_generation.go +++ b/backend/internal/service/sora_generation.go @@ -37,6 +37,7 @@ const ( // Sora 存储类型常量 const ( SoraStorageTypeS3 = "s3" + SoraStorageTypeGDrive = "gdrive" SoraStorageTypeLocal = "local" SoraStorageTypeUpstream = "upstream" SoraStorageTypeNone = "none" @@ -60,4 +61,5 @@ type SoraGenerationRepository interface { Delete(ctx context.Context, id int64) error List(ctx context.Context, params SoraGenerationListParams) ([]*SoraGeneration, int64, error) CountByUserAndStatus(ctx context.Context, userID int64, statuses []string) (int64, error) + CountByStorageType(ctx context.Context, storageType string, statuses []string) (int64, error) } diff --git a/backend/internal/service/sora_generation_service.go b/backend/internal/service/sora_generation_service.go index 22d5b51947..0eba7f91c9 100644 --- a/backend/internal/service/sora_generation_service.go +++ b/backend/internal/service/sora_generation_service.go @@ -35,21 +35,21 @@ type soraGenerationRepoConditionalUpdater interface { // SoraGenerationService 管理 Sora 客户端的生成记录 CRUD。 type SoraGenerationService struct { - genRepo SoraGenerationRepository - s3Storage *SoraS3Storage - quotaService *SoraQuotaService + genRepo SoraGenerationRepository + objectStorage SoraObjectStorage + quotaService *SoraQuotaService } // NewSoraGenerationService 创建生成记录服务。 func NewSoraGenerationService( genRepo SoraGenerationRepository, - s3Storage *SoraS3Storage, + objectStorage SoraObjectStorage, quotaService *SoraQuotaService, ) *SoraGenerationService { return &SoraGenerationService{ - genRepo: genRepo, - s3Storage: s3Storage, - quotaService: quotaService, + genRepo: genRepo, + objectStorage: objectStorage, + quotaService: quotaService, } } @@ -268,15 +268,15 @@ func (s *SoraGenerationService) Delete(ctx context.Context, id, userID int64) er return fmt.Errorf("无权删除此生成记录") } - // 清理 S3 文件 - if gen.StorageType == SoraStorageTypeS3 && len(gen.S3ObjectKeys) > 0 && s.s3Storage != nil { - if err := s.s3Storage.DeleteObjects(ctx, gen.S3ObjectKeys); err != nil { - logger.LegacyPrintf("service.sora_gen", "[SoraGen] S3 清理失败 id=%d err=%v", id, err) + // 清理存储文件(S3 / Google Drive) + if IsObjectStorageType(gen.StorageType) && len(gen.S3ObjectKeys) > 0 && s.objectStorage != nil { + if err := s.objectStorage.DeleteObjects(ctx, gen.S3ObjectKeys); err != nil { + logger.LegacyPrintf("service.sora_gen", "[SoraGen] 存储清理失败 id=%d type=%s err=%v", id, gen.StorageType, err) } } - // 释放配额(S3/本地均释放) - if gen.FileSizeBytes > 0 && (gen.StorageType == SoraStorageTypeS3 || gen.StorageType == SoraStorageTypeLocal) && s.quotaService != nil { + // 释放配额(对象存储/本地均释放) + if gen.FileSizeBytes > 0 && (IsObjectStorageType(gen.StorageType) || gen.StorageType == SoraStorageTypeLocal) && s.quotaService != nil { if err := s.quotaService.ReleaseUsage(ctx, userID, gen.FileSizeBytes); err != nil { logger.LegacyPrintf("service.sora_gen", "[SoraGen] 配额释放失败 id=%d err=%v", id, err) } @@ -290,9 +290,9 @@ func (s *SoraGenerationService) CountActiveByUser(ctx context.Context, userID in return s.genRepo.CountByUserAndStatus(ctx, userID, []string{SoraGenStatusPending, SoraGenStatusGenerating}) } -// ResolveMediaURLs 为 S3 记录动态生成预签名 URL。 +// ResolveMediaURLs 为对象存储记录动态生成访问 URL。 func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraGeneration) error { - if gen == nil || gen.StorageType != SoraStorageTypeS3 || s.s3Storage == nil { + if gen == nil || !IsObjectStorageType(gen.StorageType) || s.objectStorage == nil { return nil } if len(gen.S3ObjectKeys) == 0 { @@ -308,7 +308,7 @@ func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraG wg.Add(1) go func(i int, objectKey string) { defer wg.Done() - url, err := s.s3Storage.GetAccessURL(ctx, objectKey) + url, err := s.objectStorage.GetAccessURL(ctx, objectKey) if err != nil { errMu.Lock() if firstErr == nil { @@ -330,3 +330,22 @@ func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraG return nil } + +// StorageVideoStats 各存储类型的视频统计。 +type StorageVideoStats struct { + Completed int64 `json:"completed"` + InProgress int64 `json:"in_progress"` +} + +// CountByStorageType 按存储类型统计视频数量(completed 和 in_progress)。 +func (s *SoraGenerationService) CountByStorageType(ctx context.Context, storageType string) (completed, inProgress int64, err error) { + completed, err = s.genRepo.CountByStorageType(ctx, storageType, []string{SoraGenStatusCompleted}) + if err != nil { + return 0, 0, fmt.Errorf("count completed: %w", err) + } + inProgress, err = s.genRepo.CountByStorageType(ctx, storageType, []string{SoraGenStatusPending, SoraGenStatusGenerating}) + if err != nil { + return 0, 0, fmt.Errorf("count in_progress: %w", err) + } + return completed, inProgress, nil +} diff --git a/backend/internal/service/sora_generation_service_test.go b/backend/internal/service/sora_generation_service_test.go index 46f322c82c..0bde211f6c 100644 --- a/backend/internal/service/sora_generation_service_test.go +++ b/backend/internal/service/sora_generation_service_test.go @@ -115,6 +115,25 @@ func (r *stubGenRepo) CountByUserAndStatus(_ context.Context, userID int64, stat return count, nil } +func (r *stubGenRepo) CountByStorageType(_ context.Context, storageType string, statuses []string) (int64, error) { + if r.countErr != nil { + return 0, r.countErr + } + var count int64 + statusSet := make(map[string]struct{}) + for _, s := range statuses { + statusSet[s] = struct{}{} + } + for _, gen := range r.gens { + if gen.StorageType == storageType { + if _, ok := statusSet[gen.Status]; ok { + count++ + } + } + } + return count, nil +} + // ==================== Stub: UserRepository (用于 SoraQuotaService) ==================== var _ UserRepository = (*stubUserRepoForQuota)(nil) @@ -519,7 +538,7 @@ func TestDelete_S3Cleanup_NilS3(t *testing.T) { svc := NewSoraGenerationService(repo, nil, nil) err := svc.Delete(context.Background(), 1, 1) - require.NoError(t, err) // s3Storage 为 nil,跳过清理 + require.NoError(t, err) // objectStorage 为 nil,跳过清理 } func TestDelete_QuotaRelease_NilQuota(t *testing.T) { diff --git a/backend/internal/service/sora_object_storage.go b/backend/internal/service/sora_object_storage.go new file mode 100644 index 0000000000..f9e9a4b03e --- /dev/null +++ b/backend/internal/service/sora_object_storage.go @@ -0,0 +1,37 @@ +package service + +import "context" + +// SoraObjectStorage 是 Sora 媒体文件的通用对象存储接口。 +// S3 和 Google Drive 等存储后端均实现此接口。 +type SoraObjectStorage interface { + // Enabled 返回存储是否已启用且配置有效。 + Enabled(ctx context.Context) bool + + // IsHealthy 返回存储健康状态(带短缓存)。 + IsHealthy(ctx context.Context) bool + + // TestConnection 测试存储连接。 + TestConnection(ctx context.Context) error + + // UploadFromURL 从上游 URL 下载并上传到存储。 + // 返回 object key(S3 key 或 GDrive file ID)、文件大小、实际使用的存储类型。 + UploadFromURL(ctx context.Context, userID int64, sourceURL string) (objectKey string, sizeBytes int64, storageType string, err error) + + // DeleteObjects 删除一组存储对象。 + DeleteObjects(ctx context.Context, objectKeys []string) error + + // GetAccessURL 获取存储文件的访问 URL。 + GetAccessURL(ctx context.Context, objectKey string) (string, error) + + // RefreshClient 清除缓存客户端,配置变更时调用。 + RefreshClient() + + // StorageType 返回存储类型标识("s3" / "gdrive")。 + StorageType() string +} + +// IsObjectStorageType 判断是否为对象存储类型(S3 或 Google Drive)。 +func IsObjectStorageType(t string) bool { + return t == SoraStorageTypeS3 || t == SoraStorageTypeGDrive +} diff --git a/backend/internal/service/sora_s3_storage.go b/backend/internal/service/sora_s3_storage.go index 4c57390515..19d8b9c22f 100644 --- a/backend/internal/service/sora_s3_storage.go +++ b/backend/internal/service/sora_s3_storage.go @@ -212,29 +212,29 @@ func (s *SoraS3Storage) GenerateObjectKey(prefix string, userID int64, ext strin } // UploadFromURL 从上游 URL 下载并流式上传到 S3。 -// 返回 S3 object key。 -func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, error) { +// 返回 S3 object key、文件大小、存储类型。 +func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) { client, cfg, err := s.getClient(ctx) if err != nil { - return "", 0, err + return "", 0, "", err } // 下载源文件 req, err := http.NewRequestWithContext(ctx, http.MethodGet, sourceURL, nil) if err != nil { - return "", 0, fmt.Errorf("create download request: %w", err) + return "", 0, "", fmt.Errorf("create download request: %w", err) } httpClient := &http.Client{Timeout: 5 * time.Minute} resp, err := httpClient.Do(req) if err != nil { - return "", 0, fmt.Errorf("download from upstream: %w", err) + return "", 0, "", fmt.Errorf("download from upstream: %w", err) } defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { - return "", 0, &UpstreamDownloadError{StatusCode: resp.StatusCode} + return "", 0, "", &UpstreamDownloadError{StatusCode: resp.StatusCode} } // 推断文件扩展名 @@ -275,14 +275,14 @@ func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceU _ = writer.CloseWithError(copyErr) uploadErr := <-uploadErrCh if copyErr != nil { - return "", 0, fmt.Errorf("stream upload copy failed: %w", copyErr) + return "", 0, "", fmt.Errorf("stream upload copy failed: %w", copyErr) } if uploadErr != nil { - return "", 0, fmt.Errorf("s3 upload: %w", uploadErr) + return "", 0, "", fmt.Errorf("s3 upload: %w", uploadErr) } logger.LegacyPrintf("service.sora_s3", "[SoraS3] 上传完成 key=%s size=%d", objectKey, written) - return objectKey, written, nil + return objectKey, written, SoraStorageTypeS3, nil } func buildSoraS3Client(ctx context.Context, cfg *SoraS3Settings) (*s3.Client, string, error) { @@ -380,6 +380,11 @@ func (s *SoraS3Storage) GeneratePresignedURL(ctx context.Context, objectKey stri return result.URL, nil } +// StorageType 返回存储类型标识。 +func (s *SoraS3Storage) StorageType() string { + return SoraStorageTypeS3 +} + // GetMediaType 从 object key 推断媒体类型(image/video)。 func GetMediaTypeFromKey(objectKey string) string { ext := strings.ToLower(path.Ext(objectKey)) diff --git a/backend/internal/service/sora_s3_storage_test.go b/backend/internal/service/sora_s3_storage_test.go index 32ff9a6f0a..8ed7137d8e 100644 --- a/backend/internal/service/sora_s3_storage_test.go +++ b/backend/internal/service/sora_s3_storage_test.go @@ -238,7 +238,7 @@ func TestTestConnection_GetClientError(t *testing.T) { func TestUploadFromURL_GetClientError(t *testing.T) { s := NewSoraS3Storage(nil) - _, _, err := s.UploadFromURL(context.Background(), 1, "https://example.com/file.mp4") + _, _, _, err := s.UploadFromURL(context.Background(), 1, "https://example.com/file.mp4") require.Error(t, err) } diff --git a/backend/internal/service/sora_sdk_client.go b/backend/internal/service/sora_sdk_client.go index f9221c5b51..c839be4b03 100644 --- a/backend/internal/service/sora_sdk_client.go +++ b/backend/internal/service/sora_sdk_client.go @@ -316,20 +316,36 @@ func (c *SoraSDKClient) GetCameoStatus(ctx context.Context, account *Account, ca if err != nil { return nil, err } - sdkClient, err := c.getSDKClient(account) - if err != nil { - return nil, err - } - status, err := sdkClient.GetCameoStatus(ctx, token, cameoID) + + // 直接调用 Sora 后端 API 而非 SDK,以获取 SDK 未暴露的字段 + // (status_message、instruction_set_hint、instruction_set)。 + path := "/project_y/cameos/in_progress/" + cameoID + raw, err := c.doSoraBackendJSON(ctx, account, http.MethodGet, path, token, "", nil) if err != nil { return nil, c.wrapSDKError(err, account) } - return &SoraCameoStatus{ - Status: status.Status, - DisplayNameHint: status.DisplayNameHint, - UsernameHint: status.UsernameHint, - ProfileAssetURL: status.ProfileAssetURL, - }, nil + + return parseCameoStatusFromRaw(raw), nil +} + +// parseCameoStatusFromRaw 从原始 JSON 解析 SoraCameoStatus, +// 包含 SDK 未暴露的 status_message / instruction_set_hint / instruction_set 字段。 +func parseCameoStatusFromRaw(raw []byte) *SoraCameoStatus { + result := gjson.ParseBytes(raw) + cameoStatus := &SoraCameoStatus{ + Status: strings.TrimSpace(result.Get("status").String()), + StatusMessage: strings.TrimSpace(result.Get("status_message").String()), + DisplayNameHint: strings.TrimSpace(result.Get("display_name_hint").String()), + UsernameHint: strings.TrimSpace(result.Get("username_hint").String()), + ProfileAssetURL: strings.TrimSpace(result.Get("profile_asset_url").String()), + } + if v := result.Get("instruction_set_hint"); v.Exists() { + cameoStatus.InstructionSetHint = v.Value() + } + if v := result.Get("instruction_set"); v.Exists() { + cameoStatus.InstructionSet = v.Value() + } + return cameoStatus } func (c *SoraSDKClient) DownloadCharacterImage(ctx context.Context, account *Account, imageURL string) ([]byte, error) { @@ -925,26 +941,32 @@ func (c *SoraSDKClient) exchangeSessionToken(ctx context.Context, account *Accou return accessToken, expiresAt, nil } -// applyRecoveredToken 将恢复的 token 写入账号内存和数据库 +// applyRecoveredToken 将恢复的 token 写入账号内存和数据库。 +// 使用 copy-on-write 避免并发 map 写入 panic:创建新 map 后整体替换指针。 func (c *SoraSDKClient) applyRecoveredToken(ctx context.Context, account *Account, accessToken, refreshToken, expiresAt, sessionToken string) { if account == nil { return } - if account.Credentials == nil { - account.Credentials = make(map[string]any) + + // Copy-on-write: 复制旧 map 并写入新值,最后整体替换 + oldCreds := account.Credentials + newCreds := make(map[string]any, len(oldCreds)+4) + for k, v := range oldCreds { + newCreds[k] = v } if strings.TrimSpace(accessToken) != "" { - account.Credentials["access_token"] = accessToken + newCreds["access_token"] = accessToken } if strings.TrimSpace(refreshToken) != "" { - account.Credentials["refresh_token"] = refreshToken + newCreds["refresh_token"] = refreshToken } if strings.TrimSpace(expiresAt) != "" { - account.Credentials["expires_at"] = expiresAt + newCreds["expires_at"] = expiresAt } if strings.TrimSpace(sessionToken) != "" { - account.Credentials["session_token"] = sessionToken + newCreds["session_token"] = sessionToken } + account.Credentials = newCreds if c.accountRepo != nil { if err := c.accountRepo.Update(ctx, account); err != nil && c.debugEnabled() { diff --git a/backend/internal/service/sora_storage_router.go b/backend/internal/service/sora_storage_router.go new file mode 100644 index 0000000000..83e055f5df --- /dev/null +++ b/backend/internal/service/sora_storage_router.go @@ -0,0 +1,129 @@ +package service + +import ( + "context" + "fmt" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraStorageRouter 根据激活 profile 的 provider 字段路由到对应存储实现。 +// 实现 SoraObjectStorage 接口。 +type SoraStorageRouter struct { + settingService *SettingService + s3Storage *SoraS3Storage + gdriveStorage SoraObjectStorage // 可为 nil(GDrive 未实现时) +} + +// NewSoraStorageRouter 创建存储路由。 +func NewSoraStorageRouter( + settingService *SettingService, + s3Storage *SoraS3Storage, + gdriveStorage SoraObjectStorage, +) *SoraStorageRouter { + return &SoraStorageRouter{ + settingService: settingService, + s3Storage: s3Storage, + gdriveStorage: gdriveStorage, + } +} + +// activeBackend 返回当前激活 profile 对应的存储后端。 +func (r *SoraStorageRouter) activeBackend(ctx context.Context) SoraObjectStorage { + if r.settingService == nil { + return r.s3Storage // 默认 S3 + } + + profile, err := r.settingService.GetActiveStorageProfile(ctx) + if err != nil || profile == nil { + return r.s3Storage // 默认 S3 + } + + switch profile.GetProvider() { + case SoraStorageTypeGDrive: + if r.gdriveStorage != nil { + return r.gdriveStorage + } + logger.LegacyPrintf("service.storage_router", "[StorageRouter] GDrive 后端未初始化,降级到 S3") + return r.s3Storage + default: + return r.s3Storage + } +} + +func (r *SoraStorageRouter) Enabled(ctx context.Context) bool { + backend := r.activeBackend(ctx) + if backend == nil { + return false + } + return backend.Enabled(ctx) +} + +func (r *SoraStorageRouter) IsHealthy(ctx context.Context) bool { + backend := r.activeBackend(ctx) + if backend == nil { + return false + } + return backend.IsHealthy(ctx) +} + +func (r *SoraStorageRouter) TestConnection(ctx context.Context) error { + backend := r.activeBackend(ctx) + if backend == nil { + return fmt.Errorf("no storage backend available") + } + return backend.TestConnection(ctx) +} + +func (r *SoraStorageRouter) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) { + backend := r.activeBackend(ctx) + if backend == nil { + return "", 0, "", fmt.Errorf("no storage backend available") + } + return backend.UploadFromURL(ctx, userID, sourceURL) +} + +func (r *SoraStorageRouter) DeleteObjects(ctx context.Context, objectKeys []string) error { + backend := r.activeBackend(ctx) + if backend == nil { + return fmt.Errorf("no storage backend available") + } + return backend.DeleteObjects(ctx, objectKeys) +} + +func (r *SoraStorageRouter) GetAccessURL(ctx context.Context, objectKey string) (string, error) { + backend := r.activeBackend(ctx) + if backend == nil { + return "", fmt.Errorf("no storage backend available") + } + return backend.GetAccessURL(ctx, objectKey) +} + +func (r *SoraStorageRouter) RefreshClient() { + if r.s3Storage != nil { + r.s3Storage.RefreshClient() + } + if r.gdriveStorage != nil { + r.gdriveStorage.RefreshClient() + } +} + +// RefreshAll 刷新所有后端客户端(用作配置变更回调)。 +func (r *SoraStorageRouter) RefreshAll() { + r.RefreshClient() +} + +func (r *SoraStorageRouter) StorageType() string { + // 不带 context 的方法,返回默认值 + // 真实的 StorageType 在 activeBackend 中动态确定 + return SoraStorageTypeS3 +} + +// StorageTypeWithContext 返回当前激活后端的存储类型。 +func (r *SoraStorageRouter) StorageTypeWithContext(ctx context.Context) string { + backend := r.activeBackend(ctx) + if backend == nil { + return SoraStorageTypeS3 + } + return backend.StorageType() +} diff --git a/backend/internal/service/sora_task.go b/backend/internal/service/sora_task.go new file mode 100644 index 0000000000..6b909b9b2b --- /dev/null +++ b/backend/internal/service/sora_task.go @@ -0,0 +1,63 @@ +package service + +import ( + "context" + "time" +) + +// ── 任务状态常量 ── + +const ( + SoraTaskQueued = "queued" + SoraTaskInProgress = "in_progress" + SoraTaskCompleted = "completed" + SoraTaskFailed = "failed" +) + +// ── 对象类型常量 ── + +const ( + SoraObjectVideo = "video" + SoraObjectCharacter = "character" + SoraObjectImage = "image" +) + +// SoraTask 表示一个 Sora 异步任务记录。 +type SoraTask struct { + ID string + AccountID int64 + APIKeyID *int64 + UpstreamTaskID string + ObjectType string + Model string + Prompt string + Status string + Progress int + VideoURL string // 原始上游 URL + StoredKey string // 存储后的 key(本地路径或 S3 key) + StorageType string // local / s3 / gdrive / 空 + ShareID string + CharacterInfo *SoraCharacter + ErrorMessage string + ErrorType string + RequestBody []byte + Seconds string + Size string + CreatedAt time.Time + CompletedAt *time.Time +} + +// SoraCharacter 角色信息。 +type SoraCharacter struct { + Username string `json:"username"` + DisplayName string `json:"display_name"` +} + +// SoraTaskRepository 持久层接口。 +type SoraTaskRepository interface { + Create(ctx context.Context, task *SoraTask) error + GetByID(ctx context.Context, id string) (*SoraTask, error) + GetByIDAndAPIKey(ctx context.Context, id string, apiKeyID int64) (*SoraTask, error) + Update(ctx context.Context, task *SoraTask) error + ListPending(ctx context.Context) ([]*SoraTask, error) +} diff --git a/backend/internal/service/sora_task_service.go b/backend/internal/service/sora_task_service.go new file mode 100644 index 0000000000..93197c44c8 --- /dev/null +++ b/backend/internal/service/sora_task_service.go @@ -0,0 +1,546 @@ +package service + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraTaskService manages Sora async task creation and status queries. +type SoraTaskService struct { + repo SoraTaskRepository + accountRepo AccountRepository + soraClient SoraClient + httpUpstream HTTPUpstream +} + +func NewSoraTaskService( + repo SoraTaskRepository, + accountRepo AccountRepository, + soraClient SoraClient, + httpUpstream HTTPUpstream, +) *SoraTaskService { + return &SoraTaskService{ + repo: repo, + accountRepo: accountRepo, + soraClient: soraClient, + httpUpstream: httpUpstream, + } +} + +// ── Request structs ── + +type CreateVideoRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + StyleID string `json:"style_id,omitempty"` + Orientation string `json:"orientation,omitempty"` + Image string `json:"image,omitempty"` + Video string `json:"video,omitempty"` + RemixTargetID string `json:"remix_target_id,omitempty"` + Size string `json:"size,omitempty"` +} + +type CreateImageRequest struct { + Model string `json:"model,omitempty"` + Prompt string `json:"prompt"` + Image string `json:"image,omitempty"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + N int `json:"n,omitempty"` +} + +type EditImageRequest struct { + Image string `json:"image"` + Prompt string `json:"prompt"` + Model string `json:"model,omitempty"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` +} + +type RemixRequest struct { + Prompt string `json:"prompt"` +} + +// ── Response structs ── + +type SoraTaskResponse struct { + ID string `json:"id"` + Object string `json:"object"` + CreatedAt int64 `json:"created_at"` + Status string `json:"status"` + Model string `json:"model,omitempty"` + Prompt string `json:"prompt,omitempty"` + Progress int `json:"progress"` + VideoURL string `json:"video_url,omitempty"` + ShareID string `json:"share_id,omitempty"` + Seconds string `json:"seconds,omitempty"` + Size string `json:"size,omitempty"` + Character *SoraCharacter `json:"character,omitempty"` + URL string `json:"url,omitempty"` + CompletedAt *int64 `json:"completed_at,omitempty"` + Error *SoraTaskError `json:"error,omitempty"` +} + +type SoraTaskError struct { + Message string `json:"message"` + Type string `json:"type"` +} + +func TaskToResponse(t *SoraTask) *SoraTaskResponse { + resp := &SoraTaskResponse{ + ID: t.ID, + Object: t.ObjectType, + CreatedAt: t.CreatedAt.Unix(), + Status: t.Status, + Model: t.Model, + Prompt: t.Prompt, + Progress: t.Progress, + Seconds: t.Seconds, + Size: t.Size, + } + if t.VideoURL != "" { + resp.VideoURL = t.VideoURL + } + if t.ShareID != "" { + resp.ShareID = t.ShareID + } + if t.CharacterInfo != nil { + resp.Character = t.CharacterInfo + } + if t.CompletedAt != nil { + ts := t.CompletedAt.Unix() + resp.CompletedAt = &ts + } + if t.Status == SoraTaskFailed && t.ErrorMessage != "" { + resp.Error = &SoraTaskError{ + Message: t.ErrorMessage, + Type: t.ErrorType, + } + } + return resp +} + +// ── Task creation ── + +func (s *SoraTaskService) CreateVideoTask( + ctx context.Context, + apiKeyID int64, + account *Account, + req *CreateVideoRequest, + body []byte, +) (*SoraTask, error) { + modelCfg, ok := GetSoraModelConfig(req.Model) + + objectType := SoraObjectVideo + if req.Video != "" && req.Prompt == "" { + objectType = SoraObjectCharacter + } + + taskID := generateTaskID(objectType) + + seconds := "" + size := "" + if ok && modelCfg.Type == "video" { + seconds = fmt.Sprintf("%d", modelCfg.Frames/30) + if modelCfg.Orientation == "landscape" { + size = "1920x1080" + } else { + size = "1080x1920" + } + } + + task := &SoraTask{ + ID: taskID, + AccountID: account.ID, + APIKeyID: &apiKeyID, + ObjectType: objectType, + Model: req.Model, + Prompt: req.Prompt, + Status: SoraTaskQueued, + Progress: 0, + Seconds: seconds, + Size: size, + CreatedAt: time.Now(), + } + + if account.Type == AccountTypeAPIKey { + task.RequestBody = body + upstreamResp, err := s.forwardCreateToUpstream(ctx, account, "/v1/videos", body) + if err != nil { + return nil, fmt.Errorf("forward to upstream: %w", err) + } + s.applyUpstreamResponse(task, upstreamResp) + } else { + if s.soraClient == nil || !s.soraClient.Enabled() { + return nil, fmt.Errorf("sora SDK client not configured") + } + upstreamID, err := s.createViaSdk(ctx, account, req, modelCfg, objectType) + if err != nil { + return nil, fmt.Errorf("sdk create task: %w", err) + } + task.UpstreamTaskID = upstreamID + } + + if err := s.repo.Create(ctx, task); err != nil { + return nil, fmt.Errorf("save task: %w", err) + } + return task, nil +} + +func (s *SoraTaskService) CreateImageGeneration( + ctx context.Context, + apiKeyID int64, + account *Account, + req *CreateImageRequest, + body []byte, +) (*SoraTask, error) { + taskID := generateTaskID(SoraObjectImage) + + task := &SoraTask{ + ID: taskID, + AccountID: account.ID, + APIKeyID: &apiKeyID, + ObjectType: SoraObjectImage, + Model: req.Model, + Prompt: req.Prompt, + Status: SoraTaskQueued, + Progress: 0, + Size: req.Size, + CreatedAt: time.Now(), + } + + if account.Type == AccountTypeAPIKey { + task.RequestBody = body + upstreamResp, err := s.forwardCreateToUpstream(ctx, account, "/v1/images/generations", body) + if err != nil { + return nil, fmt.Errorf("forward to upstream: %w", err) + } + s.applyUpstreamResponse(task, upstreamResp) + } else { + if s.soraClient == nil || !s.soraClient.Enabled() { + return nil, fmt.Errorf("sora SDK client not configured") + } + modelCfg, _ := GetSoraModelConfig(req.Model) + width, height := modelCfg.Width, modelCfg.Height + if width == 0 { + width = 360 + } + if height == 0 { + height = 360 + } + upstreamID, err := s.soraClient.CreateImageTask(ctx, account, SoraImageRequest{ + Prompt: req.Prompt, + Width: width, + Height: height, + }) + if err != nil { + return nil, fmt.Errorf("sdk create image task: %w", err) + } + task.UpstreamTaskID = upstreamID + } + + if err := s.repo.Create(ctx, task); err != nil { + return nil, fmt.Errorf("save task: %w", err) + } + return task, nil +} + +func (s *SoraTaskService) GetTask(ctx context.Context, taskID string, apiKeyID int64) (*SoraTask, error) { + return s.repo.GetByIDAndAPIKey(ctx, taskID, apiKeyID) +} + +func (s *SoraTaskService) GetTaskByID(ctx context.Context, taskID string) (*SoraTask, error) { + return s.repo.GetByID(ctx, taskID) +} + +// ListPendingTasks returns tasks in queued or in_progress status. +func (s *SoraTaskService) ListPendingTasks(ctx context.Context) ([]*SoraTask, error) { + return s.repo.ListPending(ctx) +} + +// UpdateTask persists task state changes. +func (s *SoraTaskService) UpdateTask(ctx context.Context, task *SoraTask) error { + return s.repo.Update(ctx, task) +} + +func (s *SoraTaskService) GetAccountByID(ctx context.Context, accountID int64) (*Account, error) { + return s.accountRepo.GetByID(ctx, accountID) +} + +// ── Internal methods ── + +func (s *SoraTaskService) createViaSdk( + ctx context.Context, + account *Account, + req *CreateVideoRequest, + modelCfg SoraModelConfig, + objectType string, +) (string, error) { + if objectType == SoraObjectCharacter { + return "", fmt.Errorf("character creation via /v1/videos is not yet supported for OAuth accounts") + } + + orientation := modelCfg.Orientation + if req.Orientation != "" { + orientation = req.Orientation + } + + videoReq := SoraVideoRequest{ + Prompt: req.Prompt, + Orientation: orientation, + Frames: modelCfg.Frames, + Model: modelCfg.Model, + Size: modelCfg.Size, + RemixTargetID: req.RemixTargetID, + } + return s.soraClient.CreateVideoTask(ctx, account, videoReq) +} + +func (s *SoraTaskService) forwardCreateToUpstream( + ctx context.Context, + account *Account, + path string, + body []byte, +) (map[string]any, error) { + apiKey := account.GetCredential("api_key") + baseURL := account.GetBaseURL() + if apiKey == "" || baseURL == "" { + return nil, fmt.Errorf("account %d missing api_key or base_url", account.ID) + } + + upstreamURL := strings.TrimRight(baseURL, "/") + path + req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, fmt.Errorf("upstream request: %w", err) + } + defer func() { _ = resp.Body.Close() }() + + respBody, err := io.ReadAll(io.LimitReader(resp.Body, 256*1024)) + if err != nil { + return nil, fmt.Errorf("read upstream response: %w", err) + } + if resp.StatusCode >= 400 { + return nil, &SoraUpstreamError{StatusCode: resp.StatusCode, Body: respBody} + } + + var result map[string]any + if err := json.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("parse upstream response: %w", err) + } + return result, nil +} + +func (s *SoraTaskService) forwardGetToUpstream( + ctx context.Context, + account *Account, + path string, +) (map[string]any, error) { + apiKey := account.GetCredential("api_key") + baseURL := account.GetBaseURL() + if apiKey == "" || baseURL == "" { + return nil, fmt.Errorf("account %d missing api_key or base_url", account.ID) + } + + upstreamURL := strings.TrimRight(baseURL, "/") + path + req, err := http.NewRequestWithContext(ctx, http.MethodGet, upstreamURL, nil) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+apiKey) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, fmt.Errorf("upstream request: %w", err) + } + defer func() { _ = resp.Body.Close() }() + + respBody, err := io.ReadAll(io.LimitReader(resp.Body, 256*1024)) + if err != nil { + return nil, fmt.Errorf("read upstream response: %w", err) + } + if resp.StatusCode >= 400 { + return nil, fmt.Errorf("upstream error %d: %s", resp.StatusCode, string(respBody)) + } + + var result map[string]any + if err := json.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("parse upstream response: %w", err) + } + return result, nil +} + +func (s *SoraTaskService) applyUpstreamResponse(task *SoraTask, resp map[string]any) { + if id, ok := resp["id"].(string); ok && id != "" { + task.UpstreamTaskID = id + } + if status, ok := resp["status"].(string); ok && status != "" { + task.Status = status + } + if progress, ok := resp["progress"].(float64); ok { + task.Progress = int(progress) + } + if videoURL, ok := resp["video_url"].(string); ok { + task.VideoURL = videoURL + } + if shareID, ok := resp["share_id"].(string); ok { + task.ShareID = shareID + } + if seconds, ok := resp["seconds"].(string); ok { + task.Seconds = seconds + } + if size, ok := resp["size"].(string); ok { + task.Size = size + } + if obj, ok := resp["object"].(string); ok && obj != "" { + task.ObjectType = obj + } +} + +// PollTask polls a single task's latest status (called by worker). +func (s *SoraTaskService) PollTask(ctx context.Context, task *SoraTask, account *Account) error { + if account.Type == AccountTypeAPIKey { + return s.pollUpstreamTask(ctx, task, account) + } + return s.pollSdkTask(ctx, task, account) +} + +func (s *SoraTaskService) pollUpstreamTask(ctx context.Context, task *SoraTask, account *Account) error { + upstreamID := task.UpstreamTaskID + if upstreamID == "" { + upstreamID = task.ID + } + + path := fmt.Sprintf("/v1/videos/%s", upstreamID) + resp, err := s.forwardGetToUpstream(ctx, account, path) + if err != nil { + logger.LegacyPrintf("service.sora_task", "[PollUpstream] task=%s error=%v", task.ID, err) + return err + } + + s.applyUpstreamResponse(task, resp) + now := time.Now() + if task.Status == SoraTaskCompleted || task.Status == SoraTaskFailed { + task.CompletedAt = &now + } + + if errObj, ok := resp["error"].(map[string]any); ok { + if msg, ok := errObj["message"].(string); ok { + task.ErrorMessage = msg + } + if typ, ok := errObj["type"].(string); ok { + task.ErrorType = typ + } + } + + return s.repo.Update(ctx, task) +} + +func (s *SoraTaskService) pollSdkTask(ctx context.Context, task *SoraTask, account *Account) error { + if s.soraClient == nil { + return fmt.Errorf("sora SDK client not configured") + } + upstreamID := task.UpstreamTaskID + if upstreamID == "" { + return fmt.Errorf("task %s has no upstream_task_id", task.ID) + } + + now := time.Now() + + switch task.ObjectType { + case SoraObjectImage: + status, err := s.soraClient.GetImageTask(ctx, account, upstreamID) + if err != nil { + return err + } + task.Progress = int(status.ProgressPct) + switch status.Status { + case "complete", "completed": + task.Status = SoraTaskCompleted + task.Progress = 100 + task.CompletedAt = &now + if len(status.URLs) > 0 { + task.VideoURL = status.URLs[0] + } + case "failed": + task.Status = SoraTaskFailed + task.CompletedAt = &now + task.ErrorMessage = status.ErrorMsg + task.ErrorType = "server_error" + default: + task.Status = SoraTaskInProgress + } + + case SoraObjectVideo, SoraObjectCharacter: + status, err := s.soraClient.GetVideoTask(ctx, account, upstreamID) + if err != nil { + return err + } + task.Progress = status.ProgressPct + switch status.Status { + case "complete", "completed": + task.Status = SoraTaskCompleted + task.Progress = 100 + task.CompletedAt = &now + if len(status.URLs) > 0 { + task.VideoURL = status.URLs[0] + } + if status.GenerationID != "" { + task.ShareID = status.GenerationID + } + case "failed": + task.Status = SoraTaskFailed + task.CompletedAt = &now + task.ErrorMessage = status.ErrorMsg + task.ErrorType = "server_error" + default: + task.Status = SoraTaskInProgress + } + + default: + return fmt.Errorf("unknown object type: %s", task.ObjectType) + } + + return s.repo.Update(ctx, task) +} + +func generateTaskID(objectType string) string { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + b = []byte(fmt.Sprintf("%x", time.Now().UnixNano())) + } + switch objectType { + case SoraObjectCharacter: + return "char_" + hex.EncodeToString(b) + case SoraObjectImage: + return "img_" + hex.EncodeToString(b) + default: + return "video_" + hex.EncodeToString(b) + } +} diff --git a/backend/internal/service/sora_task_worker.go b/backend/internal/service/sora_task_worker.go new file mode 100644 index 0000000000..059b671031 --- /dev/null +++ b/backend/internal/service/sora_task_worker.go @@ -0,0 +1,183 @@ +package service + +import ( + "context" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraTaskWorker polls unfinished Sora tasks in the background. +// When a task completes, it downloads media to configured storage if available. +type SoraTaskWorker struct { + taskService *SoraTaskService + accountRepo AccountRepository + objectStorage SoraObjectStorage + mediaStorage *SoraMediaStorage + interval time.Duration + pollTimeout time.Duration + stopCh chan struct{} + stopOnce sync.Once + wg sync.WaitGroup +} + +func NewSoraTaskWorker( + taskService *SoraTaskService, + accountRepo AccountRepository, + objectStorage SoraObjectStorage, + mediaStorage *SoraMediaStorage, + interval time.Duration, +) *SoraTaskWorker { + if interval <= 0 { + interval = 60 * time.Second + } + return &SoraTaskWorker{ + taskService: taskService, + accountRepo: accountRepo, + objectStorage: objectStorage, + mediaStorage: mediaStorage, + interval: interval, + pollTimeout: 30 * time.Second, + stopCh: make(chan struct{}), + } +} + +func (w *SoraTaskWorker) Start() { + if w == nil || w.taskService == nil { + return + } + w.wg.Add(1) + go func() { + defer w.wg.Done() + ticker := time.NewTicker(w.interval) + defer ticker.Stop() + + logger.LegacyPrintf("service.sora_task_worker", "[Start] polling interval=%s", w.interval) + w.pollAll() + + for { + select { + case <-ticker.C: + w.pollAll() + case <-w.stopCh: + return + } + } + }() +} + +func (w *SoraTaskWorker) Stop() { + if w == nil { + return + } + w.stopOnce.Do(func() { + close(w.stopCh) + }) + w.wg.Wait() +} + +func (w *SoraTaskWorker) pollAll() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) + defer cancel() + + tasks, err := w.taskService.ListPendingTasks(ctx) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", "[PollAll] list pending tasks error: %v", err) + return + } + if len(tasks) == 0 { + return + } + + logger.LegacyPrintf("service.sora_task_worker", "[PollAll] found %d pending tasks", len(tasks)) + + for _, task := range tasks { + select { + case <-w.stopCh: + return + default: + } + w.pollOne(ctx, task) + } +} + +func (w *SoraTaskWorker) pollOne(ctx context.Context, task *SoraTask) { + pollCtx, cancel := context.WithTimeout(ctx, w.pollTimeout) + defer cancel() + + account, err := w.accountRepo.GetByID(pollCtx, task.AccountID) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s get account=%d error: %v", task.ID, task.AccountID, err) + w.markTaskFailed(pollCtx, task, "account not found") + return + } + + if err := w.taskService.PollTask(pollCtx, task, account); err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s poll error: %v", task.ID, err) + return + } + + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s status=%s progress=%d", task.ID, task.Status, task.Progress) + + if task.Status == SoraTaskCompleted && task.VideoURL != "" && task.StoredKey == "" { + w.tryStoreMedia(pollCtx, task) + } +} + +// tryStoreMedia downloads media to configured storage. +// Priority: S3 > local disk > keep upstream URL. +func (w *SoraTaskWorker) tryStoreMedia(ctx context.Context, task *SoraTask) { + mediaType := "video" + if task.ObjectType == SoraObjectImage { + mediaType = "image" + } + + if w.objectStorage != nil && w.objectStorage.Enabled(ctx) { + key, _, storageType, err := w.objectStorage.UploadFromURL(ctx, 0, task.VideoURL) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s object storage upload error: %v", task.ID, err) + } else { + task.StoredKey = key + task.StorageType = storageType + if updateErr := w.taskService.UpdateTask(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s update stored_key error: %v", task.ID, updateErr) + } + return + } + } + + if w.mediaStorage != nil && w.mediaStorage.Enabled() { + stored, err := w.mediaStorage.StoreFromURLs(ctx, mediaType, []string{task.VideoURL}) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s local storage error: %v", task.ID, err) + return + } + if len(stored) > 0 && stored[0] != task.VideoURL { + task.StoredKey = stored[0] + task.StorageType = "local" + if updateErr := w.taskService.UpdateTask(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s update stored_key error: %v", task.ID, updateErr) + } + } + } +} + +func (w *SoraTaskWorker) markTaskFailed(ctx context.Context, task *SoraTask, message string) { + now := time.Now() + task.Status = SoraTaskFailed + task.ErrorMessage = message + task.ErrorType = "server_error" + task.CompletedAt = &now + if updateErr := w.taskService.UpdateTask(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s update failed: %v", task.ID, updateErr) + } +} diff --git a/backend/internal/service/token_refresh_service_test.go b/backend/internal/service/token_refresh_service_test.go index f48de65e42..abc3cb9630 100644 --- a/backend/internal/service/token_refresh_service_test.go +++ b/backend/internal/service/token_refresh_service_test.go @@ -390,7 +390,7 @@ func TestTokenRefreshService_RefreshWithRetry_ClearsTempUnschedulable(t *testing err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour) require.NoError(t, err) require.Equal(t, 1, repo.updateCalls) - require.Equal(t, 1, repo.clearTempCalls) // DB 清除 + require.Equal(t, 1, repo.clearTempCalls) // DB 清除 require.Equal(t, 1, tempCache.deleteCalls) // Redis 缓存也应清除 } diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go index 05fe50560b..7f6c748fa6 100644 --- a/backend/internal/service/user_service_test.go +++ b/backend/internal/service/user_service_test.go @@ -47,8 +47,8 @@ func (m *mockUserRepo) RemoveGroupFromAllowedGroups(context.Context, int64) (int } func (m *mockUserRepo) AddGroupToAllowedGroups(context.Context, int64, int64) error { return nil } func (m *mockUserRepo) UpdateTotpSecret(context.Context, int64, *string) error { return nil } -func (m *mockUserRepo) EnableTotp(context.Context, int64) error { return nil } -func (m *mockUserRepo) DisableTotp(context.Context, int64) error { return nil } +func (m *mockUserRepo) EnableTotp(context.Context, int64) error { return nil } +func (m *mockUserRepo) DisableTotp(context.Context, int64) error { return nil } // --- mock: APIKeyAuthCacheInvalidator --- diff --git a/backend/internal/testutil/stubs.go b/backend/internal/testutil/stubs.go index bc572e1137..1ae72127bc 100644 --- a/backend/internal/testutil/stubs.go +++ b/backend/internal/testutil/stubs.go @@ -100,6 +100,27 @@ func (c StubGatewayCache) RefreshSessionTTL(_ context.Context, _ int64, _ string func (c StubGatewayCache) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { return nil } +func (c StubGatewayCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (c StubGatewayCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (c StubGatewayCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} +func (c StubGatewayCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (c StubGatewayCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (c StubGatewayCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { + return nil, nil +} +func (c StubGatewayCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} // ============================================================ // StubSessionLimitCache — service.SessionLimitCache 的空实现 diff --git a/backend/internal/web/embed_on.go b/backend/internal/web/embed_on.go index ffca98a54f..c8ef4c33a3 100644 --- a/backend/internal/web/embed_on.go +++ b/backend/internal/web/embed_on.go @@ -10,6 +10,8 @@ import ( "io" "io/fs" "net/http" + "os" + "path/filepath" "strings" "time" @@ -32,11 +34,12 @@ type PublicSettingsProvider interface { // FrontendServer serves the embedded frontend with settings injection type FrontendServer struct { - distFS fs.FS - fileServer http.Handler - baseHTML []byte - cache *HTMLCache - settings PublicSettingsProvider + distFS fs.FS + fileServer http.Handler + baseHTML []byte + cache *HTMLCache + settings PublicSettingsProvider + overrideDir string // local file override directory } // NewFrontendServer creates a new frontend server with settings injection @@ -62,11 +65,12 @@ func NewFrontendServer(settingsProvider PublicSettingsProvider) (*FrontendServer cache.SetBaseHTML(baseHTML) return &FrontendServer{ - distFS: distFS, - fileServer: http.FileServer(http.FS(distFS)), - baseHTML: baseHTML, - cache: cache, - settings: settingsProvider, + distFS: distFS, + fileServer: http.FileServer(http.FS(distFS)), + baseHTML: baseHTML, + cache: cache, + settings: settingsProvider, + overrideDir: filepath.Join("data", "public"), }, nil } @@ -99,6 +103,11 @@ func (s *FrontendServer) Middleware() gin.HandlerFunc { return } + // Try local override first + if s.tryServeOverride(c, cleanPath) { + return + } + // Serve static files normally s.fileServer.ServeHTTP(c.Writer, c.Request) c.Abort() @@ -114,6 +123,22 @@ func (s *FrontendServer) fileExists(path string) bool { return true } +// tryServeOverride checks if a local override file exists and serves it. +// Files in overrideDir take precedence over embedded files. +func (s *FrontendServer) tryServeOverride(c *gin.Context, cleanPath string) bool { + if s.overrideDir == "" { + return false + } + filePath := filepath.Join(s.overrideDir, filepath.Clean("/"+cleanPath)) + info, err := os.Stat(filePath) + if err != nil || info.IsDir() { + return false + } + c.File(filePath) + c.Abort() + return true +} + func (s *FrontendServer) serveIndexHTML(c *gin.Context) { // Get nonce from context (generated by SecurityHeaders middleware) nonce := middleware.GetNonceFromContext(c) @@ -226,6 +251,7 @@ func ServeEmbeddedFrontend() gin.HandlerFunc { panic("failed to get dist subdirectory: " + err.Error()) } fileServer := http.FileServer(http.FS(distFS)) + overrideDir := filepath.Join("data", "public") return func(c *gin.Context) { path := c.Request.URL.Path @@ -242,6 +268,10 @@ func ServeEmbeddedFrontend() gin.HandlerFunc { if file, err := distFS.Open(cleanPath); err == nil { _ = file.Close() + // Try local override first + if tryServeOverrideFile(c, overrideDir, cleanPath) { + return + } fileServer.ServeHTTP(c.Writer, c.Request) c.Abort() return @@ -251,6 +281,21 @@ func ServeEmbeddedFrontend() gin.HandlerFunc { } } +// tryServeOverrideFile is a standalone version of tryServeOverride for legacy usage. +func tryServeOverrideFile(c *gin.Context, overrideDir, cleanPath string) bool { + if overrideDir == "" { + return false + } + filePath := filepath.Join(overrideDir, filepath.Clean("/"+cleanPath)) + info, err := os.Stat(filePath) + if err != nil || info.IsDir() { + return false + } + c.File(filePath) + c.Abort() + return true +} + func shouldBypassEmbeddedFrontend(path string) bool { trimmed := strings.TrimSpace(path) return strings.HasPrefix(trimmed, "/api/") || diff --git a/backend/migrations/056_add_sonnet46_to_model_mapping.sql b/backend/migrations/056_add_sonnet46_to_model_mapping.sql new file mode 100644 index 0000000000..aa7657d71a --- /dev/null +++ b/backend/migrations/056_add_sonnet46_to_model_mapping.sql @@ -0,0 +1,42 @@ +-- Add claude-sonnet-4-6 to model_mapping for all Antigravity accounts +-- +-- Background: +-- Antigravity now supports claude-sonnet-4-6 +-- +-- Strategy: +-- Directly overwrite the entire model_mapping with updated mappings +-- This ensures consistency with DefaultAntigravityModelMapping in constants.go + +UPDATE accounts +SET credentials = jsonb_set( + credentials, + '{model_mapping}', + '{ + "claude-opus-4-6-thinking": "claude-opus-4-6-thinking", + "claude-opus-4-6": "claude-opus-4-6-thinking", + "claude-opus-4-5-thinking": "claude-opus-4-6-thinking", + "claude-opus-4-5-20251101": "claude-opus-4-6-thinking", + "claude-sonnet-4-6": "claude-sonnet-4-6", + "claude-sonnet-4-5": "claude-sonnet-4-5", + "claude-sonnet-4-5-thinking": "claude-sonnet-4-5-thinking", + "claude-sonnet-4-5-20250929": "claude-sonnet-4-5", + "claude-haiku-4-5": "claude-sonnet-4-5", + "claude-haiku-4-5-20251001": "claude-sonnet-4-5", + "gemini-2.5-flash": "gemini-2.5-flash", + "gemini-2.5-flash-lite": "gemini-2.5-flash-lite", + "gemini-2.5-flash-thinking": "gemini-2.5-flash-thinking", + "gemini-2.5-pro": "gemini-2.5-pro", + "gemini-3-flash": "gemini-3-flash", + "gemini-3-pro-high": "gemini-3-pro-high", + "gemini-3-pro-low": "gemini-3-pro-low", + "gemini-3-pro-image": "gemini-3-pro-image", + "gemini-3-flash-preview": "gemini-3-flash", + "gemini-3-pro-preview": "gemini-3-pro-high", + "gemini-3-pro-image-preview": "gemini-3-pro-image", + "gpt-oss-120b-medium": "gpt-oss-120b-medium", + "tab_flash_lite_preview": "tab_flash_lite_preview" + }'::jsonb +) +WHERE platform = 'antigravity' + AND deleted_at IS NULL + AND credentials->'model_mapping' IS NOT NULL; diff --git a/backend/migrations/057_add_gemini31_pro_to_model_mapping.sql b/backend/migrations/057_add_gemini31_pro_to_model_mapping.sql new file mode 100644 index 0000000000..6305e717bc --- /dev/null +++ b/backend/migrations/057_add_gemini31_pro_to_model_mapping.sql @@ -0,0 +1,45 @@ +-- Add gemini-3.1-pro-high, gemini-3.1-pro-low, gemini-3.1-pro-preview to model_mapping +-- +-- Background: +-- Antigravity now supports gemini-3.1-pro-high and gemini-3.1-pro-low +-- +-- Strategy: +-- Directly overwrite the entire model_mapping with updated mappings +-- This ensures consistency with DefaultAntigravityModelMapping in constants.go + +UPDATE accounts +SET credentials = jsonb_set( + credentials, + '{model_mapping}', + '{ + "claude-opus-4-6-thinking": "claude-opus-4-6-thinking", + "claude-opus-4-6": "claude-opus-4-6-thinking", + "claude-opus-4-5-thinking": "claude-opus-4-6-thinking", + "claude-opus-4-5-20251101": "claude-opus-4-6-thinking", + "claude-sonnet-4-6": "claude-sonnet-4-6", + "claude-sonnet-4-5": "claude-sonnet-4-5", + "claude-sonnet-4-5-thinking": "claude-sonnet-4-5-thinking", + "claude-sonnet-4-5-20250929": "claude-sonnet-4-5", + "claude-haiku-4-5": "claude-sonnet-4-5", + "claude-haiku-4-5-20251001": "claude-sonnet-4-5", + "gemini-2.5-flash": "gemini-2.5-flash", + "gemini-2.5-flash-lite": "gemini-2.5-flash-lite", + "gemini-2.5-flash-thinking": "gemini-2.5-flash-thinking", + "gemini-2.5-pro": "gemini-2.5-pro", + "gemini-3-flash": "gemini-3-flash", + "gemini-3-pro-high": "gemini-3-pro-high", + "gemini-3-pro-low": "gemini-3-pro-low", + "gemini-3-pro-image": "gemini-3-pro-image", + "gemini-3-flash-preview": "gemini-3-flash", + "gemini-3-pro-preview": "gemini-3-pro-high", + "gemini-3-pro-image-preview": "gemini-3-pro-image", + "gemini-3.1-pro-high": "gemini-3.1-pro-high", + "gemini-3.1-pro-low": "gemini-3.1-pro-low", + "gemini-3.1-pro-preview": "gemini-3.1-pro-high", + "gpt-oss-120b-medium": "gpt-oss-120b-medium", + "tab_flash_lite_preview": "tab_flash_lite_preview" + }'::jsonb +) +WHERE platform = 'antigravity' + AND deleted_at IS NULL + AND credentials->'model_mapping' IS NOT NULL; diff --git a/backend/migrations/060_add_group_simulate_claude_max.sql b/backend/migrations/060_add_group_simulate_claude_max.sql new file mode 100644 index 0000000000..55662dfdac --- /dev/null +++ b/backend/migrations/060_add_group_simulate_claude_max.sql @@ -0,0 +1,3 @@ +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS simulate_claude_max_enabled BOOLEAN NOT NULL DEFAULT FALSE; + diff --git a/backend/migrations/070_add_sora_tasks.sql b/backend/migrations/070_add_sora_tasks.sql new file mode 100644 index 0000000000..3a484976d9 --- /dev/null +++ b/backend/migrations/070_add_sora_tasks.sql @@ -0,0 +1,36 @@ +-- +migrate Up + +-- Sora 异步任务表:记录视频/角色/图片生成任务状态 +CREATE TABLE IF NOT EXISTS sora_tasks ( + id VARCHAR(64) PRIMARY KEY, -- 对外 ID (video_xxx / char_xxx / img_xxx) + account_id BIGINT NOT NULL, -- 使用的账号 ID + api_key_id BIGINT, -- 调用方 API Key ID + upstream_task_id VARCHAR(128) NOT NULL DEFAULT '', -- 上游任务 ID + object_type VARCHAR(16) NOT NULL DEFAULT 'video', -- video / character / image + model VARCHAR(64) NOT NULL, -- 请求模型名 + prompt TEXT NOT NULL DEFAULT '', -- 生成提示词 + status VARCHAR(16) NOT NULL DEFAULT 'queued', -- queued / in_progress / completed / failed + progress INT NOT NULL DEFAULT 0, -- 进度 0-100 + video_url TEXT NOT NULL DEFAULT '', -- 原始上游下载 URL + stored_key TEXT NOT NULL DEFAULT '', -- 存储后的 key(本地相对路径或 S3 object key) + storage_type VARCHAR(16) NOT NULL DEFAULT '', -- 存储类型:local / s3 / gdrive / 空=未存储 + share_id VARCHAR(128) NOT NULL DEFAULT '', -- 可分享 ID + character_info JSONB, -- 角色信息 {username, display_name} + error_message TEXT NOT NULL DEFAULT '', -- 错误信息 + error_type VARCHAR(64) NOT NULL DEFAULT '', -- 错误类型 + request_body JSONB, -- 原始请求体(用于上游重放) + seconds VARCHAR(8) NOT NULL DEFAULT '', -- 视频时长 + size VARCHAR(16) NOT NULL DEFAULT '', -- 分辨率 (如 1920x1080) + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + completed_at TIMESTAMPTZ +); + +-- 仅索引未完成任务,供后台 worker 轮询 +CREATE INDEX IF NOT EXISTS idx_sora_tasks_pending ON sora_tasks(status) WHERE status IN ('queued', 'in_progress'); + +-- 按 API Key 查询 +CREATE INDEX IF NOT EXISTS idx_sora_tasks_api_key ON sora_tasks(api_key_id) WHERE api_key_id IS NOT NULL; + +-- +migrate Down + +DROP TABLE IF EXISTS sora_tasks; diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index a0bc1a6061..e0e9e54f49 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -47,13 +47,15 @@ services: # ======================================================================= # Database Configuration (PostgreSQL) + # Default: uses local postgres container + # External DB: set DATABASE_HOST and DATABASE_SSLMODE in .env # ======================================================================= - - DATABASE_HOST=postgres - - DATABASE_PORT=5432 + - DATABASE_HOST=${DATABASE_HOST:-postgres} + - DATABASE_PORT=${DATABASE_PORT:-5432} - DATABASE_USER=${POSTGRES_USER:-sub2api} - DATABASE_PASSWORD=${POSTGRES_PASSWORD:?POSTGRES_PASSWORD is required} - DATABASE_DBNAME=${POSTGRES_DB:-sub2api} - - DATABASE_SSLMODE=disable + - DATABASE_SSLMODE=${DATABASE_SSLMODE:-disable} - DATABASE_MAX_OPEN_CONNS=${DATABASE_MAX_OPEN_CONNS:-50} - DATABASE_MAX_IDLE_CONNS=${DATABASE_MAX_IDLE_CONNS:-10} - DATABASE_CONN_MAX_LIFETIME_MINUTES=${DATABASE_CONN_MAX_LIFETIME_MINUTES:-30} @@ -139,8 +141,6 @@ services: # Examples: http://host:port, socks5://host:port - UPDATE_PROXY_URL=${UPDATE_PROXY_URL:-} depends_on: - postgres: - condition: service_healthy redis: condition: service_healthy networks: diff --git a/docs/ACCOUNT_SCHEDULING_FLOW.md b/docs/ACCOUNT_SCHEDULING_FLOW.md new file mode 100644 index 0000000000..3f7cc807ae --- /dev/null +++ b/docs/ACCOUNT_SCHEDULING_FLOW.md @@ -0,0 +1,61 @@ +# Account Scheduling Flow(SelectAccountWithLoadAwareness) + +本文档对应后端主调度入口 `SelectAccountWithLoadAwareness`,用于说明 Anthropic 网关的账号选择流程、亲和策略与关键过滤规则。 + +## 代码入口与引用 + +- 调度主入口:`backend/internal/service/gateway_service.go` 中的 `SelectAccountWithLoadAwareness` +- 亲和详情 API:`backend/internal/handler/admin/account_affinity_handler.go` 中的 `GetAffinityDetails` +- Admin 路由:`backend/internal/server/routes/admin.go`(`/api/v1/admin/accounts/:id/affinity-details`) + +## 流程概览(修正版) + +1. 解析输入参数 +`sessionHash`、`stickyAccountID`、`affinityClientID(metadata.user_id 提取)`、`affinityUserID(sub2api user id)`。 + +2. Claude Code 限制与分组降级 +先执行 `checkClaudeCodeRestriction`,必要时替换 `groupID`。 + +3. Legacy / Load-aware 分支 +若 `concurrencyService == nil || !LoadBatchEnabled` 走传统路径;否则走负载感知路径。 + +4. Layer 1(模型路由优先) +- 路由候选过滤(排除、平台、模型、配额、窗口费用、RPM) +- 路由范围内 sticky 优先 +- 路由内按(优先级 > 负载 > 亲和数 > LRU)尝试获取槽位 + +5. Layer 1.3(pinned_users 预处理) +仅做 `UpdateAffinity` 预热,不直接做调度决策。 + +6. Layer 1.4(客户端亲和调度) +按亲和记录尝试命中;`allow_switch=false` 且无一票放行时尝试等待计划,否则返回 `ErrAffinityNoSwitch`。 + +7. Layer 1.5(粘性会话) +仅在 `!affinityHit && routingAccountIDs==0` 时生效;先过 `shouldClearStickySession`。 + +8. Layer 2(负载感知选择) +分层过滤:优先级 -> 亲和区(单维客户端) -> 最低负载 -> 最少亲和客户端 -> LRU。 + +9. Layer 3(兜底排队) +按 `FallbackSelectionMode`(`last_used` 或 `random`)排序后返回等待计划。 +注意:此层不是“按最低负载”。 + +## 关键业务规则(2026-03) + +### 无客户端 ID 时过滤亲和 Anthropic 账号 + +当 `metadata.user_id` 无法提取 `client_id`(即 `affinityClientID == ""`)时: + +- 会直接过滤 **Anthropic 平台且开启客户端亲和** 的账号(覆盖 OAuth / SetupToken / API Key / Bedrock 等类型) +- 该规则只看平台与亲和开关,不再限定账号类型 + +设计目标:避免无客户端标识的请求误用客户端亲和账号。 + +## 亲和详情 API 契约 + +`GET /api/v1/admin/accounts/:id/affinity-details` 返回结构(后端已对齐前端): + +- `users[]`:包含 `user_id`、`user_email`、`client_count`、`is_pinned`、`clients[]` +- `total_users` +- `total_clients` +- `pinned_users` diff --git a/frontend/package.json b/frontend/package.json index 1b380b176b..d2a6deded7 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -16,6 +16,7 @@ }, "dependencies": { "@lobehub/icons": "^4.0.2", + "@tanstack/vue-virtual": "^3.13.23", "@vueuse/core": "^10.7.0", "axios": "^1.13.5", "chart.js": "^4.4.1", diff --git a/frontend/pnpm-lock.yaml b/frontend/pnpm-lock.yaml index 37c384b431..505b72f388 100644 --- a/frontend/pnpm-lock.yaml +++ b/frontend/pnpm-lock.yaml @@ -11,6 +11,9 @@ importers: '@lobehub/icons': specifier: ^4.0.2 version: 4.0.2(@lobehub/ui@4.9.2)(@types/react@19.2.7)(antd@6.1.3(react-dom@19.2.3(react@19.2.3))(react@19.2.3))(react-dom@19.2.3(react@19.2.3))(react@19.2.3) + '@tanstack/vue-virtual': + specifier: ^3.13.23 + version: 3.13.23(vue@3.5.26(typescript@5.6.3)) '@vueuse/core': specifier: ^10.7.0 version: 10.11.1(vue@3.5.26(typescript@5.6.3)) @@ -1376,6 +1379,14 @@ packages: peerDependencies: react: '>= 16.3.0' + '@tanstack/virtual-core@3.13.23': + resolution: {integrity: sha512-zSz2Z2HNyLjCplANTDyl3BcdQJc2k1+yyFoKhNRmCr7V7dY8o8q5m8uFTI1/Pg1kL+Hgrz6u3Xo6eFUB7l66cg==} + + '@tanstack/vue-virtual@3.13.23': + resolution: {integrity: sha512-b5jPluAR6U3eOq6GWAYSpj3ugnAIZgGR0e6aGAgyRse0Yu6MVQQ0ZWm9SArSXWtageogn6bkVD8D//c4IjW3xQ==} + peerDependencies: + vue: ^2.7.0 || ^3.0.0 + '@types/d3-array@3.2.2': resolution: {integrity: sha512-hOLWVbm7uRza0BYXpIIW5pxfrKe0W+D5lrFiAEYR+pb6w3N2SwSMaJbXdUfSEv+dT4MfHBLtn5js0LAWaO6otw==} @@ -5808,6 +5819,13 @@ snapshots: dependencies: react: 19.2.3 + '@tanstack/virtual-core@3.13.23': {} + + '@tanstack/vue-virtual@3.13.23(vue@3.5.26(typescript@5.6.3))': + dependencies: + '@tanstack/virtual-core': 3.13.23 + vue: 3.5.26(typescript@5.6.3) + '@types/d3-array@3.2.2': {} '@types/d3-axis@3.0.6': diff --git a/frontend/public/wechat-qr.jpg b/frontend/public/wechat-qr.jpg new file mode 100644 index 0000000000..659068d835 Binary files /dev/null and b/frontend/public/wechat-qr.jpg differ diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index 23d50d3acb..0c3b6f6713 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -17,7 +17,8 @@ import type { AdminDataPayload, AdminDataImportResult, CheckMixedChannelRequest, - CheckMixedChannelResponse + CheckMixedChannelResponse, + AffinityDetailsResponse } from '@/types' /** @@ -618,6 +619,30 @@ export async function batchRefresh(accountIds: number[]): Promise { + const { data } = await apiClient.get<{ client_id: string; last_active: string }[]>( + `/admin/accounts/${id}/affinity-clients` + ) + return data +} + +/** + * Get affinity details for an account with user-level grouping + * @param id - Account ID + * @returns Affinity details with user groups + */ +export async function getAffinityDetails(id: number): Promise { + const { data } = await apiClient.get( + `/admin/accounts/${id}/affinity-details` + ) + return data +} + export const accountsAPI = { list, listWithEtag, @@ -654,7 +679,9 @@ export const accountsAPI = { importData, getAntigravityDefaultModelMapping, batchClearError, - batchRefresh + batchRefresh, + getAffinityClients, + getAffinityDetails } export default accountsAPI diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index a2cd67f067..de937b0bf5 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -382,7 +382,7 @@ export async function updateBetaPolicySettings( return data } -// ==================== Sora S3 Settings ==================== +// ==================== Sora Storage Settings ==================== export interface SoraS3Settings { enabled: boolean @@ -401,7 +401,9 @@ export interface SoraS3Profile { profile_id: string name: string is_active: boolean + provider: string // "s3" | "gdrive" enabled: boolean + access_mode: string // "direct" | "proxy" endpoint: string region: string bucket: string @@ -412,6 +414,13 @@ export interface SoraS3Profile { cdn_url: string default_storage_quota_bytes: number updated_at: string + // Google Drive fields + auth_type: string // "oauth2" | "service_account" + client_id: string + client_secret_configured: boolean + refresh_token_configured: boolean + service_account_configured: boolean + folder_id: string } export interface ListSoraS3ProfilesResponse { @@ -437,30 +446,47 @@ export interface CreateSoraS3ProfileRequest { profile_id: string name: string set_active?: boolean + provider?: string enabled: boolean - endpoint: string - region: string - bucket: string - access_key_id: string + access_mode?: string + endpoint?: string + region?: string + bucket?: string + access_key_id?: string secret_access_key?: string - prefix: string - force_path_style: boolean - cdn_url: string - default_storage_quota_bytes: number + prefix?: string + force_path_style?: boolean + cdn_url?: string + default_storage_quota_bytes?: number + // Google Drive fields + auth_type?: string + client_id?: string + client_secret?: string + refresh_token?: string + service_account_json?: string + folder_id?: string } export interface UpdateSoraS3ProfileRequest { name: string enabled: boolean - endpoint: string - region: string - bucket: string - access_key_id: string + access_mode?: string + endpoint?: string + region?: string + bucket?: string + access_key_id?: string secret_access_key?: string - prefix: string - force_path_style: boolean - cdn_url: string - default_storage_quota_bytes: number + prefix?: string + force_path_style?: boolean + cdn_url?: string + default_storage_quota_bytes?: number + // Google Drive fields + auth_type?: string + client_id?: string + client_secret?: string + refresh_token?: string + service_account_json?: string + folder_id?: string } export interface TestSoraS3ConnectionRequest { @@ -477,44 +503,115 @@ export interface TestSoraS3ConnectionRequest { default_storage_quota_bytes?: number } +export interface StartGDriveOAuthRequest { + client_id: string + client_secret: string + redirect_uri: string +} + +export interface StartGDriveOAuthResponse { + auth_url: string + state: string +} + +export interface ExchangeGDriveOAuthCodeRequest { + client_id: string + client_secret: string + redirect_uri: string + code: string + profile_id?: string +} + +export interface ExchangeGDriveOAuthCodeResponse { + refresh_token: string + message: string +} + export async function getSoraS3Settings(): Promise { - const { data } = await apiClient.get('/admin/settings/sora-s3') + const { data } = await apiClient.get('/admin/settings/sora-storage') return data } export async function updateSoraS3Settings(settings: UpdateSoraS3SettingsRequest): Promise { - const { data } = await apiClient.put('/admin/settings/sora-s3', settings) + const { data } = await apiClient.put('/admin/settings/sora-storage', settings) return data } export async function testSoraS3Connection( settings: TestSoraS3ConnectionRequest ): Promise<{ message: string }> { - const { data } = await apiClient.post<{ message: string }>('/admin/settings/sora-s3/test', settings) + const { data } = await apiClient.post<{ message: string }>('/admin/settings/sora-storage/test', settings) return data } export async function listSoraS3Profiles(): Promise { - const { data } = await apiClient.get('/admin/settings/sora-s3/profiles') + const { data } = await apiClient.get('/admin/settings/sora-storage/profiles') return data } export async function createSoraS3Profile(request: CreateSoraS3ProfileRequest): Promise { - const { data } = await apiClient.post('/admin/settings/sora-s3/profiles', request) + const { data } = await apiClient.post('/admin/settings/sora-storage/profiles', request) return data } export async function updateSoraS3Profile(profileID: string, request: UpdateSoraS3ProfileRequest): Promise { - const { data } = await apiClient.put(`/admin/settings/sora-s3/profiles/${profileID}`, request) + const { data } = await apiClient.put(`/admin/settings/sora-storage/profiles/${profileID}`, request) return data } export async function deleteSoraS3Profile(profileID: string): Promise { - await apiClient.delete(`/admin/settings/sora-s3/profiles/${profileID}`) + await apiClient.delete(`/admin/settings/sora-storage/profiles/${profileID}`) } export async function setActiveSoraS3Profile(profileID: string): Promise { - const { data } = await apiClient.post(`/admin/settings/sora-s3/profiles/${profileID}/activate`) + const { data } = await apiClient.post(`/admin/settings/sora-storage/profiles/${profileID}/activate`) + return data +} + +export async function startGDriveOAuth(request: StartGDriveOAuthRequest): Promise { + const { data } = await apiClient.post('/admin/settings/sora-storage/gdrive-oauth/start', request) + return data +} + +export async function exchangeGDriveOAuthCode(request: ExchangeGDriveOAuthCodeRequest): Promise { + const { data } = await apiClient.post('/admin/settings/sora-storage/gdrive-oauth/callback', request) + return data +} + +export interface TestGDriveStorageResponse { + status: string + quota_limit_bytes?: number + quota_used_bytes?: number + uploaded_file_id?: string + uploaded_file_name?: string + uploaded_file_size?: number + access_url?: string + web_view_link?: string + deleted?: boolean + delete_warning?: string +} + +export async function testGDriveStorage(): Promise { + const { data } = await apiClient.post('/admin/settings/sora-storage/gdrive-test') + return data +} + +export interface GDriveQuotaInfo { + limit_bytes: number + used_bytes: number +} + +export interface StorageVideoStats { + [type: string]: { completed: number; in_progress: number } +} + +export async function getGDriveQuota(): Promise { + const { data } = await apiClient.get('/admin/settings/sora-storage/gdrive-quota') + return data +} + +export async function getStorageVideoStats(): Promise { + const { data } = await apiClient.get('/admin/settings/sora-storage/video-stats') return data } @@ -541,7 +638,12 @@ export const settingsAPI = { createSoraS3Profile, updateSoraS3Profile, deleteSoraS3Profile, - setActiveSoraS3Profile + setActiveSoraS3Profile, + startGDriveOAuth, + exchangeGDriveOAuthCode, + testGDriveStorage, + getGDriveQuota, + getStorageVideoStats } export default settingsAPI diff --git a/frontend/src/components/account/AccountCapacityCell.vue b/frontend/src/components/account/AccountCapacityCell.vue index f8fe4b4795..aa926d5eab 100644 --- a/frontend/src/components/account/AccountCapacityCell.vue +++ b/frontend/src/components/account/AccountCapacityCell.vue @@ -1,76 +1,35 @@ - + diff --git a/frontend/src/components/account/AffinityConfigCard.vue b/frontend/src/components/account/AffinityConfigCard.vue new file mode 100644 index 0000000000..fb458242e3 --- /dev/null +++ b/frontend/src/components/account/AffinityConfigCard.vue @@ -0,0 +1,528 @@ + + + diff --git a/frontend/src/components/account/BulkEditAccountModal.vue b/frontend/src/components/account/BulkEditAccountModal.vue index 64524d519f..b1ca533031 100644 --- a/frontend/src/components/account/BulkEditAccountModal.vue +++ b/frontend/src/components/account/BulkEditAccountModal.vue @@ -5,7 +5,7 @@ width="wide" @close="handleClose" > -
+

@@ -599,6 +599,48 @@

+ +
+
+
+ +

+ {{ t('admin.accounts.allowOveragesTooltip') }} +

+
+ +
+
+ +
+
+
@@ -843,6 +885,11 @@ const appStore = useAppStore() // Platform awareness const isMixedPlatform = computed(() => props.selectedPlatforms.length > 1) +// 是否全部为 Antigravity 平台(allow_overages 仅在此条件下显示) +const allAntigravity = computed(() => + props.selectedPlatforms.length === 1 && props.selectedPlatforms[0] === 'antigravity' +) + // 是否全部为 Anthropic OAuth/SetupToken(RPM 配置仅在此条件下显示) const allAnthropicOAuthOrSetupToken = computed(() => { return ( @@ -887,6 +934,7 @@ const enableRateMultiplier = ref(false) const enableStatus = ref(false) const enableGroups = ref(false) const enableRpmLimit = ref(false) +const enableAllowOverages = ref(false) // State - field values const submitting = ref(false) @@ -912,6 +960,7 @@ const bulkBaseRpm = ref(null) const bulkRpmStrategy = ref<'tiered' | 'sticky_exempt'>('tiered') const bulkRpmStickyBuffer = ref(null) const userMsgQueueMode = ref(null) +const allowOverages = ref(false) const umqModeOptions = computed(() => [ { value: '', label: t('admin.accounts.quotaControl.rpmLimit.umqModeOff') }, { value: 'throttle', label: t('admin.accounts.quotaControl.rpmLimit.umqModeThrottle') }, @@ -1120,6 +1169,13 @@ const buildUpdatePayload = (): Record | null => { umqExtra.user_msg_queue_enabled = false // 清理旧字段(JSONB merge) } + // Allow overages (Antigravity only) + if (enableAllowOverages.value) { + if (!updates.extra) updates.extra = {} + const overagesExtra = updates.extra as Record + overagesExtra.allow_overages = allowOverages.value + } + return Object.keys(updates).length > 0 ? updates : null } @@ -1182,6 +1238,7 @@ const handleSubmit = async () => { enableStatus.value || enableGroups.value || enableRpmLimit.value || + enableAllowOverages.value || userMsgQueueMode.value !== null if (!hasAnyFieldEnabled) { @@ -1273,6 +1330,7 @@ watch( enableStatus.value = false enableGroups.value = false enableRpmLimit.value = false + enableAllowOverages.value = false // Reset all values baseUrl.value = '' @@ -1294,6 +1352,7 @@ watch( bulkRpmStrategy.value = 'tiered' bulkRpmStickyBuffer.value = null userMsgQueueMode.value = null + allowOverages.value = false // Reset mixed channel warning state showMixedChannelWarning.value = false diff --git a/frontend/src/components/account/CapacityBadge.vue b/frontend/src/components/account/CapacityBadge.vue new file mode 100644 index 0000000000..0abbafdd4b --- /dev/null +++ b/frontend/src/components/account/CapacityBadge.vue @@ -0,0 +1,25 @@ + + + diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 6f02a9d991..ff6e57e35e 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -1556,10 +1556,64 @@
- -
+ +
-

{{ t('admin.accounts.quotaLimit') }}

+

{{ t('admin.accounts.quotaControl.title') }}

+

+ {{ t('admin.accounts.quotaControl.hint') }} +

+
+ + +
+ + +
+
+

{{ t('admin.accounts.quotaControl.title') }}

{{ t('admin.accounts.quotaLimitHint') }}

@@ -1902,7 +1956,7 @@
- +
+ +
@@ -2868,6 +2941,7 @@ import ProxySelector from '@/components/common/ProxySelector.vue' import GroupSelector from '@/components/common/GroupSelector.vue' import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.vue' import QuotaLimitCard from '@/components/account/QuotaLimitCard.vue' +import AffinityConfigCard from '@/components/account/AffinityConfigCard.vue' import { applyInterceptWarmup } from '@/components/account/credentialsBuilder' import { formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format' import { createStableObjectKeyResolver } from '@/utils/stableObjectKey' @@ -3062,6 +3136,16 @@ const antigravityMixedChannelConfirmed = ref(false) const showAdvancedOAuth = ref(false) const showGeminiHelpDialog = ref(false) +// Client affinity (all Anthropic accounts) +const clientAffinityEnabled = ref(false) +const affinityBase = ref(null) +const affinityBuffer = ref(null) +const affinityAllowSwitch = ref(true) +const affinityUserBase = ref(null) +const affinityUserBuffer = ref(null) +const perUserClientLimit = ref(null) +const pinnedUsers = ref([]) + // Quota control state (Anthropic OAuth/SetupToken only) const windowCostEnabled = ref(false) const windowCostLimit = ref(null) @@ -3711,6 +3795,14 @@ const resetForm = () => { editWeeklyResetDay.value = null editWeeklyResetHour.value = null editResetTimezone.value = null + clientAffinityEnabled.value = false + affinityBase.value = null + affinityBuffer.value = null + affinityAllowSwitch.value = true + affinityUserBase.value = null + affinityUserBuffer.value = null + perUserClientLimit.value = null + pinnedUsers.value = [] modelMappings.value = [] modelRestrictionMode.value = 'whitelist' allowedModels.value = [...claudeModels] // Default fill related models @@ -3817,10 +3909,61 @@ const buildAnthropicExtra = (base?: Record): Record 0 ? extra : undefined } +/** 将客户端亲和设置写入 extra(Anthropic 全类型通用) */ +const applyClientAffinity = (extra: Record) => { + if (clientAffinityEnabled.value) { + extra.affinity_enabled = true + extra.client_affinity_enabled = true + if (affinityBase.value != null && affinityBase.value > 0) { + extra.affinity_base = affinityBase.value + } else { + delete extra.affinity_base + } + if (affinityBase.value != null && affinityBase.value > 0 && affinityBuffer.value != null) { + extra.affinity_buffer = affinityBuffer.value + } else { + delete extra.affinity_buffer + } + // v2 fields + extra.affinity_allow_switch = affinityAllowSwitch.value + if (affinityUserBase.value != null && affinityUserBase.value > 0) { + extra.affinity_user_base = affinityUserBase.value + } else { + delete extra.affinity_user_base + } + if (affinityUserBase.value != null && affinityUserBase.value > 0 && affinityUserBuffer.value != null) { + extra.affinity_user_buffer = affinityUserBuffer.value + } else { + delete extra.affinity_user_buffer + } + if (perUserClientLimit.value != null && perUserClientLimit.value > 0) { + extra.per_user_client_limit = perUserClientLimit.value + } else { + delete extra.per_user_client_limit + } + if (pinnedUsers.value.length > 0) { + extra.pinned_users = pinnedUsers.value + } else { + delete extra.pinned_users + } + } else { + delete extra.affinity_enabled + delete extra.client_affinity_enabled + delete extra.affinity_base + delete extra.affinity_buffer + delete extra.affinity_allow_switch + delete extra.affinity_user_base + delete extra.affinity_user_buffer + delete extra.per_user_client_limit + delete extra.pinned_users + } +} + const buildSoraExtra = ( base?: Record, linkedOpenAIAccountId?: string | number @@ -4822,6 +4965,9 @@ const handleAnthropicExchange = async (authCode: string) => { extra.cache_ttl_override_target = cacheTTLOverrideTarget.value } + // Add client affinity setting + applyClientAffinity(extra) + const credentials: Record = { ...tokenInfo } applyInterceptWarmup(credentials, interceptWarmupRequests.value, 'create') await createAccountAndFinish(form.platform, addMethod.value as AccountType, credentials, extra) diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 5f3da1b7c1..85532eed58 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1027,6 +1027,7 @@
+
@@ -1149,10 +1150,63 @@
- -
+ +
-

{{ t('admin.accounts.quotaLimit') }}

+

{{ t('admin.accounts.quotaControl.title') }}

+

+ {{ t('admin.accounts.quotaControl.hint') }} +

+
+ + +
+ +
+
+

{{ t('admin.accounts.quotaControl.title') }}

{{ t('admin.accounts.quotaLimitHint') }}

@@ -1237,7 +1291,7 @@
- +
+ +
@@ -1717,6 +1790,7 @@ import ProxySelector from '@/components/common/ProxySelector.vue' import GroupSelector from '@/components/common/GroupSelector.vue' import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.vue' import QuotaLimitCard from '@/components/account/QuotaLimitCard.vue' +import AffinityConfigCard from '@/components/account/AffinityConfigCard.vue' import { applyInterceptWarmup } from '@/components/account/credentialsBuilder' import { formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format' import { createStableObjectKeyResolver } from '@/utils/stableObjectKey' @@ -1844,6 +1918,14 @@ const tlsFingerprintEnabled = ref(false) const sessionIdMaskingEnabled = ref(false) const cacheTTLOverrideEnabled = ref(false) const cacheTTLOverrideTarget = ref('5m') +const clientAffinityEnabled = ref(false) +const affinityBase = ref(null) +const affinityBuffer = ref(null) +const affinityAllowSwitch = ref(true) +const affinityUserBase = ref(null) +const affinityUserBuffer = ref(null) +const perUserClientLimit = ref(null) +const pinnedUsers = ref([]) // OpenAI 自动透传开关(OAuth/API Key) const openaiPassthroughEnabled = ref(false) @@ -2461,9 +2543,47 @@ function loadQuotaControlSettings(account: Account) { sessionIdMaskingEnabled.value = false cacheTTLOverrideEnabled.value = false cacheTTLOverrideTarget.value = '5m' + clientAffinityEnabled.value = false + affinityBase.value = null + affinityBuffer.value = null + affinityAllowSwitch.value = true + affinityUserBase.value = null + affinityUserBuffer.value = null + perUserClientLimit.value = null + pinnedUsers.value = [] - // Only applies to Anthropic OAuth/SetupToken accounts - if (account.platform !== 'anthropic' || (account.type !== 'oauth' && account.type !== 'setup-token')) { + // Remaining quota control settings only apply to Anthropic accounts + if (account.platform !== 'anthropic') { + return + } + + // Client affinity (all Anthropic accounts) + if (account.client_affinity_enabled === true) { + clientAffinityEnabled.value = true + } + // Affinity base/buffer from extra + const extra = account.extra as Record | undefined + if (extra) { + const base = extra.affinity_base + affinityBase.value = (typeof base === 'number' && base > 0) ? base : null + // buffer: null = infinite yellow, 0 = no yellow, >0 = yellow range + const buf = extra.affinity_buffer + affinityBuffer.value = (typeof buf === 'number') ? buf : null + + // New v2 fields + affinityAllowSwitch.value = extra.affinity_allow_switch !== false + const ub = extra.affinity_user_base + affinityUserBase.value = (typeof ub === 'number' && ub > 0) ? ub : null + const ubuf = extra.affinity_user_buffer + affinityUserBuffer.value = (typeof ubuf === 'number') ? ubuf : null + const pul = extra.per_user_client_limit + perUserClientLimit.value = (typeof pul === 'number' && pul > 0) ? pul : null + const pu = extra.pinned_users + pinnedUsers.value = Array.isArray(pu) ? pu : [] + } + + // Window cost / session limit only apply to Anthropic OAuth/SetupToken accounts + if (account.type !== 'oauth' && account.type !== 'setup-token') { return } @@ -2880,9 +3000,63 @@ const handleSubmit = async () => { updatePayload.extra = newExtra } + // For all Anthropic accounts, handle client_affinity in extra + if (props.account.platform === 'anthropic') { + const currentExtra = (props.account.extra as Record) || {} + const newExtra: Record = { ...currentExtra } + if (clientAffinityEnabled.value) { + newExtra.affinity_enabled = true + newExtra.client_affinity_enabled = true + if (affinityBase.value != null && affinityBase.value > 0) { + newExtra.affinity_base = affinityBase.value + } else { + delete newExtra.affinity_base + } + // buffer: null = infinite yellow, 0 = no yellow, >0 = yellow range + if (affinityBase.value != null && affinityBase.value > 0 && affinityBuffer.value != null) { + newExtra.affinity_buffer = affinityBuffer.value + } else { + delete newExtra.affinity_buffer + } + // v2 fields + newExtra.affinity_allow_switch = affinityAllowSwitch.value + if (affinityUserBase.value != null && affinityUserBase.value > 0) { + newExtra.affinity_user_base = affinityUserBase.value + } else { + delete newExtra.affinity_user_base + } + if (affinityUserBase.value != null && affinityUserBase.value > 0 && affinityUserBuffer.value != null) { + newExtra.affinity_user_buffer = affinityUserBuffer.value + } else { + delete newExtra.affinity_user_buffer + } + if (perUserClientLimit.value != null && perUserClientLimit.value > 0) { + newExtra.per_user_client_limit = perUserClientLimit.value + } else { + delete newExtra.per_user_client_limit + } + if (pinnedUsers.value.length > 0) { + newExtra.pinned_users = pinnedUsers.value + } else { + delete newExtra.pinned_users + } + } else { + newExtra.affinity_enabled = false + newExtra.client_affinity_enabled = false + delete newExtra.affinity_base + delete newExtra.affinity_buffer + delete newExtra.affinity_allow_switch + delete newExtra.affinity_user_base + delete newExtra.affinity_user_buffer + delete newExtra.per_user_client_limit + delete newExtra.pinned_users + } + updatePayload.extra = newExtra + } + // For Anthropic OAuth/SetupToken accounts, handle quota control settings in extra if (props.account.platform === 'anthropic' && (props.account.type === 'oauth' || props.account.type === 'setup-token')) { - const currentExtra = (props.account.extra as Record) || {} + const currentExtra = (updatePayload.extra as Record) || (props.account.extra as Record) || {} const newExtra: Record = { ...currentExtra } // Window cost limit settings @@ -2957,7 +3131,7 @@ const handleSubmit = async () => { // For Anthropic API Key accounts, handle passthrough mode in extra if (props.account.platform === 'anthropic' && props.account.type === 'apikey') { - const currentExtra = (props.account.extra as Record) || {} + const currentExtra = (updatePayload.extra as Record) || (props.account.extra as Record) || {} const newExtra: Record = { ...currentExtra } if (anthropicPassthroughEnabled.value) { newExtra.anthropic_passthrough = true @@ -3007,20 +3181,27 @@ const handleSubmit = async () => { const currentExtra = (updatePayload.extra as Record) || (props.account.extra as Record) || {} const newExtra: Record = { ...currentExtra } + // Total quota if (editQuotaLimit.value != null && editQuotaLimit.value > 0) { newExtra.quota_limit = editQuotaLimit.value } else { delete newExtra.quota_limit } + // Daily quota if (editQuotaDailyLimit.value != null && editQuotaDailyLimit.value > 0) { newExtra.quota_daily_limit = editQuotaDailyLimit.value } else { delete newExtra.quota_daily_limit + delete newExtra.quota_daily_used + delete newExtra.quota_daily_start } + // Weekly quota if (editQuotaWeeklyLimit.value != null && editQuotaWeeklyLimit.value > 0) { newExtra.quota_weekly_limit = editQuotaWeeklyLimit.value } else { delete newExtra.quota_weekly_limit + delete newExtra.quota_weekly_used + delete newExtra.quota_weekly_start } // Quota reset mode config if (editDailyResetMode.value === 'fixed') { diff --git a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts index 9158da64ed..7a739d442b 100644 --- a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts +++ b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts @@ -15,6 +15,10 @@ vi.mock('@/api/admin', () => ({ } })) +vi.mock('@/utils/usageLoadQueue', () => ({ + enqueueUsageRequest: (_account: unknown, fn: () => Promise) => fn() +})) + vi.mock('vue-i18n', async () => { const actual = await vi.importActual('vue-i18n') return { @@ -385,6 +389,117 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).toContain('7d|0|27700') }) + it('OpenAI OAuth 在 usage 请求失败时仍回退显示本地 codex 快照', async () => { + getUsage.mockRejectedValue(new Error('network error')) + const errorSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ + id: 2004, + platform: 'openai', + type: 'oauth', + extra: { + codex_usage_updated_at: '2099-03-07T10:00:00Z', + codex_5h_used_percent: 12, + codex_5h_reset_at: '2099-03-07T12:00:00Z', + codex_7d_used_percent: 34, + codex_7d_reset_at: '2099-03-13T12:00:00Z' + } + }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt', 'windowStats', 'color'], + template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}
' + }, + AccountQuotaInfo: true + } + } + }) + + await flushPromises() + + expect(getUsage).toHaveBeenCalledWith(2004) + expect(wrapper.text()).toContain('5h|12|2099-03-07T12:00:00.000Z') + expect(wrapper.text()).toContain('7d|34|2099-03-13T12:00:00.000Z') + errorSpy.mockRestore() + }) + + it('OpenAI OAuth 已限额时首屏优先等待重新查询的 usage,而不是先显示旧 codex 快照', async () => { + let resolveUsage: ((value: any) => void) | null = null + getUsage.mockReturnValue( + new Promise((resolve) => { + resolveUsage = resolve + }) + ) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ + id: 2005, + platform: 'openai', + type: 'oauth', + rate_limit_reset_at: '2099-03-07T12:00:00Z', + extra: { + codex_5h_used_percent: 0, + codex_5h_reset_at: '2099-03-07T12:00:00Z', + codex_7d_used_percent: 0, + codex_7d_reset_at: '2099-03-13T12:00:00Z' + } + }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt', 'windowStats', 'color'], + template: '
{{ label }}|{{ utilization }}|{{ windowStats?.tokens }}
' + }, + AccountQuotaInfo: true + } + } + }) + + await Promise.resolve() + + expect(getUsage).toHaveBeenCalledWith(2005) + expect(wrapper.text()).not.toContain('5h|0|') + expect(wrapper.text()).not.toContain('7d|0|') + + resolveUsage?.({ + five_hour: { + utilization: 100, + resets_at: '2026-03-07T12:00:00Z', + remaining_seconds: 3600, + window_stats: { + requests: 211, + tokens: 106540000, + cost: 38.13, + standard_cost: 38.13, + user_cost: 38.13 + } + }, + seven_day: { + utilization: 100, + resets_at: '2026-03-13T12:00:00Z', + remaining_seconds: 3600, + window_stats: { + requests: 211, + tokens: 106540000, + cost: 38.13, + standard_cost: 38.13, + user_cost: 38.13 + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('5h|100|106540000') + expect(wrapper.text()).toContain('7d|100|106540000') + }) + it('OpenAI OAuth 在行数据刷新但仍无 codex 快照时会重新拉取 usage', async () => { getUsage .mockResolvedValueOnce({ diff --git a/frontend/src/components/account/__tests__/AffinityBadge.spec.ts b/frontend/src/components/account/__tests__/AffinityBadge.spec.ts new file mode 100644 index 0000000000..e7561e7f76 --- /dev/null +++ b/frontend/src/components/account/__tests__/AffinityBadge.spec.ts @@ -0,0 +1,163 @@ +import { describe, expect, it, vi } from 'vitest' +import { mount } from '@vue/test-utils' +import AffinityBadge from '../AffinityBadge.vue' + +vi.mock('@/api/admin/accounts', () => ({ + getAffinityClients: vi.fn(), + getAffinityDetails: vi.fn() +})) + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ + t: (key: string, fallback?: string) => fallback ?? key + }) + } +}) + +function mountBadge(opts: { + clientCount: number + base?: number + buffer?: number | null + userCount?: number + userBase?: number + userBuffer?: number | null +}) { + return mount(AffinityBadge, { + props: { + accountId: 42, + clientCount: opts.clientCount, + base: opts.base ?? 5, + buffer: opts.buffer ?? 10, + userCount: opts.userCount ?? 0, + userBase: opts.userBase ?? 0, + userBuffer: opts.userBuffer ?? null + } + }) +} + +describe('AffinityBadge', () => { + // ====== Original tests (client dimension only) ====== + it('renders the correct client count number', () => { + const wrapper = mountBadge({ clientCount: 5 }) + expect(wrapper.text()).toContain('5') + }) + + it('renders configured limit text', () => { + const wrapper = mountBadge({ clientCount: 5, base: 5, buffer: 10 }) + expect(wrapper.text()).toContain('15') + }) + + it('renders infinity limit text when base is not configured', () => { + const wrapper = mountBadge({ clientCount: 5, base: 0, buffer: null }) + expect(wrapper.text()).toContain('\u221E') + }) + + it('applies red badge class when count exceeds base plus buffer', () => { + const wrapper = mountBadge({ clientCount: 16, base: 5, buffer: 10 }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-red-100') + expect(badge.classes()).toContain('text-red-700') + }) + + it('applies yellow badge class when count is in buffer range', () => { + const wrapper = mountBadge({ clientCount: 6, base: 5, buffer: 10 }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-yellow-100') + expect(badge.classes()).toContain('text-yellow-700') + }) + + it('applies yellow badge class when buffer is infinite', () => { + const wrapper = mountBadge({ clientCount: 6, base: 5, buffer: null }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-yellow-100') + expect(badge.classes()).toContain('text-yellow-700') + }) + + it('applies green badge class when count is within base', () => { + const wrapper = mountBadge({ clientCount: 5, base: 5, buffer: 10 }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-emerald-100') + expect(badge.classes()).toContain('text-emerald-700') + }) + + it('applies gray badge class when count is 0', () => { + const wrapper = mountBadge({ clientCount: 0, base: 5, buffer: 10 }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-gray-100') + expect(badge.classes()).toContain('text-gray-600') + }) + + it('does NOT show popover initially', () => { + const wrapper = mountBadge({ clientCount: 3 }) + expect(wrapper.html()).not.toContain('divide-y') + }) + + it('has mouseenter and mouseleave handlers on the badge', () => { + const wrapper = mountBadge({ clientCount: 3 }) + const badge = wrapper.find('span') + expect(badge.exists()).toBe(true) + badge.trigger('mouseenter') + badge.trigger('mouseleave') + }) + + // ====== Dual dimension tests ====== + it('dual dimension: user green + client yellow = yellow', () => { + const wrapper = mountBadge({ + clientCount: 6, base: 5, buffer: 10, + userCount: 3, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-yellow-100') + expect(badge.classes()).toContain('text-yellow-700') + }) + + it('dual dimension: user green + client red = red', () => { + const wrapper = mountBadge({ + clientCount: 16, base: 5, buffer: 10, + userCount: 3, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-red-100') + expect(badge.classes()).toContain('text-red-700') + }) + + it('dual dimension: user red + client green = red', () => { + const wrapper = mountBadge({ + clientCount: 3, base: 5, buffer: 10, + userCount: 11, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-red-100') + expect(badge.classes()).toContain('text-red-700') + }) + + it('user only dimension shows user count text', () => { + const wrapper = mountBadge({ + clientCount: 0, base: 0, buffer: null, + userCount: 3, userBase: 5, userBuffer: 5 + }) + expect(wrapper.text()).toContain('3') + expect(wrapper.text()).toContain('10') + }) + + it('client only dimension shows client count text', () => { + const wrapper = mountBadge({ + clientCount: 4, base: 5, buffer: 10, + userCount: 0, userBase: 0, userBuffer: null + }) + expect(wrapper.text()).toContain('4') + expect(wrapper.text()).toContain('15') + }) + + it('dual dimension limit text shows both U and C prefixes', () => { + const wrapper = mountBadge({ + clientCount: 3, base: 5, buffer: 10, + userCount: 2, userBase: 4, userBuffer: 3 + }) + expect(wrapper.text()).toContain('U2/7') + expect(wrapper.text()).toContain('C3/15') + }) +}) diff --git a/frontend/src/components/common/DataTable.vue b/frontend/src/components/common/DataTable.vue index 16aea10739..57f800a034 100644 --- a/frontend/src/components/common/DataTable.vue +++ b/frontend/src/components/common/DataTable.vue @@ -33,7 +33,7 @@