16 Commits
Author SHA1 Message Date
Sakurasan ca4dc4b3b7 fix: 缓存 token 计费按上游协议区分口径
- 新增 ComputeCost:OpenAI 系 prompt 含缓存读需扣减(异常数据钳制为 0);
  Anthropic input_tokens 不含缓存,按原值计费,缓存写按输入价 ×1.25(原实现误用输出价)
- recordUsage 传入上游协议(用量语义跟随解析它的上游响应,而非客户端协议)
- 补充单元测试覆盖两种协议口径与边界情况
2026-09-02 02:36:28 +08:00
Sakurasan 0628d5050f feat: 用量统计页改版,月度汇总支持多指标图表
- 新增 GET /api/usage/monthly?year= 年度按自然月聚合,每月含按模型分解(token 降序)
- 月度汇总改为堆叠柱状图:Token/消费金额/调用次数三指标切换,按模型分色(图例取前 8,其余归入「其他」)
- 选中月份概览卡片(金额/次数/token 分解),点击柱体或图例联动切换
- 年份切换、悬停明细 tooltip、请求明细保留
2026-09-02 02:36:28 +08:00
Sakurasan a376ac0722 feat: add Redis support for distributed passkey sessions
- Implement SessionStore interface for pluggable storage backends
- Add redisStore for distributed environment
- Add memoryStore as fallback for single instance
- Update router to initialize Redis client when configured
- Update .env.example with Redis documentation
2026-09-01 23:03:39 +08:00
Sakurasan 9733b3c20b refactor: restructure passkey with dedicated package
- Create internal/passkey/ package with challenge session management
- Add HTTP handler layer in internal/api/passkey.go
- Simplify Passkey model to use blob credential storage
- Register passkey routes in router
- Update frontend to match new API endpoints
- Add .env.example template
2026-09-01 22:34:44 +08:00
Sakurasan d28fca8ee6 chore: sqlite 默认数据库文件调整为 db/openteam.db,运行时自动创建目录 2026-09-01 21:36:15 +08:00
Sakurasan a2cef00908 fix: 未登录访问管理后台不再闪现页面,路由守卫改为服务端校验
- 守卫先经 /profile 校验 token 有效后才放行受保护路由,无效 token 直接跳登录并带 redirect 回跳
- DashboardLayout 用户信息就绪前只渲染 spinner,不渲染后台内容
- manager 路由增加 requiresAdmin 校验(role >= 10),非管理员重定向回仪表盘
- 401 拦截器改用 router.push 软跳转,避免整页刷新与双重跳转
- pinia 先于 router 安装;登录页支持 redirect 参数回跳原目标
2026-09-01 20:20:05 +08:00
Sakurasan d41bcdc371 修复backend dist文件路径 2026-09-01 17:22:26 +08:00
Sakurasan 8d949eff18 修改图标配色 2026-09-01 15:38:32 +08:00
Sakurasan c65d497551 feat: add docker build commands to Makefile 2026-09-01 15:22:37 +08:00
Sakurasan 19232567f2 chore: 忽略 paseo 任务运行时目录 .pi/ 2026-09-01 03:41:18 +08:00
Sakurasan 20f365d11a fix: 识别上游 HTTP 200 错误体,正确记账并返回 502
问题:OpenRouter 等上游超时时返回 HTTP 200 但响应体/SSE 内含
error 字段(如 {"error":{"code":504}}),网关只检查 StatusCode>=400,
导致:1) 客户端收到假 200 + 空内容;2) 用量被错误记为 success。

修复:
- bufferResponse 检测 JSON 错误体(openai 风格 error / anthropic 风格 type=error),
  识别时以 502 返回给客户端,并向 Dispatch 返回错误码
- streamResponse 扫描 SSE data 载荷中的 error 字段,识别时返回错误码
- Dispatch 对带错误码的响应按失败记账(status=error, error_code=upstream_error),
  不产生费用
2026-09-01 03:31:55 +08:00
Sakurasan 104ecc4691 fix: 原始数据弹窗改为标签页切换,手机友好
- 原始请求/响应不再上下堆叠,改为「请求 / 响应」两个标签页
- 单个内容区(pre 内滚动 + 自动换行),避免手机上需要长距离滚动
- 打开时默认停在有内容的标签(请求优先);仅有一个数据时只显示对应标签
- 弹窗标题补充模型名与协议信息
2026-09-01 03:16:22 +08:00
Sakurasan a259e5eb4d feat: 普通用户用量统计 + 管理后台用量明细
后端
- /api/usage/stats:当前用户按日聚合统计(请求数/输入/输出/缓存 tokens/费用)
- /api/usage/logs:当前用户用量明细分页
- /api/admin/usage/logs:全量明细(分页 + 协议/状态/模型/用户筛选 + 用户名关联)
- /api/admin/usage/summary:全量汇总(按用户分组)
- dao:UsageFilter 支持筛选分页;DailyUsageDAO.ListAll
- 路由加固:新增 middleware.AdminOnly,既有 /api/admin/* 全部迁移到
  带管理员角色校验的分组(此前登录即可访问,属安全隐患)

前端
- 普通用户「用量统计」页:统计卡片 + 纯 CSS 每日请求量条形图 + 明细表格分页
- 管理后台「用量明细」页:汇总卡片 + 协议/状态/模型/用户筛选 +
  明细表格 + 原始请求/响应查看弹窗(记录开关开启时)
- stores/usage.ts 与类型定义;控制台/管理菜单挂载
2026-09-01 02:55:55 +08:00
Sakurasan f9e9a1572f feat: 系统设置新增原始请求/响应记录开关(仅管理员)
- 系统配置键 log_raw_requests:开启后,仅管理员账号的每次请求
  在用量明细中保存客户端原始请求体与上游原始响应体
  (流式含全部 SSE 事件),用于排障
- UsageLog 新增 raw_request / raw_response 字段(type:text)
- AuthLLM 附带 user_role 供网关判断管理员
- gateway:10s TTL 缓存开关;streamResponse/bufferResponse
  支持累积上游原始响应;recordUsage 填充原始字段
- 前端 SystemConfig 新增开关(会显著增加存储的提示)
- 新增 doc/flow.md 网关调用流程示意图
2026-09-01 02:32:31 +08:00
Sakurasan 9f4d631fc4 feat: 网关路由对齐参考实现 + 用量落库 + e2e 全绿
路由与故障转移(参考 openteam 语义)
- channel.Candidates:绑定模型优先(携带 upstream_model 映射),
  未绑定模型回退到权重最低的健康备用渠道;新增 Pick 加权随机与 FilterHealthy 内存健康过滤
- gateway.Dispatch:遍历候选渠道,可重试失败(连接错误/429/5xx)自动故障转移,
  4xx 透传;不再使用单一 SelectChannel
- 修复 gorm default 标签把渠道 weight=0 静默改写为 1 的问题(去掉 default,
  权重 0 语义 = 不参与加权选择,仅作备用承接 unbound 流量)
- RecordFailure 连续 2 次进入 degraded 快速熔断,健康检查成功或冷却过期后复位

网关功能补全
- /v1/models 返回 DB 中启用的模型列表(替换 TODO 存根)
- 请求级 request_id 生成与用量记录接入:流式 SSE 逐块累计 usage、
  非流式从响应提取,按模型定价计算成本后经 usage.Recorder 异步落库
- 流式结束检测:chat 的 [DONE]、messages 的 message_stop、responses 的
  response.completed,避免 keep-alive 上游发完不关连接导致读阻塞到超时
- ResponsesRequest.input 兼容字符串与条目数组两种客户端写法

测试
- 修复 convert_test 对新 input 形态的断言
- 网关 e2e(/tmp/test_gateway.py + mock upstream)72/72 全部通过,连续 3 次稳定
2026-09-01 00:46:05 +08:00
Sakurasan f81b364436 feat: 三协议互转网关 + 鉴权修复 + 管理端增强
后端
- 新增 proxy/convert 三协议(chat/messages/responses)请求、响应与 SSE 流式互转,
  以 Chat 为中间模型;usage.go 统一提取三协议 token 用量(含单测)
- gateway: 跨协议调度(渠道未声明客户端协议时转为渠道首选格式),
  streamResponse 按 \n\n 分块逐行转换直通,bufferResponse 转换失败时剥非 JSON 前缀
- gateway: 新增 SetUsageRecorder 注入异步用量记录器
- auth_llm: 修复 key_prefix 查询长度错配([:8] vs 存储的 [:12])导致全部 401;
  修复长度 8-11 的 key 切片越界 panic;统一 unauthorized 响应
- usage: 日报表改为增量累加 upsert,避免多次 flush 互相清零;记录协议/错误码/时延等字段
- channel: 新增渠道并发槽 TryAcquire;健康检查支持可配置参数
- api: 新增 admin 渠道/模型/系统配置管理端点(旧端点保留兼容)

前端
- 新增渠道管理、模型管理、系统配置视图与 ChannelModelsDrawer
- 新增 ui 基础组件(Button/Badge/Input/Modal)与 protocol.ts
- 调整 Toast 样式、密钥页、路由菜单;dev 代理默认指向 3000 端口
2026-08-31 22:29:09 +08:00
73 changed files with 8223 additions and 621 deletions
+1 -1
View File
@@ -9,7 +9,7 @@ web
# Go 构建产物 # Go 构建产物
bin bin
cmd/openteam/dist backend/cmd/openteam/dist
# 文档与 CI # 文档与 CI
doc doc
+48
View File
@@ -0,0 +1,48 @@
# ===========================================
# OpenCatd-Open 配置文件
# 复制此文件为 .env 并修改相应配置
# ===========================================
# --- 服务器配置 ---
PORT=80
READ_TIMEOUT=10
WRITE_TIMEOUT=10
# --- Passkey (WebAuthn) 配置 ---
# 应用名称(显示给用户)
APP_NAME=OpenTeam
# 依赖方 ID(通常为域名,生产环境需改为实际域名)
RPID=localhost
# 依赖方来源(前端 URL,逗号分隔)
RPORIGINS=http://localhost:5173,http://localhost:3000
# --- 数据库配置 ---
# 支持: sqlite, mysql, postgres
DB_TYPE=sqlite
# DSN 连接字符串(SQLite 可留空)
DB_DSN=
DB_MAX_OPEN_CONNS=10
DB_MAX_IDLE_CONNS=5
# --- Redis 配置(可选,用于分布式 passkey session)---
# REDIS_HOST=localhost
# REDIS_PORT=6379
# REDIS_PASSWORD=
# REDIS_DB=0
# --- 日志配置 ---
LOG_LEVEL=info
LOG_PATH=./logs/
# --- 功能开关 ---
# 允许注册(false=关闭注册)
ALLOW_REGISTER=false
# 无限制配额(true=不限制)
UNLIMITED_QUOTA=true
# 新用户默认激活
DEFAULT_ACTIVE=true
# --- 用量统计 ---
USAGE_WORKER=1
USAGE_CHAN_SIZE=1000
TASK_TIME_INTERVAL=60
+3
View File
@@ -7,6 +7,9 @@ demo/
.env .env
openteam openteam
# paseo 任务运行时记录
.pi/
# 构建产物(make web 生成,由 go:embed 打进二进制);保留 .gitkeep 占位使未构建前也能编译 # 构建产物(make web 生成,由 go:embed 打进二进制);保留 .gitkeep 占位使未构建前也能编译
backend/cmd/openteam/dist/* backend/cmd/openteam/dist/*
!backend/cmd/openteam/dist/.gitkeep !backend/cmd/openteam/dist/.gitkeep
+20 -1
View File
@@ -1,4 +1,4 @@
.PHONY: build run test clean fmt lint frontend dev dev-backend dev-frontend .PHONY: build run test clean fmt lint frontend dev dev-backend dev-frontend docker docker-cn docker-multi
BINARY_NAME=openteam BINARY_NAME=openteam
BUILD_DIR=bin BUILD_DIR=bin
@@ -8,6 +8,10 @@ BACKEND_DIR=backend
build: frontend build: frontend
cd $(BACKEND_DIR) && CGO_ENABLED=0 go build -o ../$(BUILD_DIR)/$(BINARY_NAME) ./cmd/openteam cd $(BACKEND_DIR) && CGO_ENABLED=0 go build -o ../$(BUILD_DIR)/$(BINARY_NAME) ./cmd/openteam
# Build backend only (frontend dist must exist)
build-backend:
cd $(BACKEND_DIR) && CGO_ENABLED=0 go build -o ../$(BUILD_DIR)/$(BINARY_NAME) ./cmd/openteam
# Build frontend and copy dist # Build frontend and copy dist
frontend: frontend:
cd frontend && pnpm install && pnpm build cd frontend && pnpm install && pnpm build
@@ -74,3 +78,18 @@ migrate:
# Seed data (will be implemented) # Seed data (will be implemented)
seed: seed:
@echo "Seeding will be implemented in future" @echo "Seeding will be implemented in future"
# Docker build (default platform)
docker:
docker build -f deploy/docker/Dockerfile -t $(BINARY_NAME):latest .
# Docker build (China mirror accelerated)
docker-cn:
docker build -f deploy/docker/Dockerfile.cn -t $(BINARY_NAME):latest .
# Docker build multi-platform (requires: docker buildx)
docker-multi:
docker buildx build -f deploy/docker/Dockerfile \
--platform linux/amd64,linux/arm64 \
-t $(BINARY_NAME):latest --push .
+3
View File
@@ -27,6 +27,7 @@ require (
filippo.io/edwards25519 v1.1.0 // indirect filippo.io/edwards25519 v1.1.0 // indirect
github.com/bytedance/sonic v1.13.2 // indirect github.com/bytedance/sonic v1.13.2 // indirect
github.com/bytedance/sonic/loader v0.2.4 // indirect github.com/bytedance/sonic/loader v0.2.4 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/cloudwego/base64x v0.1.5 // indirect github.com/cloudwego/base64x v0.1.5 // indirect
github.com/dlclark/regexp2 v1.11.4 // indirect github.com/dlclark/regexp2 v1.11.4 // indirect
github.com/fxamacker/cbor/v2 v2.8.0 // indirect github.com/fxamacker/cbor/v2 v2.8.0 // indirect
@@ -59,10 +60,12 @@ require (
github.com/ncruces/go-sqlite3-wasm/v2 v2.1.35300 // indirect github.com/ncruces/go-sqlite3-wasm/v2 v2.1.35300 // indirect
github.com/ncruces/julianday v1.0.0 // indirect github.com/ncruces/julianday v1.0.0 // indirect
github.com/pelletier/go-toml/v2 v2.2.3 // indirect github.com/pelletier/go-toml/v2 v2.2.3 // indirect
github.com/redis/go-redis/v9 v9.22.0 // indirect
github.com/spf13/pflag v1.0.6 // indirect github.com/spf13/pflag v1.0.6 // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.2.12 // indirect github.com/ugorji/go/codec v1.2.12 // indirect
github.com/x448/float16 v0.8.4 // indirect github.com/x448/float16 v0.8.4 // indirect
go.uber.org/atomic v1.11.0 // indirect
golang.org/x/arch v0.16.0 // indirect golang.org/x/arch v0.16.0 // indirect
golang.org/x/net v0.52.0 // indirect golang.org/x/net v0.52.0 // indirect
golang.org/x/sync v0.20.0 // indirect golang.org/x/sync v0.20.0 // indirect
+6
View File
@@ -7,6 +7,8 @@ github.com/bytedance/sonic v1.13.2/go.mod h1:o68xyaF9u2gvVBuGHPlUVCy+ZfmNNO5ETf1
github.com/bytedance/sonic/loader v0.1.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU= github.com/bytedance/sonic/loader v0.1.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU=
github.com/bytedance/sonic/loader v0.2.4 h1:ZWCw4stuXUsn1/+zQDqeE7JKP+QO47tz7QCNan80NzY= github.com/bytedance/sonic/loader v0.2.4 h1:ZWCw4stuXUsn1/+zQDqeE7JKP+QO47tz7QCNan80NzY=
github.com/bytedance/sonic/loader v0.2.4/go.mod h1:N8A3vUdtUebEY2/VQC0MyhYeKUFosQU6FxH2JmUe6VI= github.com/bytedance/sonic/loader v0.2.4/go.mod h1:N8A3vUdtUebEY2/VQC0MyhYeKUFosQU6FxH2JmUe6VI=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudwego/base64x v0.1.5 h1:XPciSp1xaq2VCSt6lF0phncD4koWyULpl5bUxbfCyP4= github.com/cloudwego/base64x v0.1.5 h1:XPciSp1xaq2VCSt6lF0phncD4koWyULpl5bUxbfCyP4=
github.com/cloudwego/base64x v0.1.5/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w= github.com/cloudwego/base64x v0.1.5/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w=
github.com/cloudwego/iasm v0.2.0/go.mod h1:8rXZaNYT2n95jn+zTI1sDr+IgcD2GVs0nlbbQPiEFhY= github.com/cloudwego/iasm v0.2.0/go.mod h1:8rXZaNYT2n95jn+zTI1sDr+IgcD2GVs0nlbbQPiEFhY=
@@ -112,6 +114,8 @@ github.com/pkoukk/tiktoken-go v0.1.7 h1:qOBHXX4PHtvIvmOtyg1EeKlwFRiMKAcoMp4Q+bLQ
github.com/pkoukk/tiktoken-go v0.1.7/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg= github.com/pkoukk/tiktoken-go v0.1.7/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/redis/go-redis/v9 v9.22.0 h1:laDvpYXTJtZLloinw1fA5Kqd6HAEH2XKxOkG/PDq2F0=
github.com/redis/go-redis/v9 v9.22.0/go.mod h1:y2g0Wj8rQvuK0ELM+oxSudcLtC09JScs98I/X9gRWY4=
github.com/rogpeppe/go-internal v1.8.0 h1:FCbCCtXNOY3UtUuHUYaghJg4y7Fd14rXifAYUAtL9R8= github.com/rogpeppe/go-internal v1.8.0 h1:FCbCCtXNOY3UtUuHUYaghJg4y7Fd14rXifAYUAtL9R8=
github.com/rogpeppe/go-internal v1.8.0/go.mod h1:WmiCO8CzOY8rg0OYDC4/i/2WRWAB6poM+XZ2dLUbcbE= github.com/rogpeppe/go-internal v1.8.0/go.mod h1:WmiCO8CzOY8rg0OYDC4/i/2WRWAB6poM+XZ2dLUbcbE=
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
@@ -141,6 +145,8 @@ github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZ
github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM=
github.com/x448/float16 v0.8.4/go.mod h1:14CWIYCyZA/cWjXOioeEpHeN/83MdbZDRQHoFcYsOfg= github.com/x448/float16 v0.8.4/go.mod h1:14CWIYCyZA/cWjXOioeEpHeN/83MdbZDRQHoFcYsOfg=
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
golang.org/x/arch v0.16.0 h1:foMtLTdyOmIniqWCHjY6+JxuC54XP1fDwx4N0ASyW+U= golang.org/x/arch v0.16.0 h1:foMtLTdyOmIniqWCHjY6+JxuC54XP1fDwx4N0ASyW+U=
golang.org/x/arch v0.16.0/go.mod h1:JmwW7aLIoRUKgaTzhkiEFxvcEiQGyOg9BMonBJUS7EE= golang.org/x/arch v0.16.0/go.mod h1:JmwW7aLIoRUKgaTzhkiEFxvcEiQGyOg9BMonBJUS7EE=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
+588
View File
@@ -0,0 +1,588 @@
package api
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
"opencatd-open/internal/pkg/crypto"
"opencatd-open/internal/store"
"github.com/gin-gonic/gin"
)
// AdminChannels GET /api/admin/channels — 渠道列表(不返回加密 key,返回掩码)。
func (h *Handler) AdminChannels(c *gin.Context) {
var chs []store.Channel
if err := h.db.Order("id ASC").Find(&chs).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load channels"})
return
}
out := make([]gin.H, 0, len(chs))
for _, ch := range chs {
masked := ""
if key, err := crypto.Decrypt(ch.APIKeyEnc); err == nil && len(key) > 8 {
masked = maskAPIKey(key)
} else if err == nil {
masked = "****"
}
out = append(out, gin.H{
"id": ch.ID, "name": ch.Name, "provider": ch.Provider, "formats": ch.FormatsEffective(),
"base_url": ch.BaseURL, "base_urls": ch.BaseURLs,
"api_key_masked": masked, "weight": ch.Weight, "priority": ch.Priority,
"timeout_ms": ch.TimeoutMS, "max_concurrency": ch.MaxConcurrency,
"health_status": ch.HealthStatus, "enabled": ch.Enabled,
"created_at": ch.CreatedAt,
})
}
c.JSON(http.StatusOK, gin.H{"data": out})
}
type channelBody struct {
Name string `json:"name" binding:"required,min=1,max=64"`
Provider string `json:"provider"`
Formats []string `json:"formats"`
BaseURL string `json:"base_url"`
BaseURLs map[string]string `json:"base_urls"`
APIKey string `json:"api_key"`
Weight *int `json:"weight"`
Priority *int `json:"priority"`
TimeoutMS *int `json:"timeout_ms"`
MaxConcurrency *int `json:"max_concurrency"`
Enabled *bool `json:"enabled"`
}
// normalizeBaseURLs 校验并清理分协议 base_url。
func normalizeBaseURLs(m map[string]string) map[string]string {
if len(m) == 0 {
return nil
}
out := map[string]string{}
for k, v := range m {
if validFormats[k] && strings.TrimSpace(v) != "" {
out[k] = strings.TrimRight(strings.TrimSpace(v), "/")
}
}
if len(out) == 0 {
return nil
}
return out
}
// resolveBaseURL 渠道 base_url:留空按供应商默认;网关按内容智能识别前缀/完整端点。
func resolveBaseURL(provider, raw string) (string, error) {
base := strings.TrimRight(raw, "/")
if base == "" {
switch provider {
case store.ChannelProviderOpenAI:
base = "https://api.openai.com"
case store.ChannelProviderAnthropic:
base = "https://api.anthropic.com"
}
}
if base == "" {
return "", errors.New("base_url required for compatible channels")
}
return base, nil
}
func validateProvider(p string) bool {
return p == store.ChannelProviderOpenAI || p == store.ChannelProviderAnthropic || p == store.ChannelProviderCompatible
}
var validFormats = map[string]bool{
store.FormatChat: true, store.FormatResponses: true, store.FormatMessages: true,
}
// deriveProvider 按格式推断供应商(仅作内部字段/兼容用途,不参与路由)。
func deriveProvider(formats []string) string {
if len(formats) == 0 {
return store.ChannelProviderCompatible
}
messagesOnly, hasResponses := true, false
for _, f := range formats {
if f != store.FormatMessages {
messagesOnly = false
}
if f == store.FormatResponses {
hasResponses = true
}
}
if messagesOnly {
return store.ChannelProviderAnthropic
}
if hasResponses {
return store.ChannelProviderOpenAI
}
return store.ChannelProviderCompatible
}
// resolveFormats 渠道协议格式:显式给出则校验去重;空则按 provider 推断默认。
func resolveFormats(provider string, formats []string) ([]string, error) {
if len(formats) == 0 {
switch provider {
case store.ChannelProviderAnthropic:
return []string{store.FormatMessages}, nil
case store.ChannelProviderOpenAI:
return []string{store.FormatChat, store.FormatResponses}, nil
default:
return []string{store.FormatChat}, nil
}
}
seen := map[string]bool{}
out := make([]string, 0, len(formats))
for _, f := range formats {
if !validFormats[f] {
return nil, fmt.Errorf("unsupported format %q", f)
}
if !seen[f] {
seen[f] = true
out = append(out, f)
}
}
return out, nil
}
// AdminCreateChannel POST /api/admin/channels
func (h *Handler) AdminCreateChannel(c *gin.Context) {
var req channelBody
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input: " + err.Error()})
return
}
if req.Provider == "" {
req.Provider = deriveProvider(req.Formats)
}
if !validateProvider(req.Provider) {
c.JSON(http.StatusBadRequest, gin.H{"error": "provider must be openai, anthropic or compatible"})
return
}
if req.APIKey == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "api_key required"})
return
}
formats, err := resolveFormats(req.Provider, req.Formats)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
baseURL, err := resolveBaseURL(req.Provider, req.BaseURL)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
enc, err := crypto.Encrypt(req.APIKey)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to encrypt api key"})
return
}
ch := store.Channel{
Name: req.Name, Provider: req.Provider, Formats: formats, BaseURL: baseURL,
BaseURLs: normalizeBaseURLs(req.BaseURLs),
APIKeyEnc: enc, Weight: intOr(req.Weight, 1), Priority: intOr(req.Priority, 0),
TimeoutMS: intOr(req.TimeoutMS, 120000), MaxConcurrency: intOr(req.MaxConcurrency, 16),
HealthStatus: store.ChannelHealthHealthy, Enabled: boolOr(req.Enabled, true),
}
if err := h.db.Create(&ch).Error; err != nil {
c.JSON(http.StatusConflict, gin.H{"error": "failed to create channel (name may already exist)"})
return
}
c.JSON(http.StatusCreated, gin.H{"id": ch.ID, "name": ch.Name})
}
// AdminUpdateChannel PUT /api/admin/channels/:id
func (h *Handler) AdminUpdateChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
var body struct {
Name *string `json:"name"`
Provider *string `json:"provider"`
Formats *[]string `json:"formats"`
BaseURL *string `json:"base_url"`
BaseURLs *map[string]string `json:"base_urls"`
APIKey *string `json:"api_key"`
Weight *int `json:"weight"`
Priority *int `json:"priority"`
TimeoutMS *int `json:"timeout_ms"`
MaxConcurrency *int `json:"max_concurrency"`
HealthStatus *string `json:"health_status"`
Enabled *bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&body); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
var ch store.Channel
if err := h.db.First(&ch, id).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
updates := map[string]any{}
if body.Name != nil {
updates["name"] = *body.Name
}
if body.Provider != nil {
if !validateProvider(*body.Provider) {
c.JSON(http.StatusBadRequest, gin.H{"error": "provider must be openai, anthropic or compatible"})
return
}
updates["provider"] = *body.Provider
}
if body.BaseURL != nil {
prov := ch.Provider
if body.Provider != nil {
prov = *body.Provider
}
b, berr := resolveBaseURL(prov, *body.BaseURL)
if berr != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": berr.Error()})
return
}
updates["base_url"] = b
}
if body.BaseURLs != nil {
raw, _ := json.Marshal(normalizeBaseURLs(*body.BaseURLs))
updates["base_urls"] = string(raw)
}
if body.APIKey != nil && *body.APIKey != "" {
enc, err := crypto.Encrypt(*body.APIKey)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to encrypt api key"})
return
}
updates["api_key_enc"] = enc
}
if body.Weight != nil {
updates["weight"] = *body.Weight
}
if body.Priority != nil {
updates["priority"] = *body.Priority
}
if body.TimeoutMS != nil {
updates["timeout_ms"] = *body.TimeoutMS
}
if body.MaxConcurrency != nil {
updates["max_concurrency"] = *body.MaxConcurrency
}
if body.HealthStatus != nil {
updates["health_status"] = *body.HealthStatus
}
if body.Enabled != nil {
updates["enabled"] = *body.Enabled
}
if body.Formats != nil {
prov := ch.Provider
if body.Provider != nil {
prov = *body.Provider
}
formats, ferr := resolveFormats(prov, *body.Formats)
if ferr != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": ferr.Error()})
return
}
raw, _ := json.Marshal(formats)
updates["formats"] = string(raw)
}
if len(updates) > 0 {
if err := h.db.Model(&ch).Updates(updates).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update channel"})
return
}
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminDeleteChannel DELETE /api/admin/channels/:id
func (h *Handler) AdminDeleteChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
res := h.db.Delete(&store.Channel{}, id)
if res.Error != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete channel"})
return
}
if res.RowsAffected == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
h.db.Where("channel_id = ?", id).Delete(&store.ChannelModelBinding{})
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminTestChannel POST /api/admin/channels/:id/test — 请求渠道 /v1/models 测连通性。
func (h *Handler) AdminTestChannel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
var ch store.Channel
if err := h.db.First(&ch, id).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
key, err := crypto.Decrypt(ch.APIKeyEnc)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to decrypt channel key"})
return
}
url := ch.UpstreamURL("", "/models")
client := &http.Client{Timeout: 10 * time.Second}
req, _ := http.NewRequest(http.MethodGet, url, nil)
req.Header.Set("Authorization", "Bearer "+key)
req.Header.Set("Accept", "application/json")
start := time.Now()
resp, err := client.Do(req)
status := store.ChannelHealthHealthy
msg := "ok"
latency := 0
if err != nil {
status = store.ChannelHealthCooldown
msg = err.Error()
} else {
latency = int(time.Since(start).Milliseconds())
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
status = store.ChannelHealthCooldown
b, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
msg = fmt.Sprintf("http %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
}
resp.Body.Close()
}
h.db.Model(&store.Channel{}).Where("id = ?", ch.ID).Update("health_status", status)
if status != store.ChannelHealthHealthy {
c.JSON(http.StatusBadGateway, gin.H{"error": msg})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true, "latency_ms": latency, "message": msg})
}
// AdminChannelRemoteModels GET /api/admin/channels/:id/models/remote — 拉取远端模型列表。
// 返回本渠道尚未允许的模型(新增候选),排除已绑定的模型。
func (h *Handler) AdminChannelRemoteModels(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
var ch store.Channel
if err := h.db.First(&ch, id).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
key, err := crypto.Decrypt(ch.APIKeyEnc)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to decrypt channel key"})
return
}
url := ch.UpstreamURL("", "/models")
client := &http.Client{Timeout: 10 * time.Second}
req, _ := http.NewRequest(http.MethodGet, url, nil)
req.Header.Set("Authorization", "Bearer "+key)
req.Header.Set("Accept", "application/json")
resp, err := client.Do(req)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
defer resp.Body.Close()
body, _ := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024))
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
c.JSON(http.StatusBadGateway, gin.H{"error": fmt.Sprintf("http %d: %s", resp.StatusCode, string(body))})
return
}
// 解析 OpenAI 格式的模型列表
var result struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
if err := json.Unmarshal(body, &result); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": "failed to parse response: " + err.Error()})
return
}
// 本渠道已允许的上游模型名:不作为新增候选
var boundNames []string
h.db.Model(&store.ChannelModelBinding{}).Where("channel_id = ?", id).Pluck("upstream_model", &boundNames)
boundSet := make(map[string]bool, len(boundNames))
for _, n := range boundNames {
boundSet[strings.TrimSpace(n)] = true
}
models := make([]string, 0, len(result.Data))
for _, m := range result.Data {
name := strings.TrimSpace(m.ID)
if name != "" && !boundSet[name] {
models = append(models, name)
}
}
c.JSON(http.StatusOK, gin.H{"data": models})
}
// AdminChannelModels GET /api/admin/channels/:id/models — 渠道绑定列表。
func (h *Handler) AdminChannelModels(c *gin.Context) {
channelID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
var bindings []store.ChannelModelBinding
if err := h.db.Preload("Model").Where("channel_id = ?", channelID).Find(&bindings).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load bindings"})
return
}
out := make([]gin.H, 0, len(bindings))
for _, b := range bindings {
out = append(out, gin.H{
"id": b.ID, "model_id": b.ModelID, "model_name": b.Model.Name,
"upstream_model": b.UpstreamModel, "weight": b.Weight,
})
}
c.JSON(http.StatusOK, gin.H{"data": out})
}
// AdminChannelAddModel POST /api/admin/channels/:id/models — 手工添加模型绑定。
// 无需渠道具备 /v1/models 接口:直接填上游模型名,可选自定义名称作为客户端调用名。
func (h *Handler) AdminChannelAddModel(c *gin.Context) {
channelID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid channel id"})
return
}
var req struct {
UpstreamModel string `json:"upstream_model" binding:"required"` // 渠道侧真实模型名
CustomName string `json:"custom_name"` // 客户端调用名,空=用上游名
Weight *int `json:"weight"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input: upstream_model required"})
return
}
globalName := req.CustomName
if globalName == "" {
globalName = req.UpstreamModel
}
var ch store.Channel
if err := h.db.First(&ch, channelID).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
// 查找或创建全局模型
var m store.Model
if err := h.db.Where("name = ?", globalName).First(&m).Error; err != nil {
m = store.Model{Name: globalName, Enabled: true}
if err := h.db.Create(&m).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create model"})
return
}
}
// 查找已存在的绑定,如果存在则更新
var existing store.ChannelModelBinding
if err := h.db.Where("channel_id = ? AND model_id = ?", channelID, m.ID).First(&existing).Error; err == nil {
// 已存在,更新
existing.UpstreamModel = req.UpstreamModel
existing.Weight = intOr(req.Weight, 1)
if err := h.db.Save(&existing).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update binding"})
return
}
c.JSON(http.StatusOK, gin.H{"id": existing.ID, "model_id": m.ID, "model_name": m.Name, "upstream_model": existing.UpstreamModel, "weight": existing.Weight})
return
}
// 不存在,创建新的
b := store.ChannelModelBinding{
ChannelID: channelID, ModelID: m.ID,
UpstreamModel: req.UpstreamModel, Weight: intOr(req.Weight, 1),
}
if err := h.db.Create(&b).Error; err != nil {
c.JSON(http.StatusConflict, gin.H{"error": "binding may already exist"})
return
}
c.JSON(http.StatusCreated, gin.H{"id": b.ID, "model_id": m.ID, "model_name": m.Name, "upstream_model": req.UpstreamModel, "weight": b.Weight})
}
// AdminChannelUpdateModel PATCH /api/admin/channels/:id/models/:bid — 改映射名/权重。
func (h *Handler) AdminChannelUpdateModel(c *gin.Context) {
bid, err := strconv.ParseUint(c.Param("bid"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid binding id"})
return
}
var req struct {
UpstreamModel *string `json:"upstream_model"`
Weight *int `json:"weight"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
var b store.ChannelModelBinding
if err := h.db.First(&b, bid).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "binding not found"})
return
}
updates := map[string]any{}
if req.UpstreamModel != nil {
updates["upstream_model"] = *req.UpstreamModel
}
if req.Weight != nil {
updates["weight"] = *req.Weight
}
if len(updates) > 0 {
if err := h.db.Model(&b).Updates(updates).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update binding"})
return
}
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminChannelDeleteModel DELETE /api/admin/channels/:id/models/:bid — 解除绑定。
func (h *Handler) AdminChannelDeleteModel(c *gin.Context) {
bid, err := strconv.ParseUint(c.Param("bid"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid binding id"})
return
}
res := h.db.Delete(&store.ChannelModelBinding{}, bid)
if res.Error != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete binding"})
return
}
if res.RowsAffected == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "binding not found"})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// maskAPIKey 掩码渠道密钥:保留前 7 位与后 4 位,中间固定 ****** 遮蔽。
func maskAPIKey(key string) string {
if len(key) <= 11 {
return strings.Repeat("*", len(key)-4) + key[len(key)-4:]
}
return key[:7] + "******" + key[len(key)-4:]
}
func intOr(p *int, def int) int {
if p == nil {
return def
}
return *p
}
func boolOr(p *bool, def bool) bool {
if p == nil {
return def
}
return *p
}
+109
View File
@@ -0,0 +1,109 @@
package api
import (
"net/http"
"opencatd-open/internal/store"
"github.com/gin-gonic/gin"
)
// AdminGetConfig GET /api/admin/config — 获取系统配置。
func (h *Handler) AdminGetConfig(c *gin.Context) {
configs := map[string]string{}
var rows []store.SystemConfig
h.db.Find(&rows)
for _, r := range rows {
configs[r.Key] = r.Value
}
c.JSON(http.StatusOK, gin.H{"data": configs})
}
// AdminUpdateConfig PUT /api/admin/config — 更新系统配置。
func (h *Handler) AdminUpdateConfig(c *gin.Context) {
var req map[string]string
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
for key, value := range req {
var sc store.SystemConfig
result := h.db.Where("key = ?", key).First(&sc)
if result.Error == nil {
sc.Value = value
h.db.Save(&sc)
} else {
sc = store.SystemConfig{Key: key, Value: value}
h.db.Create(&sc)
}
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminGetRegistration GET /api/admin/config/registration — 获取注册配置。
func (h *Handler) AdminGetRegistration(c *gin.Context) {
var sc store.SystemConfig
enabled := "true"
if err := h.db.Where("key = ?", "registration_enabled").First(&sc).Error; err == nil {
enabled = sc.Value
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"enabled": enabled == "true"}})
}
// AdminUpdateRegistration PUT /api/admin/config/registration — 更新注册配置。
func (h *Handler) AdminUpdateRegistration(c *gin.Context) {
var req struct {
Enabled bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
value := "false"
if req.Enabled {
value = "true"
}
var sc store.SystemConfig
if err := h.db.Where("key = ?", "registration_enabled").First(&sc).Error; err == nil {
sc.Value = value
h.db.Save(&sc)
} else {
sc = store.SystemConfig{Key: "registration_enabled", Value: value}
h.db.Create(&sc)
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminGetPasswordLogin GET /api/admin/config/password-login — 获取密码登录配置。
func (h *Handler) AdminGetPasswordLogin(c *gin.Context) {
var sc store.SystemConfig
enabled := "true"
if err := h.db.Where("key = ?", "password_login_enabled").First(&sc).Error; err == nil {
enabled = sc.Value
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"enabled": enabled == "true"}})
}
// AdminUpdatePasswordLogin PUT /api/admin/config/password-login — 更新密码登录配置。
func (h *Handler) AdminUpdatePasswordLogin(c *gin.Context) {
var req struct {
Enabled bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
value := "false"
if req.Enabled {
value = "true"
}
var sc store.SystemConfig
if err := h.db.Where("key = ?", "password_login_enabled").First(&sc).Error; err == nil {
sc.Value = value
h.db.Save(&sc)
} else {
sc = store.SystemConfig{Key: "password_login_enabled", Value: value}
h.db.Create(&sc)
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
+291
View File
@@ -0,0 +1,291 @@
package api
import (
"encoding/json"
"net/http"
"strconv"
"opencatd-open/internal/store"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// AdminModels GET /api/admin/models — 模型列表(含价格、渠道绑定、定价/禁止状态)。
func (h *Handler) AdminModels(c *gin.Context) {
var ms []store.Model
if err := h.db.Order("sort ASC, id ASC").Find(&ms).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load models"})
return
}
allow, deny := h.modelPolicyConfig()
out := make([]gin.H, 0, len(ms))
for _, m := range ms {
var bindings []store.ChannelModelBinding
h.db.Preload("Channel").Where("model_id = ?", m.ID).Find(&bindings)
chs := make([]gin.H, 0, len(bindings))
for _, b := range bindings {
if !b.Channel.Enabled {
continue
}
chs = append(chs, gin.H{
"id": b.ID, "channel_id": b.ChannelID, "channel_name": b.Channel.Name,
"upstream_model": b.UpstreamModel, "weight": b.Weight,
})
}
used := len(chs) > 0
needsPricing := used && m.InputPrice == 0 && m.OutputPrice == 0 && m.CacheReadPrice == 0
denied := containsStr(deny, m.Name) || (len(allow) > 0 && !containsStr(allow, m.Name))
out = append(out, gin.H{
"id": m.ID, "name": m.Name, "display_name": m.DisplayName,
"input_price": m.InputPrice, "output_price": m.OutputPrice, "cache_read_price": m.CacheReadPrice,
"enabled": m.Enabled, "sort": m.Sort, "channels": chs,
"used": used, "needs_pricing": needsPricing, "denied": denied,
})
}
var orphans []struct {
ChannelName string
UpstreamModel string
ModelID uint64
}
h.db.Raw(`SELECT c.name as channel_name, b.model_id, b.upstream_model
FROM channel_model_bindings b
LEFT JOIN models m ON m.id = b.model_id
LEFT JOIN channels c ON c.id = b.channel_id
WHERE m.id IS NULL`).Scan(&orphans)
missing := make([]gin.H, 0, len(orphans))
for _, o := range orphans {
missing = append(missing, gin.H{
"channel": o.ChannelName, "model_id": o.ModelID, "upstream_model": o.UpstreamModel,
})
}
unpriced := 0
{
var usedBindings []struct {
ModelID uint64
}
h.db.Model(&store.ChannelModelBinding{}).Distinct("model_id").Scan(&usedBindings)
usedIDs := map[uint64]bool{}
for _, u := range usedBindings {
usedIDs[u.ModelID] = true
}
for _, m := range ms {
if usedIDs[m.ID] && m.InputPrice == 0 && m.OutputPrice == 0 && m.CacheReadPrice == 0 {
unpriced++
}
}
}
c.JSON(http.StatusOK, gin.H{
"data": out,
"summary": gin.H{
"total": len(ms),
"unpriced": unpriced,
"missing": missing,
"denied_count": len(deny),
},
})
}
// modelPolicyConfig 读取全局模型允许/禁止列表。
func (h *Handler) modelPolicyConfig() (allow, deny []string) {
var raw string
h.db.Model(&store.SystemConfig{}).Where("key = ?", "model_allowlist").Pluck("value", &raw)
_ = json.Unmarshal([]byte(raw), &allow)
raw = ""
h.db.Model(&store.SystemConfig{}).Where("key = ?", "model_denylist").Pluck("value", &raw)
_ = json.Unmarshal([]byte(raw), &deny)
return
}
func containsStr(list []string, s string) bool {
for _, v := range list {
if v == s {
return true
}
}
return false
}
// AdminCreateModel POST /api/admin/models
func (h *Handler) AdminCreateModel(c *gin.Context) {
var req struct {
Name string `json:"name" binding:"required,min=1,max=128"`
DisplayName string `json:"display_name"`
InputPrice float64 `json:"input_price"`
OutputPrice float64 `json:"output_price"`
CacheReadPrice float64 `json:"cache_read_price"`
Sort int `json:"sort"`
Enabled *bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input: " + err.Error()})
return
}
m := store.Model{
Name: req.Name, DisplayName: req.DisplayName,
InputPrice: req.InputPrice, OutputPrice: req.OutputPrice, CacheReadPrice: req.CacheReadPrice,
Sort: req.Sort, Enabled: boolOr(req.Enabled, true),
}
if err := h.db.Create(&m).Error; err != nil {
c.JSON(http.StatusConflict, gin.H{"error": "failed to create model (name may already exist)"})
return
}
c.JSON(http.StatusCreated, gin.H{"id": m.ID, "name": m.Name})
}
// AdminUpdateModel PUT /api/admin/models/:id — 价格/启停/排序。
func (h *Handler) AdminUpdateModel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid model id"})
return
}
var req struct {
DisplayName *string `json:"display_name"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
Enabled *bool `json:"enabled"`
Sort *int `json:"sort"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
var m store.Model
if err := h.db.First(&m, id).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "model not found"})
return
}
updates := map[string]any{}
if req.DisplayName != nil {
updates["display_name"] = *req.DisplayName
}
if req.InputPrice != nil {
updates["input_price"] = *req.InputPrice
}
if req.OutputPrice != nil {
updates["output_price"] = *req.OutputPrice
}
if req.CacheReadPrice != nil {
updates["cache_read_price"] = *req.CacheReadPrice
}
if req.Enabled != nil {
updates["enabled"] = *req.Enabled
}
if req.Sort != nil {
updates["sort"] = *req.Sort
}
if len(updates) > 0 {
if err := h.db.Model(&m).Updates(updates).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update model"})
return
}
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminDeleteModel DELETE /api/admin/models/:id
func (h *Handler) AdminDeleteModel(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid model id"})
return
}
res := h.db.Delete(&store.Model{}, id)
if res.Error != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete model"})
return
}
if res.RowsAffected == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "model not found"})
return
}
h.db.Where("model_id = ?", id).Delete(&store.ChannelModelBinding{})
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// AdminDeleteUnusedModels DELETE /api/admin/models/unused — 一键清除未绑定任何渠道的模型。
func (h *Handler) AdminDeleteUnusedModels(c *gin.Context) {
var orphans []store.Model
if err := h.db.Where("id NOT IN (SELECT DISTINCT model_id FROM channel_model_bindings)").Find(&orphans).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load models"})
return
}
names := make([]string, 0, len(orphans))
ids := make([]uint64, 0, len(orphans))
for _, m := range orphans {
names = append(names, m.Name)
ids = append(ids, m.ID)
}
if len(ids) > 0 {
if err := h.db.Delete(&store.Model{}, ids).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete models"})
return
}
}
c.JSON(http.StatusOK, gin.H{"deleted": names, "count": len(names)})
}
// AdminCreateModelBinding POST /api/admin/models/:id/bindings
func (h *Handler) AdminCreateModelBinding(c *gin.Context) {
modelID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid model id"})
return
}
var req struct {
ChannelID uint64 `json:"channel_id" binding:"required"`
UpstreamModel string `json:"upstream_model" binding:"required"`
Weight *int `json:"weight"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input: channel_id and upstream_model required"})
return
}
var m store.Model
if err := h.db.First(&m, modelID).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "model not found"})
return
}
var ch store.Channel
if err := h.db.First(&ch, req.ChannelID).Error; err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return
}
b := store.ChannelModelBinding{
ChannelID: req.ChannelID, ModelID: modelID,
UpstreamModel: req.UpstreamModel, Weight: intOr(req.Weight, 1),
}
if err := h.db.Create(&b).Error; err != nil {
c.JSON(http.StatusConflict, gin.H{"error": "binding may already exist"})
return
}
c.JSON(http.StatusCreated, gin.H{"id": b.ID})
}
// AdminDeleteModelBinding DELETE /api/admin/models/:id/bindings/:bid
func (h *Handler) AdminDeleteModelBinding(c *gin.Context) {
bid, err := strconv.ParseUint(c.Param("bid"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid binding id"})
return
}
res := h.db.Delete(&store.ChannelModelBinding{}, bid)
if res.Error != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete binding"})
return
}
if res.RowsAffected == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "binding not found"})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
var _ = gorm.ErrRecordNotFound
+7 -4
View File
@@ -3,6 +3,7 @@ package api
import ( import (
"net/http" "net/http"
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/passkey"
"opencatd-open/internal/store" "opencatd-open/internal/store"
"opencatd-open/internal/pkg/apikey" "opencatd-open/internal/pkg/apikey"
"opencatd-open/internal/pkg/crypto" "opencatd-open/internal/pkg/crypto"
@@ -23,9 +24,10 @@ type Handler struct {
modelDAO *dao.ModelDAO modelDAO *dao.ModelDAO
usageDAO *dao.UsageDAO usageDAO *dao.UsageDAO
dailyDAO *dao.DailyUsageDAO dailyDAO *dao.DailyUsageDAO
passkeys *passkey.Service
} }
func NewHandler(db *gorm.DB) *Handler { func NewHandler(db *gorm.DB, passkeys *passkey.Service) *Handler {
return &Handler{ return &Handler{
db: db, db: db,
userDAO: dao.NewUserDAO(db), userDAO: dao.NewUserDAO(db),
@@ -34,6 +36,7 @@ func NewHandler(db *gorm.DB) *Handler {
modelDAO: dao.NewModelDAO(db), modelDAO: dao.NewModelDAO(db),
usageDAO: dao.NewUsageDAO(db), usageDAO: dao.NewUsageDAO(db),
dailyDAO: dao.NewDailyUsageDAO(db), dailyDAO: dao.NewDailyUsageDAO(db),
passkeys: passkeys,
} }
} }
@@ -556,7 +559,7 @@ func (h *Handler) DeleteApiKey(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"message": "deleted"}) c.JSON(http.StatusOK, gin.H{"message": "deleted"})
} }
// --- Channels --- // --- Legacy Channel endpoints (kept for backward compatibility) ---
func (h *Handler) ListChannels(c *gin.Context) { func (h *Handler) ListChannels(c *gin.Context) {
// Support both limit/offset and pageSize/page parameters // Support both limit/offset and pageSize/page parameters
@@ -702,7 +705,7 @@ func (h *Handler) DeleteChannel(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"message": "deleted"}) c.JSON(http.StatusOK, gin.H{"message": "deleted"})
} }
// --- Models --- // --- Legacy Model endpoints (kept for backward compatibility) ---
func (h *Handler) ListModels(c *gin.Context) { func (h *Handler) ListModels(c *gin.Context) {
// Support both limit/offset and pageSize/page parameters // Support both limit/offset and pageSize/page parameters
@@ -827,7 +830,7 @@ func (h *Handler) DeleteModel(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"message": "deleted"}) c.JSON(http.StatusOK, gin.H{"message": "deleted"})
} }
// --- Channel-Model Bindings --- // --- Legacy Channel-Model Bindings (kept for backward compatibility) ---
func (h *Handler) BindChannelModels(c *gin.Context) { func (h *Handler) BindChannelModels(c *gin.Context) {
channelID, err := strconv.ParseUint(c.Param("id"), 10, 64) channelID, err := strconv.ParseUint(c.Param("id"), 10, 64)
+162
View File
@@ -0,0 +1,162 @@
package api
import (
"encoding/json"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"opencatd-open/internal/auth"
"opencatd-open/internal/pkg/jwt"
"opencatd-open/internal/store"
)
// PasskeyRegisterBegin POST /api/webauthn/register/begin — 生成注册选项。
func (h *Handler) PasskeyRegisterBegin(c *gin.Context) {
userID, _ := c.Get("user_id")
u, err := h.passkeys.GetUserByID(userID.(uint64))
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "user not found"})
return
}
creation, err := h.passkeys.BeginRegistration(u)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to begin registration: " + err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"creation": creation, "challenge": creation.Response.Challenge}})
}
// PasskeyRegisterComplete POST /api/webauthn/register/complete — 校验并保存凭据。
func (h *Handler) PasskeyRegisterComplete(c *gin.Context) {
userID, _ := c.Get("user_id")
u, err := h.passkeys.GetUserByID(userID.(uint64))
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "user not found"})
return
}
var req struct {
Challenge string `json:"challenge"`
Name string `json:"name"`
Credential json.RawMessage `json:"credential"`
}
if err := c.ShouldBindJSON(&req); err != nil || len(req.Credential) == 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
if err := h.passkeys.FinishRegistration(u, req.Challenge, req.Credential, req.Name); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "passkey 注册失败: " + err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"ok": true}})
}
// PasskeyLoginBegin POST /api/auth/passkey/begin — 生成断言选项。
// 传 username 用指定用户;不传则用可发现凭据(平台 passkey)。
func (h *Handler) PasskeyLoginBegin(c *gin.Context) {
var req struct {
Username string `json:"username"`
}
_ = c.ShouldBindJSON(&req)
if req.Username != "" {
u, err := h.passkeys.GetUserByUsername(req.Username)
if err != nil || u.Status != store.UserStatusActive {
c.JSON(http.StatusNotFound, gin.H{"error": "user not found"})
return
}
assertion, err := h.passkeys.BeginLogin(u)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to begin login: " + err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"assertion": assertion, "challenge": assertion.Response.Challenge, "user_id": u.ID}})
return
}
assertion, err := h.passkeys.BeginDiscoverableLogin()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to begin login: " + err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"assertion": assertion, "challenge": assertion.Response.Challenge}})
}
// PasskeyLoginComplete POST /api/auth/passkey/finish — 校验断言并发放令牌。
func (h *Handler) PasskeyLoginComplete(c *gin.Context) {
var req struct {
Challenge string `json:"challenge"`
Credential json.RawMessage `json:"credential"`
UserID uint64 `json:"user_id"`
}
if err := c.ShouldBindJSON(&req); err != nil || len(req.Credential) == 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid input"})
return
}
var u *store.User
if req.UserID > 0 {
var err error
u, err = h.passkeys.GetUserByID(req.UserID)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "user not found"})
return
}
if err := h.passkeys.FinishLogin(u, req.Challenge, req.Credential); err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "passkey 校验失败: " + err.Error()})
return
}
} else {
var err error
u, err = h.passkeys.FinishDiscoverableLogin(req.Challenge, req.Credential)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "passkey 校验失败: " + err.Error()})
return
}
}
if u.Status != store.UserStatusActive {
c.JSON(http.StatusForbidden, gin.H{"error": "user account disabled"})
return
}
secret := auth.GetSecretKey()
accessToken, refreshToken, err := jwt.GenerateTokenPair(u.ID, u.Username, u.Role, secret, 24*time.Hour, 7*24*time.Hour)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to issue token"})
return
}
c.JSON(http.StatusOK, gin.H{
"data": gin.H{
"token": accessToken,
"access_token": accessToken,
"refresh_token": refreshToken,
},
})
}
// PasskeyList GET /api/profile/passkeys — 当前用户的 passkey 列表。
func (h *Handler) PasskeyList(c *gin.Context) {
userID, _ := c.Get("user_id")
pks, err := h.passkeys.List(userID.(uint64))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load passkeys"})
return
}
out := make([]gin.H, 0, len(pks))
for _, pk := range pks {
out = append(out, gin.H{"id": pk.ID, "name": pk.Name, "created_at": pk.CreatedAt})
}
c.JSON(http.StatusOK, gin.H{"data": out})
}
// PasskeyDelete DELETE /api/profile/passkeys/:id — 解除绑定。
func (h *Handler) PasskeyDelete(c *gin.Context) {
userID, _ := c.Get("user_id")
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid passkey id"})
return
}
if err := h.passkeys.Delete(userID.(uint64), id); err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "passkey not found"})
return
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{"ok": true}})
}
+391
View File
@@ -0,0 +1,391 @@
package api
import (
"fmt"
"net/http"
"sort"
"strconv"
"time"
"opencatd-open/internal/dao"
"opencatd-open/internal/store"
"github.com/gin-gonic/gin"
)
// --- 普通用户:自身用量统计与明细 ---
// MyUsageStats GET /api/usage/stats?days=30 — 当前用户的每日用量聚合。
func (h *Handler) MyUsageStats(c *gin.Context) {
userID, _ := c.Get("user_id")
uid, _ := userID.(uint64)
days := 30
if d := c.Query("days"); d != "" {
if n, err := strconv.Atoi(d); err == nil && n > 0 && n <= 365 {
days = n
}
}
end := time.Now()
start := end.AddDate(0, 0, -days)
dailies, err := h.dailyDAO.ListByDateRange(c.Request.Context(), uid, start, end)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load usage"})
return
}
// 按日期聚合(每日可能多模型多行)
byDate := map[string]*store.UsageDaily{}
var dates []string
for i := range dailies {
d := dailies[i]
agg, ok := byDate[d.Date]
if !ok {
agg = &store.UsageDaily{Date: d.Date}
byDate[d.Date] = agg
dates = append(dates, d.Date)
}
agg.Requests += d.Requests
agg.InputTokens += d.InputTokens
agg.OutputTokens += d.OutputTokens
agg.CacheReadTokens += d.CacheReadTokens
agg.Cost += d.Cost
}
// 汇总
var totalRequests, totalInput, totalOutput, totalCache int64
var totalCost float64
for _, d := range byDate {
totalRequests += d.Requests
totalInput += d.InputTokens
totalOutput += d.OutputTokens
totalCache += d.CacheReadTokens
totalCost += d.Cost
}
c.JSON(http.StatusOK, gin.H{
"data": gin.H{
"dates": dates,
"daily": byDate,
"totals": gin.H{
"requests": totalRequests,
"input_tokens": totalInput,
"output_tokens": totalOutput,
"cache_read_tokens": totalCache,
"cost": totalCost,
},
},
})
}
// MyUsageMonthly GET /api/usage/monthly?year=2026 — 当前用户年度按自然月聚合,
// 每月含按模型分解(供月度堆叠柱状图使用)。
func (h *Handler) MyUsageMonthly(c *gin.Context) {
userID, _ := c.Get("user_id")
uid, _ := userID.(uint64)
year := time.Now().Year()
if y := c.Query("year"); y != "" {
if n, err := strconv.Atoi(y); err == nil && n >= 2000 && n <= 2100 {
year = n
}
}
start := time.Date(year, 1, 1, 0, 0, 0, 0, time.Local)
end := start.AddDate(1, 0, -1)
dailies, err := h.dailyDAO.ListByDateRange(c.Request.Context(), uid, start, end)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load usage"})
return
}
// 补齐模型名(模型可能已被删除,回退为 模型#id)
modelIDs := make([]uint64, 0, len(dailies))
seen := map[uint64]bool{}
for _, d := range dailies {
if !seen[d.ModelID] {
seen[d.ModelID] = true
modelIDs = append(modelIDs, d.ModelID)
}
}
modelNames := map[uint64]string{}
if len(modelIDs) > 0 {
var models []store.Model
if err := h.db.Where("id IN ?", modelIDs).Find(&models).Error; err == nil {
for _, m := range models {
modelNames[m.ID] = m.Name
}
}
}
type modelAgg struct {
ModelID uint64 `json:"model_id"`
ModelName string `json:"model_name"`
Requests int64 `json:"requests"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
CacheReadTokens int64 `json:"cache_read_tokens"`
Cost float64 `json:"cost"`
}
type monthAgg struct {
Month string `json:"month"`
Requests int64 `json:"requests"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
CacheReadTokens int64 `json:"cache_read_tokens"`
Cost float64 `json:"cost"`
Models map[uint64]*modelAgg `json:"-"`
}
months := make([]*monthAgg, 12)
for i := range months {
months[i] = &monthAgg{
Month: fmt.Sprintf("%d-%02d", year, i+1),
Models: map[uint64]*modelAgg{},
}
}
for _, d := range dailies {
mm, err := strconv.Atoi(d.Date[5:7])
if err != nil || mm < 1 || mm > 12 {
continue
}
m := months[mm-1]
m.Requests += d.Requests
m.InputTokens += d.InputTokens
m.OutputTokens += d.OutputTokens
m.CacheReadTokens += d.CacheReadTokens
m.Cost += d.Cost
ma, ok := m.Models[d.ModelID]
if !ok {
name := modelNames[d.ModelID]
if name == "" {
name = fmt.Sprintf("模型#%d", d.ModelID)
}
ma = &modelAgg{ModelID: d.ModelID, ModelName: name}
m.Models[d.ModelID] = ma
}
ma.Requests += d.Requests
ma.InputTokens += d.InputTokens
ma.OutputTokens += d.OutputTokens
ma.CacheReadTokens += d.CacheReadTokens
ma.Cost += d.Cost
}
out := make([]gin.H, 12)
for i, m := range months {
modelList := make([]*modelAgg, 0, len(m.Models))
for _, ma := range m.Models {
modelList = append(modelList, ma)
}
// 模型按 token 总量降序,柱状图图例顺序与之一致
sort.Slice(modelList, func(a, b int) bool {
ta := modelList[a].InputTokens + modelList[a].OutputTokens + modelList[a].CacheReadTokens
tb := modelList[b].InputTokens + modelList[b].OutputTokens + modelList[b].CacheReadTokens
return ta > tb
})
out[i] = gin.H{
"month": m.Month,
"requests": m.Requests,
"input_tokens": m.InputTokens,
"output_tokens": m.OutputTokens,
"cache_read_tokens": m.CacheReadTokens,
"cost": m.Cost,
"models": modelList,
}
}
c.JSON(http.StatusOK, gin.H{
"data": gin.H{
"year": year,
"months": out,
},
})
}
// MyUsageLogs GET /api/usage/logs?page=1&pageSize=20 — 当前用户的用量明细(分页)。
func (h *Handler) MyUsageLogs(c *gin.Context) {
userID, _ := c.Get("user_id")
uid, _ := userID.(uint64)
limit, offset := paginate(c, 20)
logs, err := h.usageDAO.ListByUserID(c.Request.Context(), uid, limit, offset)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load usage logs"})
return
}
total, err := h.usageDAO.CountByUserID(c.Request.Context(), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to count usage logs"})
return
}
valLogs := make([]store.UsageLog, len(logs))
for i, l := range logs {
valLogs[i] = *l
}
c.JSON(http.StatusOK, gin.H{"data": usageLogsToResp(valLogs, nil), "total": total})
}
// --- 管理后台:全量用量明细 ---
// AdminUsageLogs GET /api/admin/usage/logs?page=&pageSize=&protocol=&status=&model=&user_id=
func (h *Handler) AdminUsageLogs(c *gin.Context) {
f := daoUsageFilter(c)
logs, err := h.usageDAO.ListAll(c.Request.Context(), f)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load usage logs"})
return
}
total, err := h.usageDAO.CountAll(c.Request.Context(), f)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to count usage logs"})
return
}
names := h.userNames(logs)
c.JSON(http.StatusOK, gin.H{
"data": usageLogsToResp(logs, names),
"total": total,
})
}
// AdminUsageSummary GET /api/admin/usage/summary?start=&end=&user_id= — 全量汇总。
func (h *Handler) AdminUsageSummary(c *gin.Context) {
var uidPtr *uint64
if v := c.Query("user_id"); v != "" {
if n, err := strconv.ParseUint(v, 10, 64); err == nil && n > 0 {
uidPtr = &n
}
}
dailies, err := h.dailyDAO.ListAll(c.Request.Context(), uidPtr, c.Query("start"), c.Query("end"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load usage"})
return
}
var totalRequests, totalInput, totalOutput, totalCache int64
var totalCost float64
perUser := map[uint64]*gin.H{}
for _, d := range dailies {
totalRequests += d.Requests
totalInput += d.InputTokens
totalOutput += d.OutputTokens
totalCache += d.CacheReadTokens
totalCost += d.Cost
u, ok := perUser[d.UserID]
if !ok {
u = &gin.H{"user_id": d.UserID, "requests": int64(0), "input_tokens": int64(0), "output_tokens": int64(0), "cost": float64(0)}
perUser[d.UserID] = u
}
(*u)["requests"] = (*u)["requests"].(int64) + d.Requests
(*u)["input_tokens"] = (*u)["input_tokens"].(int64) + d.InputTokens
(*u)["output_tokens"] = (*u)["output_tokens"].(int64) + d.OutputTokens
(*u)["cost"] = (*u)["cost"].(float64) + d.Cost
}
c.JSON(http.StatusOK, gin.H{
"data": gin.H{
"totals": gin.H{
"requests": totalRequests,
"input_tokens": totalInput,
"output_tokens": totalOutput,
"cache_read_tokens": totalCache,
"cost": totalCost,
},
"per_user": perUser,
},
})
}
// --- helpers ---
// paginate 解析 page/pageSize 查询参数,返回 limit/offset。
func paginate(c *gin.Context, defSize int) (int, int) {
limit := defSize
offset := 0
if pageSize := c.Query("pageSize"); pageSize != "" {
if n, err := strconv.Atoi(pageSize); err == nil && n > 0 && n <= 100 {
limit = n
}
}
if page := c.Query("page"); page != "" {
if p, err := strconv.Atoi(page); err == nil && p > 0 {
offset = (p - 1) * limit
}
}
return limit, offset
}
func daoUsageFilter(c *gin.Context) dao.UsageFilter {
limit, offset := paginate(c, 20)
f := dao.UsageFilter{Limit: limit, Offset: offset}
f.Protocol = c.Query("protocol")
f.Status = c.Query("status")
f.ModelName = c.Query("model")
if v := c.Query("user_id"); v != "" {
if n, err := strconv.ParseUint(v, 10, 64); err == nil && n > 0 {
f.UserID = &n
}
}
return f
}
func usageLogsToResp(logs []store.UsageLog, names map[uint64]string) []gin.H {
out := make([]gin.H, 0, len(logs))
for _, l := range logs {
row := gin.H{
"id": l.ID,
"request_id": l.RequestID,
"user_id": l.UserID,
"channel_id": l.ChannelID,
"model_id": l.ModelID,
"model_name": l.ModelName,
"protocol": l.Protocol,
"input_tokens": l.InputTokens,
"output_tokens": l.OutputTokens,
"cache_read_tokens": l.CacheReadTokens,
"cache_creation_tokens": l.CacheCreationTokens,
"cost": l.Cost,
"latency_ms": l.LatencyMS,
"status": l.Status,
"error_code": l.ErrorCode,
"created_at": l.CreatedAt,
}
if names != nil {
if u, ok := names[l.UserID]; ok {
row["username"] = u
}
}
if l.RawRequest != "" {
row["raw_request"] = l.RawRequest
}
if l.RawResponse != "" {
row["raw_response"] = l.RawResponse
}
out = append(out, row)
}
return out
}
// userNames 批量查询 user_id → username 映射。
func (h *Handler) userNames(logs []store.UsageLog) map[uint64]string {
ids := map[uint64]bool{}
for _, l := range logs {
ids[l.UserID] = true
}
if len(ids) == 0 {
return nil
}
idList := make([]uint64, 0, len(ids))
for id := range ids {
idList = append(idList, id)
}
var users []store.User
if err := h.db.Where("id IN ?", idList).Find(&users).Error; err != nil {
return nil
}
out := map[uint64]string{}
for _, u := range users {
out[u.ID] = u.Username
}
return out
}
+135 -64
View File
@@ -2,8 +2,6 @@ package channel
import ( import (
"context" "context"
"fmt"
"log"
"math/rand" "math/rand"
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/store" "opencatd-open/internal/store"
@@ -19,6 +17,9 @@ type Service struct {
// Health tracking // Health tracking
mu sync.RWMutex mu sync.RWMutex
healthStatus map[uint64]*channelHealth healthStatus map[uint64]*channelHealth
// Concurrency control per channel
sems map[uint64]chan struct{}
} }
type channelHealth struct { type channelHealth struct {
@@ -33,44 +34,78 @@ func NewService(channelDAO *dao.ChannelDAO, modelDAO *dao.ModelDAO) *Service {
channelDAO: channelDAO, channelDAO: channelDAO,
modelDAO: modelDAO, modelDAO: modelDAO,
healthStatus: make(map[uint64]*channelHealth), healthStatus: make(map[uint64]*channelHealth),
sems: make(map[uint64]chan struct{}),
} }
} }
// SelectChannel selects the best channel for a given model using weighted random selection // SelectedRoute 一次路由决策的完整结果:渠道 + 命中的模型绑定。
func (s *Service) SelectChannel(ctx context.Context, modelName string) (*store.Channel, error) { // Binding 可能为 nil(渠道经回退路径选中、无绑定记录)。
channels, err := s.channelDAO.GetEnabledChannelsByModel(modelName) type SelectedRoute struct {
if err != nil { Channel *store.Channel
return nil, fmt.Errorf("failed to get channels for model %s: %w", modelName, err) Binding *store.ChannelModelBinding
} }
if len(channels) == 0 {
return nil, fmt.Errorf("no enabled channels for model: %s", modelName)
}
// Filter out unhealthy channels // Candidates 返回可用渠道候选:健康 + 启用。
candidates := s.filterHealthy(channels) // model 非空时优先取绑定该模型的渠道(携带 upstream_model 映射,权重降序);
if len(candidates) == 0 { // 无绑定则回退到未绑定模型路径:按权重升序(闲置渠道优先探活)。
// If all channels are unhealthy, try the first one anyway func (s *Service) Candidates(model string) []Candidate {
candidates = channels[:1] if model != "" {
var b []store.ChannelModelBinding
var modelIDs []uint64
s.modelDAO.DB().Model(&store.Model{}).Where("name = ? AND enabled = ?", model, true).Pluck("id", &modelIDs)
if len(modelIDs) > 0 {
s.channelDAO.DB().Where("model_id IN ?", modelIDs).Find(&b)
if cands := s.loadBound(b); len(cands) > 0 {
return cands
} }
}
}
// 未绑定模型回退:取优先级最低的空闲健康渠道作为"备用渠道"承接搭车流量
// (排序与绑定候选一致:priority ASC, weight DESC, id ASC,取末位)。
// weight=0 的渠道不被加权随机选中,但可作为最后备用承接 unbound 流量。
var chs []store.Channel
s.channelDAO.DB().Where("enabled = ?", true).
Order("priority ASC, weight DESC, id ASC").Find(&chs)
all := make([]Candidate, 0, len(chs))
for i := range chs {
all = append(all, Candidate{Channel: &chs[i]})
}
healthy := s.FilterHealthy(all)
if len(healthy) == 0 {
return nil
}
return healthy[len(healthy)-1:]
}
// Weighted random selection // loadBound 按绑定顺序加载渠道候选,过滤健康/启用,携带 upstream_model 映射。
totalWeight := 0 func (s *Service) loadBound(bindings []store.ChannelModelBinding) []Candidate {
for _, ch := range candidates { if len(bindings) == 0 {
totalWeight += ch.Weight return nil
} }
if totalWeight == 0 { // channel_id -> 绑定(取该渠道对该模型的映射)
return candidates[0], nil byChannel := map[uint64]store.ChannelModelBinding{}
ids := make([]uint64, 0, len(bindings))
for _, b := range bindings {
if _, ok := byChannel[b.ChannelID]; !ok {
ids = append(ids, b.ChannelID)
} }
byChannel[b.ChannelID] = b
r := rand.Intn(totalWeight) }
for _, ch := range candidates { var chs []store.Channel
r -= ch.Weight s.channelDAO.DB().Where("id IN ? AND enabled = ? AND health_status = ?", ids, true, store.ChannelHealthHealthy).
if r < 0 { Order("priority ASC, weight DESC, id ASC").Find(&chs)
return ch, nil byID := map[uint64]*store.Channel{}
for i := range chs {
byID[chs[i].ID] = &chs[i]
}
out := make([]Candidate, 0, len(ids))
for _, id := range ids {
if ch, ok := byID[id]; ok {
b := byChannel[id]
out = append(out, Candidate{Channel: ch, Binding: &b})
} }
} }
return out
return candidates[0], nil
} }
// GetChannelByKeyID decrypts the API key for a channel // GetChannelByKeyID decrypts the API key for a channel
@@ -98,7 +133,9 @@ func (s *Service) RecordSuccess(channelID uint64) {
h.lastCheck = time.Now() h.lastCheck = time.Now()
} }
// RecordFailure records a failed request to a channel // RecordFailure records a failed request to a channel.
// 连续 2 次失败进入 degraded(快速熔断):失败过的渠道让位给健康渠道,
// 健康检查成功或冷却过期后复位。
func (s *Service) RecordFailure(channelID uint64) { func (s *Service) RecordFailure(channelID uint64) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
@@ -107,7 +144,7 @@ func (s *Service) RecordFailure(channelID uint64) {
h.consecutive++ h.consecutive++
h.lastCheck = time.Now() h.lastCheck = time.Now()
if h.consecutive >= 3 { if h.consecutive >= 2 {
h.status = store.ChannelHealthDegraded h.status = store.ChannelHealthDegraded
h.cooldown = time.Now().Add(5 * time.Minute) h.cooldown = time.Now().Add(5 * time.Minute)
} }
@@ -175,47 +212,81 @@ func (s *Service) GetHealthStatus(channelID uint64) string {
return h.status return h.status
} }
// ChannelCandidate represents a channel with its resolved API key // Candidate 一个候选渠道 + 该模型的映射关系。
type ChannelCandidate struct { type Candidate struct {
Channel *store.Channel Channel *store.Channel
APIKey string Binding *store.ChannelModelBinding // 全局模型在此渠道的映射(无绑定则 nil)
Format string
} }
// SelectCandidates returns candidates for a model, sorted by priority // Pick 按权重加权随机选一个候选渠道(负载均衡;weight<=0 按 1 计)。
func (s *Service) SelectCandidates(ctx context.Context, modelName string, preferredFormat string) ([]ChannelCandidate, error) { func (s *Service) Pick(cands []Candidate) *Candidate {
channels, err := s.channelDAO.GetEnabledChannelsByModel(modelName) if len(cands) == 0 {
if err != nil { return nil
return nil, err
} }
total := 0
var candidates []ChannelCandidate for _, c := range cands {
for _, ch := range channels { w := c.Channel.Weight
// Check if channel supports the preferred format if w <= 0 {
formats := ch.FormatsEffective() w = 1
supported := false }
for _, f := range formats { total += w
if f == preferredFormat || preferredFormat == "" { }
supported = true r := rand.Intn(total)
break acc := 0
for i := range cands {
w := cands[i].Channel.Weight
if w <= 0 {
w = 1
}
acc += w
if r < acc {
return &cands[i]
} }
} }
if !supported { return &cands[len(cands)-1]
}
// FilterHealthy 过滤掉内存健康状态异常的渠道候选(degraded/cooldown 均排除,
// 冷却/降级过期后复位放行)。degraded 由单次请求失败触发,作为快速熔断:
// 后续请求先走其他渠道,健康检查成功后恢复。
func (s *Service) FilterHealthy(cands []Candidate) []Candidate {
s.mu.RLock()
defer s.mu.RUnlock()
out := make([]Candidate, 0, len(cands))
now := time.Now()
for _, c := range cands {
h, ok := s.healthStatus[c.Channel.ID]
if !ok || h.status == store.ChannelHealthHealthy {
out = append(out, c)
continue continue
}
// 冷却/降级已过期:复位并放行
if !h.cooldown.IsZero() && now.After(h.cooldown) {
h.status = store.ChannelHealthHealthy
h.consecutive = 0
out = append(out, c)
} }
}
return out
}
apiKey, err := crypto.Decrypt(ch.APIKeyEnc) // TryAcquire 尝试获取渠道并发槽;渠道满载返回 false(调用方可溢出到其他渠道)。
if err != nil { // MaxConcurrency<=0 视为不限制。
log.Printf("Failed to decrypt API key for channel %s: %v", ch.Name, err) func (s *Service) TryAcquire(ch *store.Channel) (func(), bool) {
continue if ch.MaxConcurrency <= 0 {
return func() {}, true
}
s.mu.Lock()
sem, ok := s.sems[ch.ID]
if !ok {
sem = make(chan struct{}, ch.MaxConcurrency)
s.sems[ch.ID] = sem
} }
s.mu.Unlock()
candidates = append(candidates, ChannelCandidate{ select {
Channel: ch, case sem <- struct{}{}:
APIKey: apiKey, return func() { <-sem }, true
Format: preferredFormat, default:
}) return nil, false
} }
return candidates, nil
} }
+34 -4
View File
@@ -10,19 +10,45 @@ import (
"time" "time"
) )
// HealthConfig 健康检查配置
type HealthConfig struct {
Interval time.Duration // 检查间隔
Timeout time.Duration // 请求超时
FailureThreshold int // 连续失败次数阈值
DegradedCooldown time.Duration // degraded 冷却时间
CooldownCooldown time.Duration // cooldown 冷却时间
}
// DefaultHealthConfig 返回默认健康检查配置
func DefaultHealthConfig() HealthConfig {
return HealthConfig{
Interval: 5 * time.Minute,
Timeout: 10 * time.Second,
FailureThreshold: 3,
DegradedCooldown: 5 * time.Minute,
CooldownCooldown: 15 * time.Minute,
}
}
type HealthChecker struct { type HealthChecker struct {
channelDAO *dao.ChannelDAO channelDAO *dao.ChannelDAO
service *Service service *Service
client *http.Client client *http.Client
config HealthConfig
} }
func NewHealthChecker(channelDAO *dao.ChannelDAO, service *Service) *HealthChecker { func NewHealthChecker(channelDAO *dao.ChannelDAO, service *Service, config ...HealthConfig) *HealthChecker {
cfg := DefaultHealthConfig()
if len(config) > 0 {
cfg = config[0]
}
return &HealthChecker{ return &HealthChecker{
channelDAO: channelDAO, channelDAO: channelDAO,
service: service, service: service,
client: &http.Client{ client: &http.Client{
Timeout: 10 * time.Second, Timeout: cfg.Timeout,
}, },
config: cfg,
} }
} }
@@ -91,8 +117,12 @@ func (hc *HealthChecker) CheckAllChannels(ctx context.Context) error {
} }
// StartPeriodicCheck starts periodic health checks // StartPeriodicCheck starts periodic health checks
func (hc *HealthChecker) StartPeriodicCheck(ctx context.Context, interval time.Duration) { func (hc *HealthChecker) StartPeriodicCheck(ctx context.Context, interval ...time.Duration) {
ticker := time.NewTicker(interval) interval_ := hc.config.Interval
if len(interval) > 0 {
interval_ = interval[0]
}
ticker := time.NewTicker(interval_)
defer ticker.Stop() defer ticker.Stop()
for { for {
+1 -3
View File
@@ -13,18 +13,16 @@ type Api struct {
userService *service.UserServiceImpl userService *service.UserServiceImpl
tokenService *service.TokenServiceImpl tokenService *service.TokenServiceImpl
keyService *service.ApiKeyServiceImpl keyService *service.ApiKeyServiceImpl
webAuthService *service.WebAuthnService
usageService *service.UsageService usageService *service.UsageService
} }
func NewApi(cfg *config.Config, db *gorm.DB, userService *service.UserServiceImpl, tokenService *service.TokenServiceImpl, keyService *service.ApiKeyServiceImpl, webAuthService *service.WebAuthnService, usageService *service.UsageService) *Api { func NewApi(cfg *config.Config, db *gorm.DB, userService *service.UserServiceImpl, tokenService *service.TokenServiceImpl, keyService *service.ApiKeyServiceImpl, usageService *service.UsageService) *Api {
return &Api{ return &Api{
cfg: cfg, cfg: cfg,
db: db, db: db,
userService: userService, userService: userService,
tokenService: tokenService, tokenService: tokenService,
keyService: keyService, keyService: keyService,
webAuthService: webAuthService,
usageService: usageService, usageService: usageService,
} }
} }
+6 -1
View File
@@ -97,7 +97,12 @@ func (p *Proxy) SelectChannel(modelName string) (*store.Channel, error) {
if p.channelSvc == nil { if p.channelSvc == nil {
return nil, fmt.Errorf("channel service not initialized") return nil, fmt.Errorf("channel service not initialized")
} }
return p.channelSvc.SelectChannel(p.ctx, modelName) cands := p.channelSvc.Candidates(modelName)
picked := p.channelSvc.Pick(cands)
if picked == nil {
return nil, fmt.Errorf("no enabled channels for model: %s", modelName)
}
return picked.Channel, nil
} }
// RecordSuccess records a successful request // RecordSuccess records a successful request
+5
View File
@@ -14,6 +14,11 @@ func NewChannelDAO(db *gorm.DB) *ChannelDAO {
return &ChannelDAO{db: db} return &ChannelDAO{db: db}
} }
// DB 暴露底层连接,供聚合查询使用(如渠道候选联表过滤)。
func (d *ChannelDAO) DB() *gorm.DB {
return d.db
}
func (d *ChannelDAO) Create(channel *store.Channel) error { func (d *ChannelDAO) Create(channel *store.Channel) error {
return d.db.Create(channel).Error return d.db.Create(channel).Error
} }
+5
View File
@@ -14,6 +14,11 @@ func NewModelDAO(db *gorm.DB) *ModelDAO {
return &ModelDAO{db: db} return &ModelDAO{db: db}
} }
// DB 暴露底层连接,供聚合查询使用(如模型候选联表过滤)。
func (d *ModelDAO) DB() *gorm.DB {
return d.db
}
func (d *ModelDAO) Create(model *store.Model) error { func (d *ModelDAO) Create(model *store.Model) error {
return d.db.Create(model).Error return d.db.Create(model).Error
} }
+71 -1
View File
@@ -55,6 +55,50 @@ func (d *UsageDAO) CountByUserID(ctx context.Context, userID uint64) (int64, err
return count, err return count, err
} }
// UsageFilter 用量明细筛选条件(管理后台)。
type UsageFilter struct {
UserID *uint64 // 指定用户(nil=全部)
Protocol string // 协议 chat/messages/responses(空=全部)
Status string // success/error/canceled(空=全部)
ModelName string // 模型名模糊(空=全部)
Limit int
Offset int
}
// ListAll 管理后台全量用量明细(分页 + 筛选),并带用户名。
func (d *UsageDAO) ListAll(ctx context.Context, f UsageFilter) ([]store.UsageLog, error) {
q := d.db.WithContext(ctx).Model(&store.UsageLog{})
q = applyUsageFilter(q, f)
var logs []store.UsageLog
err := q.Order("created_at DESC").Limit(f.Limit).Offset(f.Offset).Find(&logs).Error
return logs, err
}
// CountAll 统计符合筛选条件的明细总数。
func (d *UsageDAO) CountAll(ctx context.Context, f UsageFilter) (int64, error) {
q := d.db.WithContext(ctx).Model(&store.UsageLog{})
q = applyUsageFilter(q, f)
var count int64
err := q.Count(&count).Error
return count, err
}
func applyUsageFilter(q *gorm.DB, f UsageFilter) *gorm.DB {
if f.UserID != nil {
q = q.Where("user_id = ?", *f.UserID)
}
if f.Protocol != "" {
q = q.Where("protocol = ?", f.Protocol)
}
if f.Status != "" {
q = q.Where("status = ?", f.Status)
}
if f.ModelName != "" {
q = q.Where("model_name LIKE ?", "%"+f.ModelName+"%")
}
return q
}
// UsageDaily DAO // UsageDaily DAO
func (d *DailyUsageDAO) Create(ctx context.Context, log *store.UsageDaily) error { func (d *DailyUsageDAO) Create(ctx context.Context, log *store.UsageDaily) error {
return d.db.WithContext(ctx).Create(log).Error return d.db.WithContext(ctx).Create(log).Error
@@ -82,10 +126,19 @@ func (d *DailyUsageDAO) GetByDate(ctx context.Context, userID uint64, date strin
return &log, nil return &log, nil
} }
// UpsertDailyUsage 按 (user_id, model_id, date) 累加式 upsert:
// 行不存在则插入;存在则在原值基础上增量累加(不能用 AssignmentColumns 覆盖,
// 否则多次 flush 会互相清零)。非限定列名在 SQLite/MySQL/PG 的 upsert 语义下都指向目标行。
func (d *DailyUsageDAO) UpsertDailyUsage(ctx context.Context, log *store.UsageDaily) error { func (d *DailyUsageDAO) UpsertDailyUsage(ctx context.Context, log *store.UsageDaily) error {
return d.db.WithContext(ctx).Clauses(clause.OnConflict{ return d.db.WithContext(ctx).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "user_id"}, {Name: "model_id"}, {Name: "date"}}, Columns: []clause.Column{{Name: "user_id"}, {Name: "model_id"}, {Name: "date"}},
DoUpdates: clause.AssignmentColumns([]string{"requests", "input_tokens", "output_tokens", "cache_read_tokens", "cost"}), DoUpdates: clause.Assignments(map[string]interface{}{
"requests": gorm.Expr("requests + ?", log.Requests),
"input_tokens": gorm.Expr("input_tokens + ?", log.InputTokens),
"output_tokens": gorm.Expr("output_tokens + ?", log.OutputTokens),
"cache_read_tokens": gorm.Expr("cache_read_tokens + ?", log.CacheReadTokens),
"cost": gorm.Expr("cost + ?", log.Cost),
}),
}).Create(log).Error }).Create(log).Error
} }
@@ -97,3 +150,20 @@ func (d *DailyUsageDAO) ListByDateRange(ctx context.Context, userID uint64, star
Find(&logs).Error Find(&logs).Error
return logs, err return logs, err
} }
// ListAll 管理后台:全部用户的日聚合(可选按用户/日期范围筛选),按日期倒序。
func (d *DailyUsageDAO) ListAll(ctx context.Context, userID *uint64, start, end string) ([]store.UsageDaily, error) {
q := d.db.WithContext(ctx).Model(&store.UsageDaily{})
if userID != nil {
q = q.Where("user_id = ?", *userID)
}
if start != "" {
q = q.Where("date >= ?", start)
}
if end != "" {
q = q.Where("date <= ?", end)
}
var logs []store.UsageDaily
err := q.Order("date DESC, user_id ASC").Find(&logs).Error
return logs, err
}
-11
View File
@@ -1,11 +0,0 @@
package dto
type Passkey struct {
ID int64 `json:"id" gorm:"column:id;primaryKey;autoIncrement"`
Name string `json:"name" gorm:"column:name"` // 凭证名称,用于用户识别不同的设备
SignCount uint32 `json:"sign_count" gorm:"column:sign_count"` // 签名计数器,用于防止重放攻击
DeviceType string `json:"device_type" gorm:"column:device_type"` // 设备类型,如"platform"或"cross-platform"
LastUsedAt int64 `json:"last_used_at" gorm:"column:last_used_at"` // 最后使用时间
CreatedAt int64 `json:"created_at,omitempty" gorm:"autoCreateTime"`
UpdatedAt int64 `json:"updated_at,omitempty" gorm:"autoUpdateTime"`
}
+367
View File
@@ -0,0 +1,367 @@
// Package passkey 封装 WebAuthn(passkey)注册与登录。
package passkey
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
"sync"
"time"
"github.com/go-webauthn/webauthn/protocol"
"github.com/go-webauthn/webauthn/webauthn"
"github.com/redis/go-redis/v9"
"opencatd-open/internal/store"
"gorm.io/gorm"
)
const (
sessionPrefix = "passkey:session:"
sessionTTL = 5 * time.Minute
)
type Config struct {
RPID string
Origins []string
Name string
Redis *redis.Client // 可选,nil 时使用内存存储
}
// SessionStore challenge 会话存储接口
type SessionStore interface {
Set(ctx context.Context, session *webauthn.SessionData) error
Get(ctx context.Context, challenge string) (*webauthn.SessionData, bool, error)
Delete(ctx context.Context, challenge string) error
}
// memoryStore 内存存储(单实例)
type memoryStore struct {
mu sync.Mutex
sessions map[string]webauthn.SessionData
}
func newMemoryStore() *memoryStore {
return &memoryStore{sessions: make(map[string]webauthn.SessionData)}
}
func (m *memoryStore) Set(_ context.Context, session *webauthn.SessionData) error {
m.mu.Lock()
m.sessions[session.Challenge] = *session
m.mu.Unlock()
return nil
}
func (m *memoryStore) Get(_ context.Context, challenge string) (*webauthn.SessionData, bool, error) {
m.mu.Lock()
sess, ok := m.sessions[challenge]
m.mu.Unlock()
if !ok {
return nil, false, nil
}
// 检查过期
if !sess.Expires.IsZero() && time.Now().After(sess.Expires) {
return nil, false, nil
}
return &sess, true, nil
}
func (m *memoryStore) Delete(_ context.Context, challenge string) error {
m.mu.Lock()
delete(m.sessions, challenge)
m.mu.Unlock()
return nil
}
// redisStore Redis 存储(分布式)
type redisStore struct {
rdb *redis.Client
}
func newRedisStore(rdb *redis.Client) *redisStore {
return &redisStore{rdb: rdb}
}
func (r *redisStore) Set(ctx context.Context, session *webauthn.SessionData) error {
data, err := json.Marshal(session)
if err != nil {
return fmt.Errorf("marshal session: %w", err)
}
key := sessionPrefix + session.Challenge
return r.rdb.Set(ctx, key, data, sessionTTL).Err()
}
func (r *redisStore) Get(ctx context.Context, challenge string) (*webauthn.SessionData, bool, error) {
key := sessionPrefix + challenge
data, err := r.rdb.Get(ctx, key).Bytes()
if err == redis.Nil {
return nil, false, nil
}
if err != nil {
return nil, false, fmt.Errorf("redis get: %w", err)
}
var sess webauthn.SessionData
if err := json.Unmarshal(data, &sess); err != nil {
return nil, false, fmt.Errorf("unmarshal session: %w", err)
}
return &sess, true, nil
}
func (r *redisStore) Delete(ctx context.Context, challenge string) error {
key := sessionPrefix + challenge
return r.rdb.Del(ctx, key).Err()
}
// Service WebAuthn 服务:凭据存储 + challenge 会话。
type Service struct {
wa *webauthn.WebAuthn
db *gorm.DB
sessions SessionStore
}
func New(db *gorm.DB, cfg Config) (*Service, error) {
wa, err := webauthn.New(&webauthn.Config{
RPDisplayName: cfg.Name,
RPID: cfg.RPID,
RPOrigins: cfg.Origins,
})
if err != nil {
return nil, err
}
// 根据配置选择存储后端
var store SessionStore
if cfg.Redis != nil {
store = newRedisStore(cfg.Redis)
} else {
store = newMemoryStore()
}
return &Service{wa: wa, db: db, sessions: store}, nil
}
// webUser 实现 go-webauthn 的 User 接口。
type webUser struct {
id uint64
name string
displayName string
credentials []webauthn.Credential
}
func (u *webUser) WebAuthnID() []byte { return []byte(strconv.FormatUint(u.id, 10)) }
func (u *webUser) WebAuthnName() string { return u.name }
func (u *webUser) WebAuthnDisplayName() string { return u.displayName }
func (u *webUser) WebAuthnIcon() string { return "" }
func (u *webUser) WebAuthnCredentials() []webauthn.Credential { return u.credentials }
func (s *Service) loadWebUser(u *store.User) (*webUser, error) {
var pks []store.Passkey
s.db.Where("user_id = ?", u.ID).Find(&pks)
creds := make([]webauthn.Credential, 0, len(pks))
for _, pk := range pks {
var c webauthn.Credential
if err := json.Unmarshal(pk.Credential, &c); err == nil {
creds = append(creds, c)
}
}
return &webUser{id: u.ID, name: u.Username, displayName: u.Username, credentials: creds}, nil
}
// GetUserByUsername 通过用户名或邮箱查找用户
func (s *Service) GetUserByUsername(username string) (*store.User, error) {
var u store.User
if err := s.db.Where("username = ? OR email = ?", username, username).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// GetUserByID 通过 ID 查找用户
func (s *Service) GetUserByID(id uint64) (*store.User, error) {
var u store.User
if err := s.db.First(&u, id).Error; err != nil {
return nil, err
}
return &u, nil
}
// ---------------------------------------------------------------------------
// 注册
// BeginRegistration 生成注册选项并暂存 challenge。
func (s *Service) BeginRegistration(u *store.User) (*protocol.CredentialCreation, error) {
wu, err := s.loadWebUser(u)
if err != nil {
return nil, err
}
creation, session, err := s.wa.BeginRegistration(wu)
if err != nil {
return nil, err
}
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return creation, nil
}
// FinishRegistration 校验浏览器返回的凭据并落库。
func (s *Service) FinishRegistration(u *store.User, challenge string, body []byte, name string) error {
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return err
}
if !ok {
return errors.New("challenge 已过期或不存在")
}
// 删除已使用的 challenge
_ = s.sessions.Delete(context.Background(), challenge)
wu, err := s.loadWebUser(u)
if err != nil {
return err
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
cred, err := s.wa.FinishRegistration(wu, *session, req)
if err != nil {
return err
}
raw, _ := json.Marshal(cred)
nm := name
if nm == "" {
nm = "passkey"
}
return s.db.Create(&store.Passkey{
UserID: u.ID, Name: nm, CredentialID: cred.ID, Credential: raw,
}).Error
}
// ---------------------------------------------------------------------------
// 登录
// BeginLogin 已知用户(按用户名)发起断言。
func (s *Service) BeginLogin(u *store.User) (*protocol.CredentialAssertion, error) {
wu, err := s.loadWebUser(u)
if err != nil {
return nil, err
}
assertion, session, err := s.wa.BeginLogin(wu)
if err != nil {
return nil, err
}
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return assertion, nil
}
// BeginDiscoverableLogin 无用户名(使用平台/漫游器上的可发现凭据)。
func (s *Service) BeginDiscoverableLogin() (*protocol.CredentialAssertion, error) {
assertion, session, err := s.wa.BeginDiscoverableLogin()
if err != nil {
return nil, err
}
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return assertion, nil
}
// FinishLogin 校验断言并更新签名计数。
func (s *Service) FinishLogin(u *store.User, challenge string, body []byte) error {
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return err
}
if !ok {
return errors.New("challenge 已过期或不存在")
}
_ = s.sessions.Delete(context.Background(), challenge)
wu, err := s.loadWebUser(u)
if err != nil {
return err
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
cred, err := s.wa.FinishLogin(wu, *session, req)
if err != nil {
return err
}
return s.updateCredential(u.ID, cred)
}
// FinishDiscoverableLogin 通过凭据定位用户并校验断言。
func (s *Service) FinishDiscoverableLogin(challenge string, body []byte) (*store.User, error) {
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return nil, err
}
if !ok {
return nil, errors.New("challenge 已过期或不存在")
}
_ = s.sessions.Delete(context.Background(), challenge)
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
// 先通过 credential_id 反查用户
var pks []store.Passkey
if err := s.db.Find(&pks).Error; err != nil {
return nil, err
}
// 遍历所有 passkey 找到匹配的
for _, pk := range pks {
var c webauthn.Credential
if err := json.Unmarshal(pk.Credential, &c); err != nil {
continue
}
// 尝试用这个用户的凭据进行登录
var u store.User
if err := s.db.First(&u, pk.UserID).Error; err != nil {
continue
}
wu, err := s.loadWebUser(&u)
if err != nil {
continue
}
cred, err := s.wa.FinishLogin(wu, *session, req)
if err != nil {
continue
}
_ = s.updateCredential(u.ID, cred)
return &u, nil
}
return nil, errors.New("no matching passkey found")
}
// ---------------------------------------------------------------------------
// 管理
// List 列出用户的 passkey。
func (s *Service) List(userID uint64) ([]store.Passkey, error) {
var pks []store.Passkey
err := s.db.Where("user_id = ?", userID).Order("id DESC").Find(&pks).Error
return pks, err
}
// Delete 删除用户的 passkey。
func (s *Service) Delete(userID, id uint64) error {
res := s.db.Where("id = ? AND user_id = ?", id, userID).Delete(&store.Passkey{})
if res.Error != nil {
return res.Error
}
if res.RowsAffected == 0 {
return gorm.ErrRecordNotFound
}
return nil
}
func (s *Service) updateCredential(userID uint64, cred *webauthn.Credential) error {
raw, _ := json.Marshal(cred)
return s.db.Model(&store.Passkey{}).
Where("user_id = ? AND credential_id = ?", userID, cred.ID).
Update("credential", raw).Error
}
@@ -30,7 +30,7 @@ func ChatToResponses(req *ChatCompletionRequest) (*ResponsesRequest, error) {
out := &ResponsesRequest{ out := &ResponsesRequest{
Model: req.Model, Model: req.Model,
Input: inputItems, Input: marshalInputItems(inputItems),
Instructions: instructions, Instructions: instructions,
Stream: req.Stream, Stream: req.Stream,
} }
@@ -220,7 +220,7 @@ func MessagesToResponses(req *MessagesRequest) (*ResponsesRequest, error) {
out := &ResponsesRequest{ out := &ResponsesRequest{
Model: req.Model, Model: req.Model,
Input: inputItems, Input: marshalInputItems(inputItems),
Instructions: instructions, Instructions: instructions,
Stream: req.Stream, Stream: req.Stream,
} }
+203
View File
@@ -0,0 +1,203 @@
// 三协议互转注册表:OpenAI Chat / OpenAI Responses / Anthropic Messages。
// 网关以 Chat 形状作为标准中间模型:非跨 chat 的转换经 chat 中转。
// 请求/响应(非流式)走 JSON 转换;流式走逐行 SSE 转换(stream_transform.go)。
package convert
import (
"bytes"
"encoding/json"
"fmt"
)
// 协议标识。
const (
ProtoChat = "chat"
ProtoMessages = "messages"
ProtoResponses = "responses"
)
// trimBody 去掉首尾空白。部分上游(如 OpenRouter)会在 JSON 前输出空白或
// SSE 注释行再跟正文,直接 Unmarshal 会失败。
func trimBody(body []byte) []byte {
return bytes.TrimSpace(body)
}
// CleanJSON 剥离非 JSON 前缀(空白、SSE 注释、`data:` 行)并压缩为标准 JSON。
// 部分上游(如 OpenRouter)的 non-stream 响应在 JSON 前夹带空白/注释;
// 原样透传会让客户端解析失败。找不到 JSON 对象时原样返回。
func CleanJSON(body []byte) []byte {
i := bytes.IndexByte(body, '{')
if i < 0 {
return body
}
var v any
if err := json.Unmarshal(bytes.TrimSpace(body[i:]), &v); err != nil {
return body
}
out, err := json.Marshal(v)
if err != nil {
return body
}
return out
}
// ConvertRequest 转换请求体。from==to 时原样返回。
func ConvertRequest(body []byte, from, to string) ([]byte, error) {
if from == to {
return body, nil
}
body = trimBody(body)
switch {
case from == ProtoMessages && to == ProtoChat:
return messagesToChatReq(body)
case from == ProtoChat && to == ProtoMessages:
return chatToMessagesReq(body)
case from == ProtoResponses && to == ProtoChat:
return responsesToChatReq(body)
case from == ProtoChat && to == ProtoResponses:
return chatToResponsesReq(body)
case from == ProtoResponses && to == ProtoMessages:
mid, err := responsesToChatReq(body)
if err != nil {
return nil, err
}
return chatToMessagesReq(mid)
case from == ProtoMessages && to == ProtoResponses:
mid, err := messagesToChatReq(body)
if err != nil {
return nil, err
}
return chatToResponsesReq(mid)
}
return nil, fmt.Errorf("unsupported request conversion %s->%s", from, to)
}
// ConvertResponse 转换响应体(非流式)。from==to 时原样返回。
func ConvertResponse(body []byte, from, to string) ([]byte, error) {
if from == to {
return body, nil
}
body = trimBody(body)
switch {
case from == ProtoMessages && to == ProtoChat:
return messagesToChatResp(body)
case from == ProtoChat && to == ProtoMessages:
return chatToMessagesResp(body)
case from == ProtoResponses && to == ProtoChat:
return responsesToChatResp(body)
case from == ProtoChat && to == ProtoResponses:
return chatToResponsesResp(body)
case from == ProtoResponses && to == ProtoMessages:
mid, err := responsesToChatResp(body)
if err != nil {
return nil, err
}
return chatToMessagesResp(mid)
case from == ProtoMessages && to == ProtoResponses:
mid, err := messagesToChatResp(body)
if err != nil {
return nil, err
}
return chatToResponsesResp(mid)
}
return nil, fmt.Errorf("unsupported response conversion %s->%s", from, to)
}
// NewStreamTransformer 构造流式逐行转换器:输入上游 SSE 一行,返回客户端 SSE 行。
// 返回 nil 表示丢弃该行或无需转换(from==to)。
func NewStreamTransformer(from, to string) func([]byte) []byte {
switch {
case from == ProtoMessages && to == ProtoChat:
return newMessagesToChat().line
case from == ProtoChat && to == ProtoMessages:
return newChatToMessages().line
case from == ProtoResponses && to == ProtoChat:
return newResponsesToChat().line
case from == ProtoChat && to == ProtoResponses:
return newChatToResponses().line
case from == ProtoResponses && to == ProtoMessages:
return newResponsesToMessages().line
case from == ProtoMessages && to == ProtoResponses:
return newMessagesToResponses().line
}
return nil
}
// ---------------------------------------------------------------------------
// 工具函数
// str 返回字符串字段;json.RawMessage 为字符串字面量时去引号。
func str(raw json.RawMessage) string {
if len(raw) == 0 || string(raw) == "null" {
return ""
}
var s string
if json.Unmarshal(raw, &s) == nil {
return s
}
// 数组/对象:尝试取 type=text 的 text
var arr []map[string]any
if json.Unmarshal(raw, &arr) == nil {
var parts []string
for _, b := range arr {
if t, _ := b["type"].(string); t == "text" || t == "input_text" || t == "output_text" {
if txt, _ := b["text"].(string); txt != "" {
parts = append(parts, txt)
}
}
}
return joinNonEmpty(parts, "\n")
}
return ""
}
func joinNonEmpty(parts []string, sep string) string {
out := ""
for _, p := range parts {
if p == "" {
continue
}
if out != "" {
out += sep
}
out += p
}
return out
}
// rawJSON 安全取字段;不存在或 null 返回 nil。
func rawJSON(m map[string]json.RawMessage, key string) json.RawMessage {
raw, ok := m[key]
if !ok || string(raw) == "null" {
return nil
}
return raw
}
// rawOrObject 把 RawMessage 解为 map;非对象返回空对象。
func rawOrObject(raw json.RawMessage) any {
if len(raw) == 0 || string(raw) == "null" {
return map[string]any{}
}
var m map[string]any
if json.Unmarshal(raw, &m) == nil {
return m
}
return map[string]any{}
}
// intOrNil 取指针值,nil 时返回默认值。
func intOrNil(p *int, def int) any {
if p == nil {
return def
}
return *p
}
// strField 取 any 中的字符串字段。
func strField(v any) string {
if s, ok := v.(string); ok {
return s
}
return ""
}
+12 -5
View File
@@ -1,6 +1,7 @@
package convert package convert
import ( import (
"encoding/json"
"testing" "testing"
) )
@@ -114,12 +115,18 @@ func TestChatToResponses(t *testing.T) {
t.Errorf("Model = %q, want %q", result.Model, "gpt-4o") t.Errorf("Model = %q, want %q", result.Model, "gpt-4o")
} }
if len(result.Input) != 1 { if len(result.Input) == 0 {
t.Errorf("Input length = %d, want 1", len(result.Input)) t.Errorf("Input empty, want 1 item")
} else {
var items []InputItem
if err := json.Unmarshal(result.Input, &items); err != nil {
t.Fatalf("Input unmarshal = %v", err)
}
if len(items) != 1 {
t.Errorf("Input length = %d, want 1", len(items))
} else if items[0].Role != "user" {
t.Errorf("Input[0].Role = %q, want %q", items[0].Role, "user")
} }
if result.Input[0].Role != "user" {
t.Errorf("Input[0].Role = %q, want %q", result.Input[0].Role, "user")
} }
if result.Instructions != "You are a helpful assistant." { if result.Instructions != "You are a helpful assistant." {
+504
View File
@@ -0,0 +1,504 @@
package convert
import (
"encoding/json"
"strings"
)
// ---------------------------------------------------------------------------
// 请求:Chat → Messages
type chatTool struct {
Type string `json:"type"`
Function struct {
Name string `json:"name"`
Description string `json:"description"`
Parameters json.RawMessage `json:"parameters"`
} `json:"function"`
}
type chatMsg struct {
Role string `json:"role"`
Content json.RawMessage `json:"content"`
ToolCallID string `json:"tool_call_id"`
ToolCalls []struct {
ID string `json:"id"`
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"function"`
} `json:"tool_calls"`
}
type chatReq struct {
Model string `json:"model"`
Messages []chatMsg `json:"messages"`
Tools []chatTool `json:"tools"`
Temperature *float64 `json:"temperature"`
TopP *float64 `json:"top_p"`
MaxTokens *int `json:"max_tokens"`
Stop []string `json:"stop"`
Stream bool `json:"stream"`
}
// chatToMessagesReq 将 OpenAI Chat 请求转为 Anthropic Messages 请求。
func chatToMessagesReq(body []byte) ([]byte, error) {
var req chatReq
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
out := map[string]any{
"model": req.Model,
"max_tokens": intOrNil(req.MaxTokens, 1024), // Anthropic 必填
}
if req.Stream {
out["stream"] = true
}
if req.Temperature != nil {
out["temperature"] = *req.Temperature
}
if req.TopP != nil {
out["top_p"] = *req.TopP
}
if len(req.Stop) > 0 {
out["stop_sequences"] = req.Stop
}
var system []string
msgs := make([]any, 0, len(req.Messages))
for _, m := range req.Messages {
if m.Role == "system" {
if s := str(m.Content); s != "" {
system = append(system, s)
}
continue
}
msgs = append(msgs, chatMsgToAnthropic(m))
}
if len(system) > 0 {
out["system"] = strings.Join(system, "\n")
}
out["messages"] = msgs
if len(req.Tools) > 0 {
tools := make([]any, 0, len(req.Tools))
for _, t := range req.Tools {
var params any
if len(t.Function.Parameters) > 0 && string(t.Function.Parameters) != "null" {
_ = json.Unmarshal(t.Function.Parameters, &params)
}
tools = append(tools, map[string]any{
"name": t.Function.Name,
"description": t.Function.Description,
"input_schema": params,
})
}
out["tools"] = tools
}
return json.Marshal(out)
}
// chatMsgToAnthropic 单条消息转 Anthropic 内容。
func chatMsgToAnthropic(m chatMsg) any {
switch m.Role {
case "assistant":
content := make([]any, 0, 2)
if s := str(m.Content); s != "" {
content = append(content, map[string]any{"type": "text", "text": s})
}
for _, tc := range m.ToolCalls {
var input any
if tc.Function.Arguments != "" {
_ = json.Unmarshal([]byte(tc.Function.Arguments), &input)
}
content = append(content, map[string]any{
"type": "tool_use",
"id": tc.ID,
"name": tc.Function.Name,
"input": input,
})
}
return map[string]any{"role": "assistant", "content": content}
case "tool":
return map[string]any{"role": "user", "content": []any{
map[string]any{"type": "tool_result", "tool_use_id": m.ToolCallID, "content": str(m.Content)},
}}
default: // user
var arr []map[string]any
if json.Unmarshal(m.Content, &arr) == nil && arr != nil {
blocks := make([]any, 0, len(arr))
for _, b := range arr {
switch b["type"] {
case "text", "input_text":
if t, _ := b["text"].(string); t != "" {
blocks = append(blocks, map[string]any{"type": "text", "text": t})
}
case "image_url":
var url string
if iu, ok := b["image_url"].(map[string]any); ok {
url, _ = iu["url"].(string)
} else if s, ok := b["image_url"].(string); ok {
url = s
}
if url != "" {
blocks = append(blocks, anthropicImageBlock(url))
}
}
}
if len(blocks) > 0 {
return map[string]any{"role": "user", "content": blocks}
}
}
return map[string]any{"role": "user", "content": str(m.Content)}
}
}
// ---------------------------------------------------------------------------
// 请求:Messages → Chat
type messagesReq struct {
Model string `json:"model"`
System json.RawMessage `json:"system"`
Messages []struct {
Role string `json:"role"`
Content json.RawMessage `json:"content"`
} `json:"messages"`
Tools []struct {
Name string `json:"name"`
Description string `json:"description"`
InputSchema json.RawMessage `json:"input_schema"`
} `json:"tools"`
Temperature *float64 `json:"temperature"`
TopP *float64 `json:"top_p"`
MaxTokens *int `json:"max_tokens"`
StopSequence []string `json:"stop_sequences"`
Stream bool `json:"stream"`
}
// messagesToChatReq 将 Anthropic Messages 请求转为 OpenAI Chat 请求。
func messagesToChatReq(body []byte) ([]byte, error) {
var req messagesReq
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
out := map[string]any{"model": req.Model}
if req.Stream {
out["stream"] = true
}
if req.Temperature != nil {
out["temperature"] = *req.Temperature
}
if req.TopP != nil {
out["top_p"] = *req.TopP
}
if req.MaxTokens != nil {
out["max_tokens"] = *req.MaxTokens
}
if len(req.StopSequence) > 0 {
out["stop"] = req.StopSequence
}
msgs := make([]any, 0, len(req.Messages)+1)
if s := str(req.System); s != "" {
msgs = append(msgs, map[string]any{"role": "system", "content": s})
}
for _, m := range req.Messages {
msgs = append(msgs, anthropicMsgToChat(m.Role, m.Content)...)
}
out["messages"] = msgs
if len(req.Tools) > 0 {
tools := make([]any, 0, len(req.Tools))
for _, t := range req.Tools {
tools = append(tools, map[string]any{
"type": "function",
"function": map[string]any{
"name": t.Name,
"description": t.Description,
"parameters": rawOrObject(t.InputSchema),
},
})
}
out["tools"] = tools
}
return json.Marshal(out)
}
// anthropicMsgToChat 将一条 Anthropic 消息拆成 0..N 条 Chat 消息。
func anthropicMsgToChat(role string, content json.RawMessage) []any {
// 块数组优先(tool_use / tool_result 需要分块解析)
var blocks []map[string]any
if json.Unmarshal(content, &blocks) == nil && blocks != nil {
var out []any
var toolMsgs []any // tool_result 单独收集,保证排在 assistant(tool_calls) 之后
var textParts []string
var contentBlocks []any // text / image_url 块,保留原始顺序
var toolCalls []any
for _, b := range blocks {
switch b["type"] {
case "text":
if t, _ := b["text"].(string); t != "" {
textParts = append(textParts, t)
contentBlocks = append(contentBlocks, map[string]any{"type": "text", "text": t})
}
case "image":
if cb := chatImageBlock(b); cb != nil {
contentBlocks = append(contentBlocks, cb)
}
case "tool_use":
id, _ := b["id"].(string)
name, _ := b["name"].(string)
args, _ := json.Marshal(b["input"])
toolCalls = append(toolCalls, map[string]any{
"id": id,
"type": "function",
"function": map[string]any{
"name": name,
"arguments": string(args),
},
})
case "tool_result":
callID, _ := b["tool_use_id"].(string)
res := strField(b["content"])
toolMsgs = append(toolMsgs, map[string]any{"role": "tool", "tool_call_id": callID, "content": res})
}
}
hasImage := false
for _, cb := range contentBlocks {
if m, _ := cb.(map[string]any); m["type"] == "image_url" {
hasImage = true
break
}
}
if hasImage || len(textParts) > 0 || len(toolCalls) > 0 {
msg := map[string]any{"role": role}
switch {
case hasImage:
msg["content"] = contentBlocks
case len(textParts) > 0:
msg["content"] = strings.Join(textParts, "")
}
if len(toolCalls) > 0 {
msg["tool_calls"] = toolCalls
}
out = append(out, msg)
}
out = append(out, toolMsgs...)
if len(out) > 0 {
return out
}
}
// 纯文本
if s := str(content); s != "" {
return []any{map[string]any{"role": role, "content": s}}
}
return nil
}
// ---------------------------------------------------------------------------
// 响应:Messages → Chat
type messagesResp struct {
ID string `json:"id"`
Model string `json:"model"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
ID string `json:"id"`
Name string `json:"name"`
Input json.RawMessage `json:"input"`
} `json:"content"`
StopReason string `json:"stop_reason"`
Usage struct {
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
} `json:"usage"`
}
// messagesToChatResp 将 Anthropic Messages 响应(非流式)转为 Chat 响应。
func messagesToChatResp(body []byte) ([]byte, error) {
var r messagesResp
if err := json.Unmarshal(body, &r); err != nil {
return nil, err
}
var text string
var toolCalls []any
for _, c := range r.Content {
switch c.Type {
case "text":
text += c.Text
case "tool_use":
args, _ := json.Marshal(c.Input)
toolCalls = append(toolCalls, map[string]any{
"id": c.ID,
"type": "function",
"function": map[string]any{
"name": c.Name,
"arguments": string(args),
},
})
}
}
msg := map[string]any{"role": "assistant", "content": text}
if len(toolCalls) > 0 {
msg["tool_calls"] = toolCalls
}
return json.Marshal(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(r.ID, "msg_"),
"object": "chat.completion",
"model": r.Model,
"created": 0,
"choices": []any{map[string]any{
"index": 0,
"message": msg,
"finish_reason": messagesStopToChat(r.StopReason),
}},
"usage": map[string]any{
"prompt_tokens": r.Usage.InputTokens,
"completion_tokens": r.Usage.OutputTokens,
"total_tokens": r.Usage.InputTokens + r.Usage.OutputTokens,
},
})
}
// ---------------------------------------------------------------------------
// 响应:Chat → Messages
type chatResp struct {
ID string `json:"id"`
Model string `json:"model"`
Choices []struct {
Message struct {
Role string `json:"role"`
Content string `json:"content"`
ToolCalls []struct {
ID string `json:"id"`
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"function"`
} `json:"tool_calls"`
} `json:"message"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
Usage struct {
PromptTokens int64 `json:"prompt_tokens"`
CompletionTokens int64 `json:"completion_tokens"`
} `json:"usage"`
}
// chatToMessagesResp 将 Chat 响应(非流式)转为 Messages 响应。
func chatToMessagesResp(body []byte) ([]byte, error) {
var r chatResp
if err := json.Unmarshal(body, &r); err != nil {
return nil, err
}
content := make([]any, 0, 2)
var finish = "end_turn"
if len(r.Choices) > 0 {
msg := r.Choices[0].Message
if msg.Content != "" {
content = append(content, map[string]any{"type": "text", "text": msg.Content})
}
for _, tc := range msg.ToolCalls {
var input any
_ = json.Unmarshal([]byte(tc.Function.Arguments), &input)
content = append(content, map[string]any{
"type": "tool_use",
"id": tc.ID,
"name": tc.Function.Name,
"input": input,
})
}
finish = chatStopToMessages(r.Choices[0].FinishReason)
}
return json.Marshal(map[string]any{
"id": "msg_" + strings.TrimPrefix(r.ID, "chatcmpl-"),
"type": "message",
"role": "assistant",
"model": r.Model,
"content": content,
"stop_reason": finish,
"usage": map[string]any{
"input_tokens": r.Usage.PromptTokens,
"output_tokens": r.Usage.CompletionTokens,
},
})
}
// ---------------------------------------------------------------------------
// 辅助
// splitDataURL 解析 data:media_type;base64,data 形式的 URL;非该形式返回 ok=false。
func splitDataURL(url string) (media, data string, ok bool) {
if !strings.HasPrefix(url, "data:") {
return "", "", false
}
i := strings.Index(url, ";base64,")
if i < 0 {
return "", "", false
}
return url[len("data:"):i], url[i+len(";base64,"):], true
}
// chatImageBlock 把 Anthropic image 块转 OpenAI image_url 块。
// 仅支持 base64 与 url source;其他类型(如 Files API 的 file_id)不支持,跳过。
func chatImageBlock(b map[string]any) any {
src, ok := b["source"].(map[string]any)
if !ok {
return nil
}
switch src["type"] {
case "base64":
media, _ := src["media_type"].(string)
data, _ := src["data"].(string)
if data == "" {
return nil
}
if media == "" {
media = "image/png"
}
return map[string]any{"type": "image_url", "image_url": map[string]any{"url": "data:" + media + ";base64," + data}}
case "url":
url, _ := src["url"].(string)
if url == "" {
return nil
}
return map[string]any{"type": "image_url", "image_url": map[string]any{"url": url}}
}
return nil
}
// anthropicImageBlock 把 OpenAI image_url 的 url 转 Anthropic image 块。
// data URL → base64 source;http(s) URL → url source。
func anthropicImageBlock(url string) any {
if media, data, ok := splitDataURL(url); ok {
if media == "" {
media = "image/png"
}
return map[string]any{"type": "image", "source": map[string]any{"type": "base64", "media_type": media, "data": data}}
}
return map[string]any{"type": "image", "source": map[string]any{"type": "url", "url": url}}
}
func messagesStopToChat(s string) string {
switch s {
case "tool_use":
return "tool_calls"
case "max_tokens":
return "length"
default:
return "stop"
}
}
func chatStopToMessages(s string) string {
switch s {
case "tool_calls":
return "tool_use"
case "length":
return "max_tokens"
default:
return "end_turn"
}
}
@@ -0,0 +1,376 @@
package convert
import (
"encoding/json"
"strings"
)
// ---------------------------------------------------------------------------
// 请求:Responses → Chat
// responsesToChatReq 将 OpenAI Responses 请求转为 Chat 请求。
func responsesToChatReq(body []byte) ([]byte, error) {
var m map[string]json.RawMessage
if err := json.Unmarshal(body, &m); err != nil {
return nil, err
}
out := map[string]any{"model": str(rawJSON(m, "model"))}
if v, ok := m["stream"]; ok && string(v) == "true" {
out["stream"] = true
}
if v, ok := m["temperature"]; ok {
out["temperature"] = v
}
if v, ok := m["top_p"]; ok {
out["top_p"] = v
}
if v, ok := m["max_output_tokens"]; ok {
out["max_tokens"] = v
}
var msgs []any
if ins := str(rawJSON(m, "instructions")); ins != "" {
msgs = append(msgs, map[string]any{"role": "system", "content": ins})
}
msgs = append(msgs, responsesInputToChat(rawJSON(m, "input"))...)
out["messages"] = msgs
if raw := rawJSON(m, "tools"); raw != nil {
var tools []map[string]any
if json.Unmarshal(raw, &tools) == nil {
chatTools := make([]any, 0, len(tools))
for _, t := range tools {
chatTools = append(chatTools, map[string]any{
"type": "function",
"function": map[string]any{
"name": t["name"],
"description": t["description"],
"parameters": t["parameters"],
},
})
}
out["tools"] = chatTools
}
}
return json.Marshal(out)
}
// responsesInputToChat 把 Responses input 转成 Chat messages。
// input 支持字符串或条目数组(message / function_call / function_call_output)。
func responsesInputToChat(raw json.RawMessage) []any {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
// 字符串输入
if s := str(raw); s != "" {
return []any{map[string]any{"role": "user", "content": s}}
}
var items []map[string]any
if err := json.Unmarshal(raw, &items); err != nil || items == nil {
return nil
}
var out []any
for _, item := range items {
switch item["type"] {
case "function_call":
out = append(out, map[string]any{
"role": "assistant",
"content": "",
"tool_calls": []any{map[string]any{
"id": strField(item["call_id"]),
"type": "function",
"function": map[string]any{
"name": strField(item["name"]),
"arguments": strField(item["arguments"]),
},
}},
})
case "function_call_output":
out = append(out, map[string]any{
"role": "tool",
"tool_call_id": strField(item["call_id"]),
"content": strField(item["output"]),
})
default: // message 条目
role, _ := item["role"].(string)
if role == "" {
role = "user"
}
if content, ok := item["content"].(string); ok {
out = append(out, map[string]any{"role": role, "content": content})
} else if blocks, ok := item["content"].([]any); ok {
var text []string
var contentBlocks []any
for _, b := range blocks {
bm, ok := b.(map[string]any)
if !ok {
continue
}
switch bm["type"] {
case "input_text", "text":
if t, _ := bm["text"].(string); t != "" {
text = append(text, t)
contentBlocks = append(contentBlocks, map[string]any{"type": "text", "text": t})
}
case "input_image":
var url string
if s, ok := bm["image_url"].(string); ok {
url = s
} else if m, ok := bm["image_url"].(map[string]any); ok {
url, _ = m["url"].(string)
}
if url != "" {
contentBlocks = append(contentBlocks, map[string]any{"type": "image_url", "image_url": map[string]any{"url": url}})
}
}
}
hasImage := false
for _, cb := range contentBlocks {
if m, _ := cb.(map[string]any); m["type"] == "image_url" {
hasImage = true
break
}
}
if hasImage {
out = append(out, map[string]any{"role": role, "content": contentBlocks})
} else {
out = append(out, map[string]any{"role": role, "content": strings.Join(text, "")})
}
}
}
}
return out
}
// chatContentToResponsesBlocks 把 Chat 用户消息 content 转 Responses input 块数组(input_text / input_image)。
func chatContentToResponsesBlocks(content json.RawMessage) []any {
// 纯字符串 → 单个 input_text
var s string
if json.Unmarshal(content, &s) == nil && s != "" {
return []any{map[string]any{"type": "input_text", "text": s}}
}
// 数组 → 按块转换(text / image_url)
var arr []map[string]any
if json.Unmarshal(content, &arr) == nil && arr != nil {
var out []any
for _, b := range arr {
switch b["type"] {
case "text", "input_text":
if t, _ := b["text"].(string); t != "" {
out = append(out, map[string]any{"type": "input_text", "text": t})
}
case "image_url":
var url string
if iu, ok := b["image_url"].(map[string]any); ok {
url, _ = iu["url"].(string)
} else if s, ok := b["image_url"].(string); ok {
url = s
}
if url != "" {
out = append(out, map[string]any{"type": "input_image", "image_url": url})
}
}
}
return out
}
return nil
}
// ---------------------------------------------------------------------------
// 请求:Chat → Responses
// chatToResponsesReq 将 Chat 请求转为 Responses 请求。
func chatToResponsesReq(body []byte) ([]byte, error) {
var req chatReq
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
out := map[string]any{"model": req.Model}
if req.Stream {
out["stream"] = true
}
if req.Temperature != nil {
out["temperature"] = *req.Temperature
}
if req.TopP != nil {
out["top_p"] = *req.TopP
}
if req.MaxTokens != nil {
out["max_output_tokens"] = *req.MaxTokens
}
var system []string
var input []any
for _, m := range req.Messages {
if m.Role == "system" {
if s := str(m.Content); s != "" {
system = append(system, s)
}
continue
}
switch m.Role {
case "tool":
input = append(input, map[string]any{
"type": "function_call_output",
"call_id": m.ToolCallID,
"output": str(m.Content),
})
case "assistant":
if len(m.ToolCalls) > 0 {
for _, tc := range m.ToolCalls {
input = append(input, map[string]any{
"type": "function_call",
"call_id": tc.ID,
"name": tc.Function.Name,
"arguments": tc.Function.Arguments,
})
}
} else if s := str(m.Content); s != "" {
input = append(input, map[string]any{"type": "message", "role": "assistant", "content": []any{
map[string]any{"type": "input_text", "text": s},
}})
}
default:
if blocks := chatContentToResponsesBlocks(m.Content); len(blocks) > 0 {
input = append(input, map[string]any{"type": "message", "role": "user", "content": blocks})
}
}
}
if len(system) > 0 {
out["instructions"] = strings.Join(system, "\n")
}
// input 必须是数组:部分上游只接受数组,单对象会被拒(400 Mismatch type)。
out["input"] = input
if len(req.Tools) > 0 {
tools := make([]any, 0, len(req.Tools))
for _, t := range req.Tools {
tools = append(tools, map[string]any{
"type": "function",
"name": t.Function.Name,
"description": t.Function.Description,
"parameters": rawOrObject(t.Function.Parameters),
})
}
out["tools"] = tools
}
return json.Marshal(out)
}
// ---------------------------------------------------------------------------
// 响应:Responses → Chat
// responsesToChatResp 将 Responses 响应(非流式)转为 Chat 响应。
func responsesToChatResp(body []byte) ([]byte, error) {
var m map[string]json.RawMessage
if err := json.Unmarshal(body, &m); err != nil {
return nil, err
}
var text string
var toolCalls []any
if raw := rawJSON(m, "output"); raw != nil {
var outputs []map[string]any
if json.Unmarshal(raw, &outputs) == nil {
for _, o := range outputs {
switch o["type"] {
case "message":
if content, ok := o["content"].([]any); ok {
for _, c := range content {
if cm, ok := c.(map[string]any); ok {
if t, _ := cm["text"].(string); t != "" {
text += t
}
}
}
}
case "function_call":
toolCalls = append(toolCalls, map[string]any{
"id": strField(o["call_id"]),
"type": "function",
"function": map[string]any{
"name": strField(o["name"]),
"arguments": strField(o["arguments"]),
},
})
}
}
}
}
msg := map[string]any{"role": "assistant", "content": text}
if len(toolCalls) > 0 {
msg["tool_calls"] = toolCalls
}
finish := "stop"
switch {
case string(rawJSON(m, "status")) == `"incomplete"`:
finish = "length" // 截断优先,客户端可据此区分
case len(toolCalls) > 0:
finish = "tool_calls"
}
var prompt, completion int64
if u := rawJSON(m, "usage"); u != nil {
var us struct {
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
}
_ = json.Unmarshal(u, &us)
prompt, completion = us.InputTokens, us.OutputTokens
}
return json.Marshal(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(str(rawJSON(m, "id")), "resp_"),
"object": "chat.completion",
"model": str(rawJSON(m, "model")),
"choices": []any{map[string]any{"index": 0, "message": msg, "finish_reason": finish}},
"usage": map[string]any{
"prompt_tokens": prompt,
"completion_tokens": completion,
"total_tokens": prompt + completion,
},
})
}
// ---------------------------------------------------------------------------
// 响应:Chat → Responses
// chatToResponsesResp 将 Chat 响应(非流式)转为 Responses 响应。
func chatToResponsesResp(body []byte) ([]byte, error) {
var r chatResp
if err := json.Unmarshal(body, &r); err != nil {
return nil, err
}
output := make([]any, 0, 2)
var finish = "completed"
if len(r.Choices) > 0 {
msg := r.Choices[0].Message
if msg.Content != "" {
output = append(output, map[string]any{
"type": "message",
"role": "assistant",
"content": []any{map[string]any{"type": "output_text", "text": msg.Content}},
})
}
for _, tc := range msg.ToolCalls {
output = append(output, map[string]any{
"type": "function_call",
"call_id": tc.ID,
"name": tc.Function.Name,
"arguments": tc.Function.Arguments,
})
}
if r.Choices[0].FinishReason == "length" {
finish = "incomplete"
}
}
return json.Marshal(map[string]any{
"id": "resp_" + strings.TrimPrefix(r.ID, "chatcmpl-"),
"object": "response",
"model": r.Model,
"status": finish,
"output": output,
"usage": map[string]any{
"input_tokens": r.Usage.PromptTokens,
"output_tokens": r.Usage.CompletionTokens,
"total_tokens": r.Usage.PromptTokens + r.Usage.CompletionTokens,
},
})
}
+16 -1
View File
@@ -1,9 +1,11 @@
package convert package convert
import "encoding/json"
// ResponsesRequest represents an OpenAI Responses API request // ResponsesRequest represents an OpenAI Responses API request
type ResponsesRequest struct { type ResponsesRequest struct {
Model string `json:"model"` Model string `json:"model"`
Input []InputItem `json:"input"` Input json.RawMessage `json:"input,omitempty"`
Instructions string `json:"instructions,omitempty"` Instructions string `json:"instructions,omitempty"`
MaxOutputTokens *int `json:"max_output_tokens,omitempty"` MaxOutputTokens *int `json:"max_output_tokens,omitempty"`
Tools []Tool `json:"tools,omitempty"` Tools []Tool `json:"tools,omitempty"`
@@ -20,6 +22,19 @@ type InputItem struct {
Content interface{} `json:"content,omitempty"` Content interface{} `json:"content,omitempty"`
} }
// marshalInputItems 把 input 条目序列化为 Responses input 的 json.RawMessage 形态。
// Input 字段用 RawMessage 以兼容字符串与条目数组两种客户端写法。
func marshalInputItems(items []InputItem) json.RawMessage {
if len(items) == 0 {
return nil
}
b, err := json.Marshal(items)
if err != nil {
return nil
}
return b
}
// ResponsesResponse represents an OpenAI Responses API response // ResponsesResponse represents an OpenAI Responses API response
type ResponsesResponse struct { type ResponsesResponse struct {
ID string `json:"id"` ID string `json:"id"`
@@ -0,0 +1,642 @@
package convert
import (
"encoding/json"
"strings"
)
// sseState 记录上一行 event 名。
type sseState struct {
event string
}
// parseLine 解析一行 SSE;返回是否 data 行及其内容、是否 [DONE]。
// data: 后可跟空格(标准)或紧贴 JSON(部分上游会省略空格)。
func (s *sseState) parseLine(line []byte) (isData bool, data string, done bool) {
strLine := strings.TrimRight(string(line), "\r\n")
switch {
case strings.HasPrefix(strLine, "event: "):
s.event = strings.TrimSpace(strings.TrimPrefix(strLine, "event: "))
return false, "", false
case strLine == "data: [DONE]" || strLine == "data:[DONE]":
return true, "[DONE]", true
case strings.HasPrefix(strLine, "data:"):
return true, strings.TrimLeft(strings.TrimPrefix(strLine, "data:"), " "), false
default:
return false, "", false
}
}
func eventData(line string) map[string]any {
var m map[string]any
_ = json.Unmarshal([]byte(line), &m)
return m
}
func dataLine(obj any) []byte {
b, _ := json.Marshal(obj)
return append(append([]byte("data: "), b...), '\n', '\n')
}
func eventLine(name string, obj any) []byte {
b, _ := json.Marshal(obj)
out := append([]byte("event: "+name+"\ndata: "), b...)
return append(out, '\n', '\n')
}
// joinLines 拼接多条 SSE 行。
func joinLines(lines [][]byte) []byte {
var s []string
for _, l := range lines {
s = append(s, string(l))
}
return []byte(strings.Join(s, ""))
}
// ---------------------------------------------------------------------------
// Messages → Chat
type messagesToChat struct {
sseState
id, model string
toolIdx map[int]int // messages content block index → chat tool_calls index(顺序编号,避开文本块)
nextTool int
}
func newMessagesToChat() *messagesToChat { return &messagesToChat{toolIdx: map[int]int{}} }
func (t *messagesToChat) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData {
return nil
}
if done {
return []byte("data: [DONE]\n\n")
}
m := eventData(data)
evt, _ := m["type"].(string)
switch evt {
case "message_start":
msg, _ := m["message"].(map[string]any)
t.id, _ = msg["id"].(string)
t.model, _ = msg["model"].(string)
return dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{"role": "assistant"}, "finish_reason": nil}},
})
case "content_block_start":
cb, _ := m["content_block"].(map[string]any)
if cb == nil || cb["type"] != "tool_use" {
return nil
}
blockIdx, _ := m["index"].(float64)
tool := t.nextTool
t.nextTool++
t.toolIdx[int(blockIdx)] = tool
toolID, _ := cb["id"].(string)
name, _ := cb["name"].(string)
return dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{
"tool_calls": []any{map[string]any{"index": tool, "id": toolID, "type": "function", "function": map[string]any{"name": name, "arguments": ""}}},
}, "finish_reason": nil}},
})
case "content_block_delta":
delta, _ := m["delta"].(map[string]any)
deltaType, _ := delta["type"].(string)
if deltaType == "input_json_delta" {
blockIdx, _ := m["index"].(float64)
tool, ok := t.toolIdx[int(blockIdx)]
if !ok {
return nil
}
partial, _ := delta["partial_json"].(string)
if partial == "" {
return nil
}
return dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{
"tool_calls": []any{map[string]any{"index": tool, "function": map[string]any{"arguments": partial}}},
}, "finish_reason": nil}},
})
}
text, _ := delta["text"].(string)
if text == "" {
return nil
}
return dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{"content": text}, "finish_reason": nil}},
})
case "message_delta":
delta, _ := m["delta"].(map[string]any)
stop, _ := delta["stop_reason"].(string)
var out [][]byte
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{}, "finish_reason": messagesStopToChat(stop)}},
}))
if u, ok := m["usage"]; ok {
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "msg_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{}, "usage": u,
}))
}
return joinLines(out)
case "message_stop":
return []byte("data: [DONE]\n\n")
}
return nil
}
// ---------------------------------------------------------------------------
// Chat → Messages
type chatToMessages struct {
sseState
started bool // message_start 已发出
nextIndex int // 下一个 content block index(顺序分配)
textIndex int // 文本块 index;-1 = 未开始
toolIdx map[int]int // chat delta.tool_calls[].index → messages block index
openBlocks []int // 已开始未停止的 block index,按开始顺序
model string
stopReason string
usage any
}
func newChatToMessages() *chatToMessages {
return &chatToMessages{textIndex: -1, toolIdx: map[int]int{}}
}
func (t *chatToMessages) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData {
return nil
}
if done {
// 汇聚最终:先对每个已开始未停止的块发 content_block_stop,再 message_delta + message_stop
var out [][]byte
for _, idx := range t.openBlocks {
out = append(out, eventLine("content_block_stop", map[string]any{"type": "content_block_stop", "index": idx}))
}
md := map[string]any{"type": "message_delta", "delta": map[string]any{
"stop_reason": stopReasonOrEnd(t.stopReason), "stop_sequence": nil,
}}
if t.usage != nil {
md["usage"] = t.usage
}
out = append(out, eventLine("message_delta", md))
out = append(out, eventLine("message_stop", map[string]any{"type": "message_stop"}))
return joinLines(out)
}
m := eventData(data)
// chat 块:delta / finish_reason 在 choices[0] 内
delta := map[string]any{}
if choices, ok := m["choices"].([]any); ok && len(choices) > 0 {
if c0, ok := choices[0].(map[string]any); ok {
if d, ok := c0["delta"].(map[string]any); ok {
delta = d
}
if fr, _ := c0["finish_reason"].(string); fr != "" {
t.stopReason = fr
}
}
}
if t.model == "" {
t.model, _ = m["model"].(string)
}
id, _ := m["id"].(string)
var out [][]byte
// message_start 只在实际有内容(文本或工具)时发出,避免 reasoning_content 块
//(带 role 无 content)提前开出一个空文本块。
ensureStarted := func() {
if t.started {
return
}
t.started = true
out = append(out, eventLine("message_start", map[string]any{
"type": "message_start",
"message": map[string]any{
"id": "msg_" + strings.TrimPrefix(id, "chatcmpl-"), "type": "message", "role": "assistant",
"model": t.model, "content": []any{}, "usage": map[string]any{"input_tokens": 0, "output_tokens": 0},
},
}))
}
// 文本:delta.content(string;兼容 {type:text,text} 数组)
if content := deltaText(delta); content != "" {
if t.textIndex < 0 {
t.textIndex = t.nextIndex
t.nextIndex++
ensureStarted()
out = append(out, eventLine("content_block_start", map[string]any{
"type": "content_block_start", "index": t.textIndex, "content_block": map[string]any{"type": "text", "text": ""},
}))
t.openBlocks = append(t.openBlocks, t.textIndex)
}
out = append(out, eventLine("content_block_delta", map[string]any{
"type": "content_block_delta", "index": t.textIndex, "delta": map[string]any{"type": "text_delta", "text": content},
}))
}
// 工具调用:delta.tool_calls(并行调用各 index 独立成块;arguments 支持整段/分段两种流式)
if tcs, ok := delta["tool_calls"].([]any); ok {
for _, tc := range tcs {
call, ok := tc.(map[string]any)
if !ok {
continue
}
idx, _ := call["index"].(float64)
tcIdx := int(idx)
fn, _ := call["function"].(map[string]any)
name, _ := fn["name"].(string)
args, _ := fn["arguments"].(string)
blockIdx, seen := t.toolIdx[tcIdx]
if !seen {
blockIdx = t.nextIndex
t.nextIndex++
t.toolIdx[tcIdx] = blockIdx
toolID, _ := call["id"].(string)
ensureStarted()
out = append(out, eventLine("content_block_start", map[string]any{
"type": "content_block_start", "index": blockIdx, "content_block": map[string]any{
"type": "tool_use", "id": toolID, "name": name, "input": map[string]any{},
},
}))
t.openBlocks = append(t.openBlocks, blockIdx)
}
if args != "" {
out = append(out, eventLine("content_block_delta", map[string]any{
"type": "content_block_delta", "index": blockIdx, "delta": map[string]any{"type": "input_json_delta", "partial_json": args},
}))
}
}
}
if u, ok := m["usage"]; ok {
t.usage = u
}
return joinLines(out)
}
// deltaText 取 chat delta.content 文本(string 或 [{type:text,text}] 数组拼接)。
func deltaText(delta map[string]any) string {
if s, ok := delta["content"].(string); ok {
return s
}
if arr, ok := delta["content"].([]any); ok {
var parts []string
for _, b := range arr {
if bm, ok := b.(map[string]any); ok {
if t, _ := bm["text"].(string); t != "" {
parts = append(parts, t)
}
}
}
return strings.Join(parts, "")
}
return ""
}
func stopReasonOrEnd(s string) string {
if s == "" {
return "end_turn"
}
return chatStopToMessages(s)
}
// ---------------------------------------------------------------------------
// Responses → Messages
type responsesToMessages struct {
sseState
started bool
model string
usage any
nextIndex int // 下一个 content block index(顺序分配)
textIndex int // 文本块 index;-1 = 未开始
toolIdx map[string]int // function_call item_id → messages block index
openBlocks []int // 已开始未停止的 block index,按开始顺序
anyTool bool
}
func newResponsesToMessages() *responsesToMessages {
return &responsesToMessages{textIndex: -1, toolIdx: map[string]int{}}
}
func (t *responsesToMessages) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData || done {
return nil
}
m := eventData(data)
evt, _ := m["type"].(string)
if resp, ok := m["response"].(map[string]any); ok {
if t.model == "" {
t.model, _ = resp["model"].(string)
}
if u, ok := resp["usage"]; ok {
t.usage = u
}
}
var out [][]byte
// message_start 只在 response.created 时发出;文本/工具块在对应事件到达时再开,
// 避免纯函数调用响应提前开出一个空文本块。
ensureStarted := func() {
if t.started {
return
}
t.started = true
rid := ""
if resp, ok := m["response"].(map[string]any); ok {
rid, _ = resp["id"].(string)
}
out = append(out, eventLine("message_start", map[string]any{
"type": "message_start",
"message": map[string]any{
"id": "msg_" + strings.TrimPrefix(rid, "resp_"), "type": "message", "role": "assistant",
"model": t.model, "content": []any{},
},
}))
}
switch evt {
case "response.created":
ensureStarted()
case "response.output_text.delta":
delta, _ := m["delta"].(string)
if delta == "" {
return nil
}
if t.textIndex < 0 {
t.textIndex = t.nextIndex
t.nextIndex++
ensureStarted()
out = append(out, eventLine("content_block_start", map[string]any{
"type": "content_block_start", "index": t.textIndex, "content_block": map[string]any{"type": "text", "text": ""},
}))
t.openBlocks = append(t.openBlocks, t.textIndex)
}
out = append(out, eventLine("content_block_delta", map[string]any{
"type": "content_block_delta", "index": t.textIndex, "delta": map[string]any{"type": "text_delta", "text": delta},
}))
case "response.output_item.added":
item, _ := m["item"].(map[string]any)
if item == nil || item["type"] != "function_call" {
return nil
}
blockIdx := t.nextIndex
t.nextIndex++
t.anyTool = true
itemID, _ := item["id"].(string)
t.toolIdx[itemID] = blockIdx
toolUseID, _ := item["call_id"].(string)
if toolUseID == "" {
toolUseID = itemID
}
name, _ := item["name"].(string)
ensureStarted()
out = append(out, eventLine("content_block_start", map[string]any{
"type": "content_block_start", "index": blockIdx, "content_block": map[string]any{
"type": "tool_use", "id": toolUseID, "name": name, "input": map[string]any{},
},
}))
t.openBlocks = append(t.openBlocks, blockIdx)
case "response.function_call_arguments.delta":
itemID, _ := m["item_id"].(string)
blockIdx, ok := t.toolIdx[itemID]
if !ok {
return nil
}
delta, _ := m["delta"].(string)
if delta == "" {
return nil
}
out = append(out, eventLine("content_block_delta", map[string]any{
"type": "content_block_delta", "index": blockIdx, "delta": map[string]any{"type": "input_json_delta", "partial_json": delta},
}))
case "response.completed":
for _, idx := range t.openBlocks {
out = append(out, eventLine("content_block_stop", map[string]any{"type": "content_block_stop", "index": idx}))
}
stop := "end_turn"
if t.anyTool {
stop = "tool_use"
}
md := map[string]any{"type": "message_delta", "delta": map[string]any{"stop_reason": stop, "stop_sequence": nil}}
if t.usage != nil {
md["usage"] = t.usage
}
out = append(out, eventLine("message_delta", md))
out = append(out, eventLine("message_stop", map[string]any{"type": "message_stop"}))
}
return joinLines(out)
}
// ---------------------------------------------------------------------------
// Messages → Responses
type messagesToResponses struct {
sseState
model string
usage any
done bool
}
func newMessagesToResponses() *messagesToResponses { return &messagesToResponses{} }
func (t *messagesToResponses) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData || done {
return nil
}
m := eventData(data)
evt, _ := m["type"].(string)
if msg, ok := m["message"].(map[string]any); ok {
if t.model == "" {
t.model, _ = msg["model"].(string)
}
if u, ok := msg["usage"]; ok {
t.usage = u
}
}
if u, ok := m["usage"]; ok {
t.usage = u
}
var out [][]byte
switch evt {
case "message_start":
id, _ := m["message"].(map[string]any)
rid := ""
if id != nil {
rid, _ = id["id"].(string)
}
out = append(out, eventLine("response.created", map[string]any{
"type": "response.created",
"response": map[string]any{
"id": "resp_" + strings.TrimPrefix(rid, "msg_"), "object": "response", "model": t.model, "status": "in_progress",
},
}))
case "content_block_delta":
delta, _ := m["delta"].(map[string]any)
text, _ := delta["text"].(string)
if text != "" {
out = append(out, eventLine("response.output_text.delta", map[string]any{
"type": "response.output_text.delta", "delta": text, "item_id": "msg_1", "output_index": 0, "content_index": 0,
}))
}
case "message_stop":
if !t.done {
t.done = true
out = append(out, eventLine("response.completed", map[string]any{
"type": "response.completed",
"response": map[string]any{
"id": "resp_stream", "object": "response", "model": t.model, "status": "completed", "usage": t.usage,
},
}))
}
}
return joinLines(out)
}
// ---------------------------------------------------------------------------
// Responses → Chat
type responsesToChat struct {
sseState
id, model string
}
func newResponsesToChat() *responsesToChat { return &responsesToChat{} }
func (t *responsesToChat) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData {
return nil
}
if done {
return nil
}
m := eventData(data)
evt, _ := m["type"].(string)
if resp, ok := m["response"].(map[string]any); ok {
if t.model == "" {
t.model, _ = resp["model"].(string)
}
if t.id == "" {
t.id, _ = resp["id"].(string)
}
}
var out [][]byte
switch evt {
case "response.created":
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "resp_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{"role": "assistant"}, "finish_reason": nil}},
}))
case "response.output_text.delta":
delta, _ := m["delta"].(string)
if delta != "" {
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "resp_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{"content": delta}, "finish_reason": nil}},
}))
}
case "response.completed":
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "resp_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{map[string]any{"index": 0, "delta": map[string]any{}, "finish_reason": "stop"}},
}))
if u, ok := m["response"].(map[string]any); ok {
if usage, ok := u["usage"]; ok {
out = append(out, dataLine(map[string]any{
"id": "chatcmpl-" + strings.TrimPrefix(t.id, "resp_"), "object": "chat.completion.chunk", "model": t.model,
"choices": []any{}, "usage": usage,
}))
}
}
out = append(out, []byte("data: [DONE]\n\n"))
}
return joinLines(out)
}
// ---------------------------------------------------------------------------
// Chat → Responses
type chatToResponses struct {
sseState
model string
usage any
finishSeen bool
done bool
createdSent bool
}
func newChatToResponses() *chatToResponses { return &chatToResponses{} }
func (t *chatToResponses) line(line []byte) []byte {
isData, data, done := t.parseLine(line)
if !isData {
return nil
}
if done {
// 流结束兜底:finish 后 usage 未随块到达时在此补发 completed
if !t.done {
t.done = true
return eventLine("response.completed", map[string]any{
"type": "response.completed",
"response": map[string]any{
"id": "resp_stream", "object": "response", "model": t.model, "status": "completed", "usage": t.usage,
},
})
}
return nil
}
m := eventData(data)
if t.model == "" {
t.model, _ = m["model"].(string)
}
if u, ok := m["usage"]; ok {
t.usage = u
}
delta := map[string]any{}
var finish string
if choices, ok := m["choices"].([]any); ok && len(choices) > 0 {
if c0, ok := choices[0].(map[string]any); ok {
if d, ok := c0["delta"].(map[string]any); ok {
delta = d
}
finish, _ = c0["finish_reason"].(string)
}
}
if finish != "" {
t.finishSeen = true
}
var out [][]byte
// 只发一次 response.created:部分上游(如 OpenRouter 的 reasoning 模型)会在
// 每个 chunk 的 delta 里都带 role:"assistant",不加守卫会刷出数十条 created。
if !t.createdSent && delta["role"] == "assistant" {
t.createdSent = true
out = append(out, eventLine("response.created", map[string]any{
"type": "response.created",
"response": map[string]any{"id": "resp_stream", "object": "response", "model": t.model, "status": "in_progress"},
}))
}
if content, _ := delta["content"].(string); content != "" {
out = append(out, eventLine("response.output_text.delta", map[string]any{
"type": "response.output_text.delta", "delta": content, "item_id": "msg_1", "output_index": 0, "content_index": 0,
}))
}
// 上游 usage 块(choices 为空)通常晚于 finish_reason:此时再发 completed,携带 usage
if _, hasUsage := m["usage"]; hasUsage && t.finishSeen && !t.done {
t.done = true
out = append(out, eventLine("response.completed", map[string]any{
"type": "response.completed",
"response": map[string]any{
"id": "resp_stream", "object": "response", "model": t.model, "status": "completed", "usage": t.usage,
},
}))
}
return joinLines(out)
}
+131
View File
@@ -0,0 +1,131 @@
package convert
import (
"encoding/json"
)
// TokenUsage 从上游响应提取的 token 用量。
// 三种协议的字段名不同,此处统一为:input / output / cache_read / cache_creation,
// 供用量记录与计费使用。
type TokenUsage struct {
InputTokens int
OutputTokens int
CacheReadTokens int
CacheCreationTokens int
}
// has 判断是否真的拿到了非零用量(过滤掉没有 usage 字段的响应)。
func (u *TokenUsage) has() bool {
return u.InputTokens > 0 || u.OutputTokens > 0 ||
u.CacheReadTokens > 0 || u.CacheCreationTokens > 0
}
// mergeJSON 把一张 usage 对象并入累计值。proto 决定字段名(chat/responses 与 messages 不同)。
func (u *TokenUsage) mergeJSON(raw map[string]any, proto string) {
switch proto {
case ProtoChat, ProtoResponses:
in, _ := raw["prompt_tokens"].(float64)
out, _ := raw["completion_tokens"].(float64)
if in == 0 && out == 0 {
in, _ = raw["input_tokens"].(float64)
out, _ = raw["output_tokens"].(float64)
}
u.InputTokens += int(in)
u.OutputTokens += int(out)
if d, ok := raw["prompt_tokens_details"].(map[string]any); ok {
if c, _ := d["cached_tokens"].(float64); c > 0 {
u.CacheReadTokens += int(c)
}
}
if d, ok := raw["input_tokens_details"].(map[string]any); ok {
if c, _ := d["cached_tokens"].(float64); c > 0 {
u.CacheReadTokens += int(c)
}
}
case ProtoMessages:
in, _ := raw["input_tokens"].(float64)
out, _ := raw["output_tokens"].(float64)
u.InputTokens += int(in)
u.OutputTokens += int(out)
if c, _ := raw["cache_read_input_tokens"].(float64); c > 0 {
u.CacheReadTokens += int(c)
}
if c, _ := raw["cache_creation_input_tokens"].(float64); c > 0 {
u.CacheCreationTokens += int(c)
}
}
}
// ExtractUsageJSON 从完整非流式响应体中提取用量。proto 为上游协议。
// 返回 (用量, 是否有效)。
func ExtractUsageJSON(body []byte, proto string) (TokenUsage, bool) {
var top map[string]any
if err := json.Unmarshal(body, &top); err != nil {
return TokenUsage{}, false
}
var u TokenUsage
if usage, ok := top["usage"].(map[string]any); ok {
u.mergeJSON(usage, proto)
}
return u, u.has()
}
// StreamUsageAccum 流式用量累计器。逐行喂入上游 SSE 的 data 载荷,
// 按协议分别取各事件里的 usage 字段(各事件只会携带一部分字段,取最大值合并)。
type StreamUsageAccum struct {
u TokenUsage
}
// NewStreamUsageAccum 创建一个流式用量累计器。
func NewStreamUsageAccum() *StreamUsageAccum {
return &StreamUsageAccum{}
}
// Feed 喂入一行 SSE data 载荷(不含 "data:" 前缀与换行)。
func (a *StreamUsageAccum) Feed(payload []byte, proto string) {
var top map[string]any
if json.Unmarshal(payload, &top) != nil {
return
}
var t TokenUsage
switch proto {
case ProtoChat:
if usage, ok := top["usage"].(map[string]any); ok {
t.mergeJSON(usage, proto)
}
case ProtoResponses:
// response.completed 事件把用量放在 response.usage 下。
if resp, ok := top["response"].(map[string]any); ok {
if usage, ok := resp["usage"].(map[string]any); ok {
t.mergeJSON(usage, proto)
}
}
case ProtoMessages:
// message_start: {message: {usage: {input_tokens, cache_*}}}
// message_delta: {usage: {output_tokens}}
if msg, ok := top["message"].(map[string]any); ok {
if usage, ok := msg["usage"].(map[string]any); ok {
t.mergeJSON(usage, proto)
}
}
if usage, ok := top["usage"].(map[string]any); ok {
var t2 TokenUsage
t2.mergeJSON(usage, proto)
t.InputTokens = max(t.InputTokens, t2.InputTokens)
t.OutputTokens = max(t.OutputTokens, t2.OutputTokens)
t.CacheReadTokens = max(t.CacheReadTokens, t2.CacheReadTokens)
t.CacheCreationTokens = max(t.CacheCreationTokens, t2.CacheCreationTokens)
}
default:
return
}
a.u.InputTokens = max(a.u.InputTokens, t.InputTokens)
a.u.OutputTokens = max(a.u.OutputTokens, t.OutputTokens)
a.u.CacheReadTokens = max(a.u.CacheReadTokens, t.CacheReadTokens)
a.u.CacheCreationTokens = max(a.u.CacheCreationTokens, t.CacheCreationTokens)
}
// Usage 返回当前累计用量。
func (a *StreamUsageAccum) Usage() TokenUsage {
return a.u
}
@@ -0,0 +1,170 @@
package convert
import (
"encoding/json"
"testing"
)
// ---- ExtractUsageJSON: 非流式各协议 ----
func TestExtractUsageJSONChat(t *testing.T) {
body := []byte(`{
"id": "chatcmpl-1",
"object": "chat.completion",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}}],
"usage": {
"prompt_tokens": 11,
"completion_tokens": 7,
"total_tokens": 18,
"prompt_tokens_details": {"cached_tokens": 4}
}
}`)
u, ok := ExtractUsageJSON(body, ProtoChat)
if !ok {
t.Fatalf("expected ok=true")
}
if u.InputTokens != 11 || u.OutputTokens != 7 {
t.Fatalf("chat usage = %+v, want input=11 output=7", u)
}
if u.CacheReadTokens != 4 {
t.Fatalf("chat cacheRead = %d, want 4", u.CacheReadTokens)
}
}
func TestExtractUsageJSONMessages(t *testing.T) {
body := []byte(`{
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "hi"}],
"usage": {
"input_tokens": 15,
"output_tokens": 8,
"cache_read_input_tokens": 3,
"cache_creation_input_tokens": 2
}
}`)
u, ok := ExtractUsageJSON(body, ProtoMessages)
if !ok {
t.Fatalf("expected ok=true")
}
if u.InputTokens != 15 || u.OutputTokens != 8 || u.CacheReadTokens != 3 || u.CacheCreationTokens != 2 {
t.Fatalf("messages usage = %+v", u)
}
}
func TestExtractUsageJSONResponses(t *testing.T) {
body := []byte(`{
"id": "resp_1",
"object": "response",
"output": [],
"usage": {
"input_tokens": 13,
"output_tokens": 9,
"input_tokens_details": {"cached_tokens": 5}
}
}`)
u, ok := ExtractUsageJSON(body, ProtoResponses)
if !ok {
t.Fatalf("expected ok=true")
}
if u.InputTokens != 13 || u.OutputTokens != 9 || u.CacheReadTokens != 5 {
t.Fatalf("responses usage = %+v", u)
}
}
func TestExtractUsageJSONInvalidAndMissing(t *testing.T) {
if _, ok := ExtractUsageJSON([]byte("not json"), ProtoChat); ok {
t.Fatalf("invalid json should not report ok")
}
if _, ok := ExtractUsageJSON([]byte(`{"id": "x"}`), ProtoChat); ok {
t.Fatalf("missing usage should not report ok")
}
// 空对象 usage:全 0 视为无效
if _, ok := ExtractUsageJSON([]byte(`{"usage": {}}`), ProtoChat); ok {
t.Fatalf("empty usage should not report ok")
}
}
// ---- StreamUsageAccum: 流式各协议 ----
func feedLines(t *testing.T, proto string, lines ...string) TokenUsage {
t.Helper()
acc := NewStreamUsageAccum()
for _, ln := range lines {
acc.Feed([]byte(ln), proto)
}
return acc.Usage()
}
func TestStreamUsageChatFinalChunk(t *testing.T) {
// 前面的 chunk 不带 usage;最后一个 chunk 带完整 usage
u := feedLines(t, ProtoChat,
`{"id":"c1","object":"chat.completion.chunk","choices":[{"delta":{"content":"he"}}]}`,
`{"id":"c1","object":"chat.completion.chunk","choices":[{"delta":{"content":"llo"}}]}`,
`{"id":"c1","object":"chat.completion.chunk","choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"prompt_tokens_details":{"cached_tokens":4}}}`,
)
if u.InputTokens != 11 || u.OutputTokens != 7 || u.CacheReadTokens != 4 {
t.Fatalf("chat stream usage = %+v", u)
}
}
func TestStreamUsageMessagesStartAndDelta(t *testing.T) {
// message_start 带 input/cache,message_delta 带 output;逐字段取 max 合并
u := feedLines(t, ProtoMessages,
`{"type":"message_start","message":{"id":"msg_1","usage":{"input_tokens":15,"cache_read_input_tokens":3,"cache_creation_input_tokens":2}}}`,
`{"type":"content_block_delta","delta":{"type":"text_delta","text":"hi"}}`,
`{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":8}}`,
)
if u.InputTokens != 15 || u.OutputTokens != 8 || u.CacheReadTokens != 3 || u.CacheCreationTokens != 2 {
t.Fatalf("messages stream usage = %+v", u)
}
}
func TestStreamUsageResponsesCompleted(t *testing.T) {
// response.completed 事件的用量嵌在 response.usage 下
u := feedLines(t, ProtoResponses,
`{"type":"response.output_text.delta","delta":"hi"}`,
`{"type":"response.completed","response":{"id":"resp_1","usage":{"input_tokens":13,"output_tokens":9,"input_tokens_details":{"cached_tokens":5}}}}`,
)
if u.InputTokens != 13 || u.OutputTokens != 9 || u.CacheReadTokens != 5 {
t.Fatalf("responses stream usage = %+v", u)
}
}
func TestStreamUsageIgnoresNonDataPayloads(t *testing.T) {
// [DONE]、垃圾行、空对象都不应产生用量
u := feedLines(t, ProtoChat, `[DONE]`, `{`, ``, `{"choices":[]}`)
if u.has() {
t.Fatalf("expected zero usage, got %+v", u)
}
}
func TestStreamUsageFeedKeepsMaxAcrossEvents(t *testing.T) {
// 同一字段在多个事件出现时取较大值(防乱序/重复)
u := feedLines(t, ProtoMessages,
`{"type":"message_start","message":{"usage":{"input_tokens":15}}}`,
`{"type":"message_delta","usage":{"output_tokens":5}}`,
`{"type":"message_delta","usage":{"output_tokens":8}}`,
)
if u.InputTokens != 15 || u.OutputTokens != 8 {
t.Fatalf("max-merge usage = %+v", u)
}
}
// ---- usage JSON 结构合法性(防止手写 struct 漂移)----
func TestUsageJSONRoundTrip(t *testing.T) {
u := TokenUsage{InputTokens: 10, OutputTokens: 5, CacheReadTokens: 2, CacheCreationTokens: 1}
b, err := json.Marshal(u)
if err != nil {
t.Fatalf("marshal: %v", err)
}
var back TokenUsage
if err := json.Unmarshal(b, &back); err != nil {
t.Fatalf("unmarshal: %v", err)
}
if back != u {
t.Fatalf("round trip = %+v, want %+v", back, u)
}
}
+29
View File
@@ -0,0 +1,29 @@
package proxy
import (
"opencatd-open/internal/proxy/convert"
)
// cacheWriteInputMultiplier 缓存写(cache creation)相对输入价的倍数。
// Anthropic 官方口径:缓存写按基础输入价的 1.25 倍计费(5m TTL);OpenAI 系无缓存写概念。
const cacheWriteInputMultiplier = 1.25
// ComputeCost 按上游协议的 token 语义计算一次请求的费用(USD)。
// 价格均为每百万 token 的 USD 单价。tok 的 token 语义由解析它的上游协议决定:
// - chat / responses(OpenAI 系):prompt_tokens 包含缓存读,
// 非缓存输入 = input − cacheRead;该协议没有缓存写,cacheCreation 恒为 0。
// - messages(Anthropic):input_tokens 不含缓存读/写(三个字段相互独立),
// 非缓存输入 = input 原值,不得再扣减;缓存写按输入价 ×1.25。
func ComputeCost(upstreamProto string, input, output, cacheRead, cacheCreation int, inputPrice, outputPrice, cacheReadPrice float64) float64 {
uncached := input
if upstreamProto != convert.ProtoMessages {
uncached -= cacheRead
if uncached < 0 {
uncached = 0
}
}
return (float64(uncached)*inputPrice +
float64(cacheRead)*cacheReadPrice +
float64(cacheCreation)*inputPrice*cacheWriteInputMultiplier +
float64(output)*outputPrice) / 1e6
}
+75
View File
@@ -0,0 +1,75 @@
package proxy
import "testing"
func TestComputeCost(t *testing.T) {
// 标准三价:输入 0.5 / 输出 1.5 / 缓存读 0.05($/M)
const (
inPrice = 0.5
outPrice = 1.5
cchPrice = 0.05
)
tests := []struct {
name string
upstreamProto string
input int
output int
cacheRead int
cacheCreation int
inP float64
outP float64
cchP float64
want float64
}{
{
// OpenAI:prompt 含缓存读,需扣减:(11000−10000)×0.5 + 10000×0.05 + 500×1.5
name: "openai prompt includes cache read",
upstreamProto: "chat",
input: 11000, output: 500, cacheRead: 10000,
inP: inPrice, outP: outPrice, cchP: cchPrice,
want: (1000*inPrice + 10000*cchPrice + 500*outPrice) / 1e6,
},
{
// Anthropic:input 不含缓存;缓存写按输入价 ×1.25
name: "anthropic cache write at 1.25x input price",
upstreamProto: "messages",
input: 1000, output: 500, cacheRead: 10000, cacheCreation: 2000,
inP: inPrice, outP: outPrice, cchP: cchPrice,
want: (1000*inPrice + 10000*cchPrice + 2000*inPrice*1.25 + 500*outPrice) / 1e6,
},
{
// Anthropic 口径不得扣减缓存读(否则这里非缓存输入会算成负数)
name: "anthropic does not subtract cache read",
upstreamProto: "messages",
input: 100, output: 0, cacheRead: 1000,
inP: inPrice, outP: outPrice, cchP: cchPrice,
want: (100*inPrice + 1000*cchPrice) / 1e6,
},
{
// OpenAI 异常数据:cached > prompt 时非缓存输入钳制为 0,不出现负费用
name: "openai clamps negative uncached input",
upstreamProto: "responses",
input: 100, output: 0, cacheRead: 5000,
inP: inPrice, outP: outPrice, cchP: cchPrice,
want: (5000 * cchPrice) / 1e6,
},
{
// 未配置价格时费用为 0
name: "no prices no cost",
upstreamProto: "messages",
input: 1000, output: 1000, cacheRead: 1000, cacheCreation: 1000,
want: 0,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := ComputeCost(tt.upstreamProto, tt.input, tt.output, tt.cacheRead, tt.cacheCreation, tt.inP, tt.outP, tt.cchP)
if diff := got - tt.want; diff > 1e-12 || diff < -1e-12 {
t.Fatalf("ComputeCost(%q, in=%d, out=%d, cr=%d, cw=%d) = %v, want %v",
tt.upstreamProto, tt.input, tt.output, tt.cacheRead, tt.cacheCreation, got, tt.want)
}
})
}
}
+409 -123
View File
@@ -1,8 +1,11 @@
package proxy package proxy
import ( import (
"bufio"
"bytes" "bytes"
"context" "context"
"crypto/rand"
"encoding/hex"
"encoding/json" "encoding/json"
"fmt" "fmt"
"io" "io"
@@ -13,6 +16,7 @@ import (
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/proxy/convert" "opencatd-open/internal/proxy/convert"
"opencatd-open/internal/store" "opencatd-open/internal/store"
"opencatd-open/internal/usage"
"opencatd-open/pkg/config" "opencatd-open/pkg/config"
"os" "os"
"strings" "strings"
@@ -34,7 +38,14 @@ type Gateway struct {
apiKeyDAO *dao.ApiKeyDAO apiKeyDAO *dao.ApiKeyDAO
usageDAO *dao.UsageDAO usageDAO *dao.UsageDAO
dailyDAO *dao.DailyUsageDAO dailyDAO *dao.DailyUsageDAO
modelDAO *dao.ModelDAO
channelSvc *channel.Service channelSvc *channel.Service
usageRec *usage.Recorder
// 原始请求/响应记录开关(系统配置 log_raw_requests,带 TTL 缓存避免每次查库)。
rawLogMu sync.Mutex
rawLogVal bool
rawLogSet time.Time
} }
func NewGateway(ctx context.Context, cfg *config.Config, db *gorm.DB, wg *sync.WaitGroup, userDAO *dao.UserDAO, apiKeyDAO *dao.ApiKeyDAO, usageDAO *dao.UsageDAO, dailyDAO *dao.DailyUsageDAO) *Gateway { func NewGateway(ctx context.Context, cfg *config.Config, db *gorm.DB, wg *sync.WaitGroup, userDAO *dao.UserDAO, apiKeyDAO *dao.ApiKeyDAO, usageDAO *dao.UsageDAO, dailyDAO *dao.DailyUsageDAO) *Gateway {
@@ -59,6 +70,7 @@ func NewGateway(ctx context.Context, cfg *config.Config, db *gorm.DB, wg *sync.W
apiKeyDAO: apiKeyDAO, apiKeyDAO: apiKeyDAO,
usageDAO: usageDAO, usageDAO: usageDAO,
dailyDAO: dailyDAO, dailyDAO: dailyDAO,
modelDAO: dao.NewModelDAO(db),
channelSvc: nil, channelSvc: nil,
} }
} }
@@ -67,6 +79,36 @@ func (g *Gateway) SetChannelService(svc *channel.Service) {
g.channelSvc = svc g.channelSvc = svc
} }
// SetUsageRecorder 注入异步用量记录器;nil 时网关跳过用量上报。
func (g *Gateway) SetUsageRecorder(r *usage.Recorder) {
g.usageRec = r
}
// rawLogEnabled 读取系统配置 log_raw_requests(10s TTL 缓存),决定是否记录原始请求/响应。
func (g *Gateway) rawLogEnabled() bool {
g.rawLogMu.Lock()
defer g.rawLogMu.Unlock()
if time.Since(g.rawLogSet) < 10*time.Second {
return g.rawLogVal
}
var sc store.SystemConfig
g.rawLogVal = false
if err := g.db.Where("key = ?", "log_raw_requests").First(&sc).Error; err == nil {
g.rawLogVal = strings.TrimSpace(sc.Value) == "true"
}
g.rawLogSet = time.Now()
return g.rawLogVal
}
// generateRequestID 生成请求级唯一 ID,用于用量明细关联与排障。
func generateRequestID() string {
b := make([]byte, 12)
if _, err := rand.Read(b); err != nil {
return fmt.Sprintf("req-%d", time.Now().UnixNano())
}
return "req-" + hex.EncodeToString(b)
}
// Request represents a parsed incoming request // Request represents a parsed incoming request
type Request struct { type Request struct {
Model string Model string
@@ -75,6 +117,12 @@ type Request struct {
Body []byte Body []byte
APIKey *store.APIKey APIKey *store.APIKey
UserID uint64 UserID uint64
KeyID uint64
RequestID string
CaptureRaw bool // 原始请求/响应记录(管理员 + 系统开关开启)
rawBuf *strings.Builder // 上游原始响应累积器(仅 CaptureRaw 时非 nil)
} }
// ParseRequest parses the incoming request and extracts key fields // ParseRequest parses the incoming request and extracts key fields
@@ -86,17 +134,27 @@ func (g *Gateway) ParseRequest(c *gin.Context, protocol string) (*Request, error
apiKey, _ := c.Get("api_key") apiKey, _ := c.Get("api_key")
userID, _ := c.Get("user_id") userID, _ := c.Get("user_id")
userRole, _ := c.Get("user_role")
req := &Request{ req := &Request{
Protocol: protocol, Protocol: protocol,
Body: body, Body: body,
UserID: userID.(uint64), UserID: userID.(uint64),
RequestID: c.GetHeader("X-Request-Id"),
}
if req.RequestID == "" {
req.RequestID = generateRequestID()
} }
if ak, ok := apiKey.(*store.APIKey); ok { if ak, ok := apiKey.(*store.APIKey); ok {
req.APIKey = ak req.APIKey = ak
} }
// 原始请求/响应记录:仅管理员 且 系统开关 log_raw_requests 开启。
if role, _ := userRole.(string); role == store.RoleAdmin && g.rawLogEnabled() {
req.CaptureRaw = true
}
// Parse model and stream based on protocol // Parse model and stream based on protocol
switch protocol { switch protocol {
case "chat": case "chat":
@@ -125,91 +183,237 @@ func (g *Gateway) ParseRequest(c *gin.Context, protocol string) (*Request, error
return req, nil return req, nil
} }
// Dispatch routes the request to the appropriate upstream // Dispatch routes the request to the appropriate upstream.
// 遍历候选渠道(绑定优先,全局回退;按优先级/权重排序),可重试性失败自动故障转移。
func (g *Gateway) Dispatch(c *gin.Context, req *Request) { func (g *Gateway) Dispatch(c *gin.Context, req *Request) {
if g.channelSvc == nil { if g.channelSvc == nil {
g.writeError(c, http.StatusBadGateway, "channel service not available") g.writeError(c, http.StatusBadGateway, "channel service not available")
return return
} }
ch, err := g.channelSvc.SelectChannel(g.ctx, req.Model) // 原始请求/响应捕获:仅管理员 + 系统开关开启(req.CaptureRaw 已在 ParseRequest 判定)。
if err != nil { // 客户端原始请求体即 req.Body;上游原始响应由 stream/bufferResponse 累积进 rawBuf。
g.writeError(c, http.StatusBadGateway, err.Error()) if req.CaptureRaw {
req.rawBuf = &strings.Builder{}
}
cands := g.channelSvc.Candidates(req.Model)
// 内存健康过滤:连续失败进入 cooldown 的渠道不再尝试(渠道级健康自愈靠冷却过期)。
cands = g.channelSvc.FilterHealthy(cands)
if len(cands) == 0 {
g.writeError(c, http.StatusServiceUnavailable, "no enabled channels for model: "+req.Model)
g.recordUsage(req, nil, nil, usage.Event{
IsError: true, ErrorCode: "no_channel",
}, convert.TokenUsage{}, "")
return return
} }
var lastCh *store.Channel
_ = lastCh // 保留变量名便于断点排查;失败渠道已在循环内各自 RecordFailure
lastErrStatus := http.StatusBadGateway
lastErrBody := "all upstream channels failed"
for i := range cands {
cand := &cands[i]
ch := cand.Channel
lastCh = ch
apiKey, err := g.channelSvc.GetAPIKey(ch) apiKey, err := g.channelSvc.GetAPIKey(ch)
if err != nil { if err != nil {
g.writeError(c, http.StatusBadGateway, "failed to decrypt API key") lastErrStatus, lastErrBody = http.StatusBadGateway, "failed to decrypt API key"
return continue
} }
// Determine target format and convert if needed // Determine target format: channel declares support for the client protocol
targetFormat := req.Protocol // then passthrough, otherwise convert to its first supported protocol
if len(ch.FormatsEffective()) > 0 { // (chat > messages > responses).
// Prefer the channel's native format targetFormat := g.conversionTarget(ch, req.Protocol)
for _, f := range ch.FormatsEffective() { if targetFormat == "" {
if f == req.Protocol { continue // 渠道不支持该协议,换下一个
targetFormat = f
break
}
}
} }
// Build upstream URL // Build upstream URL
upstreamPath := g.getUpstreamPath(req.Protocol) upstreamURL := ch.UpstreamURL(targetFormat, g.getUpstreamPath(targetFormat))
upstreamURL := ch.UpstreamURL(req.Protocol, upstreamPath)
// Convert request if needed // Convert request if needed
var requestBody []byte var requestBody []byte
if targetFormat != req.Protocol { if targetFormat != req.Protocol {
requestBody, err = g.convertRequest(req.Body, req.Protocol, targetFormat) var err error
requestBody, err = convert.ConvertRequest(req.Body, req.Protocol, targetFormat)
if err != nil { if err != nil {
g.writeError(c, http.StatusBadRequest, "conversion failed: "+err.Error()) lastErrStatus, lastErrBody = http.StatusBadRequest, "conversion failed: "+err.Error()
return continue
} }
} else { } else {
requestBody = req.Body requestBody = req.Body
} }
// 绑定了 upstream_model 时把请求体里的 model 重写为上游模型名(别名映射)。
if cand.Binding != nil && cand.Binding.UpstreamModel != "" &&
cand.Binding.UpstreamModel != req.Model {
requestBody = rewriteModel(requestBody, cand.Binding.UpstreamModel)
}
// Create upstream request // Create upstream request
httpReq, err := http.NewRequestWithContext(g.ctx, "POST", upstreamURL, bytes.NewReader(requestBody)) httpReq, err := http.NewRequestWithContext(g.ctx, "POST", upstreamURL, bytes.NewReader(requestBody))
if err != nil { if err != nil {
g.writeError(c, http.StatusBadGateway, "failed to create request") lastErrStatus, lastErrBody = http.StatusBadGateway, "failed to create request"
return continue
} }
// Set headers
g.setHeaders(httpReq, ch, apiKey, targetFormat) g.setHeaders(httpReq, ch, apiKey, targetFormat)
// Execute request // Execute request
start := time.Now() start := time.Now()
resp, err := g.httpClient.Do(httpReq) resp, err := g.httpClient.Do(httpReq)
latency := time.Since(start)
if err != nil { if err != nil {
g.channelSvc.RecordFailure(ch.ID) g.channelSvc.RecordFailure(ch.ID)
g.writeError(c, http.StatusBadGateway, fmt.Sprintf("upstream error: %v (latency: %v)", err, latency)) lastErrStatus = http.StatusBadGateway
return lastErrBody = fmt.Sprintf("upstream error: %v", err)
g.recordUsage(req, cand, ch, usage.Event{
IsError: true,
ErrorCode: "upstream_error",
LatencyMS: int(time.Since(start).Milliseconds()),
}, convert.TokenUsage{}, targetFormat)
continue // 可重试:换下一个渠道
} }
defer resp.Body.Close()
// Record success // Handle upstream error responses
g.channelSvc.RecordSuccess(ch.ID)
// Handle response
if resp.StatusCode >= 400 { if resp.StatusCode >= 400 {
body, _ := io.ReadAll(resp.Body) body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
log.Printf("Upstream error: status=%d body=%s", resp.StatusCode, string(body)) log.Printf("Upstream error: status=%d body=%s", resp.StatusCode, string(body))
if req.rawBuf != nil {
req.rawBuf.Write(body)
}
g.recordUsage(req, cand, ch, usage.Event{
IsError: true,
ErrorCode: fmt.Sprintf("upstream_%d", resp.StatusCode),
LatencyMS: int(time.Since(start).Milliseconds()),
}, convert.TokenUsage{}, targetFormat)
// 429/5xx 可换渠道重试;4xx 直接透传
if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 {
lastErrStatus, lastErrBody = resp.StatusCode, string(body)
continue
}
c.Data(resp.StatusCode, "application/json", body) c.Data(resp.StatusCode, "application/json", body)
return return
} }
// Stream or buffer response g.channelSvc.RecordSuccess(ch.ID)
// Stream or buffer response;tok 从上游响应(SSE usage 块或非流式 JSON)提取。
// 上游可能返回 HTTP 200 但 body/SSE 内带 error(OpenRouter 超时等),
// 此时按失败记账(errCode 非空),非流式错误体以 502 返回给客户端。
var tok convert.TokenUsage
var errCode string
if req.Stream { if req.Stream {
g.streamResponse(c, resp, req.Protocol, ch) tok, errCode = g.streamResponse(c, resp, req.Protocol, targetFormat, req.rawBuf)
} else { } else {
g.bufferResponse(c, resp, req.Protocol, ch) tok, errCode, _ = g.bufferResponse(c, resp, req.Protocol, targetFormat, req.rawBuf)
} }
resp.Body.Close()
if errCode != "" {
// 记账为失败(错误码),不产生费用;响应内容已由 buffer/stream 写出
g.recordUsage(req, cand, ch, usage.Event{
IsError: true,
ErrorCode: errCode,
LatencyMS: int(time.Since(start).Milliseconds()),
}, tok, targetFormat)
return
}
// 成功记录:用量 + 定价计费。
g.recordUsage(req, cand, ch, usage.Event{
LatencyMS: int(time.Since(start).Milliseconds()),
}, tok, targetFormat)
return
}
// 全部候选失败(每个候选失败时已各自 RecordFailure,不再重复计数)
g.writeError(c, lastErrStatus, lastErrBody)
}
// rewriteModel 把 JSON 请求体顶层的 model 字段替换为 upstreamModel。
func rewriteModel(body []byte, upstreamModel string) []byte {
var m map[string]json.RawMessage
if json.Unmarshal(body, &m) != nil {
return body
}
if _, ok := m["model"]; !ok {
return body
}
m["model"], _ = json.Marshal(upstreamModel)
out, err := json.Marshal(m)
if err != nil {
return body
}
return out
}
// recordUsage 汇总一次请求的用量事件并异步落库。tok 为从上游响应提取的用量,
// 其 token 语义由 upstreamProto(渠道实际使用的上游协议)决定。
// cand/ch 可为 nil(无可用渠道的失败场景)。
func (g *Gateway) recordUsage(req *Request, cand *channel.Candidate, ch *store.Channel, ev usage.Event, tok convert.TokenUsage, upstreamProto string) {
if g.usageRec == nil {
return
}
ev.UserID = req.UserID
ev.ModelName = req.Model
ev.Protocol = req.Protocol
ev.RequestID = req.RequestID
if req.APIKey != nil {
ev.KeyID = req.APIKey.ID
}
if ch != nil {
ev.ChannelID = ch.ID
}
if cand != nil && cand.Binding != nil {
ev.ModelID = cand.Binding.ModelID
}
ev.PromptTokens = tok.InputTokens
ev.CompletionTokens = tok.OutputTokens
ev.CacheReadTokens = tok.CacheReadTokens
ev.CacheCreationTokens = tok.CacheCreationTokens
// 原始请求/响应(仅管理员+开关开启时捕获)。
if req.CaptureRaw {
ev.RawRequest = string(req.Body)
if req.rawBuf != nil {
ev.RawResponse = req.rawBuf.String()
}
}
// 定价与成本(价格按每百万 token 的 USD 单价)。
// 成本口径按上游协议区分(详见 ComputeCost):OpenAI 系 prompt 含缓存读需扣减;
// Anthropic 的 input_tokens 不含缓存,缓存写按输入价 ×1.25。
if ev.ModelID != 0 {
if m, err := g.modelDAO.GetByID(ev.ModelID); err == nil {
ev.InputPrice = m.InputPrice
ev.OutputPrice = m.OutputPrice
ev.CacheReadPrice = m.CacheReadPrice
}
}
if !ev.IsError {
ev.Cost = ComputeCost(upstreamProto, tok.InputTokens, tok.OutputTokens, tok.CacheReadTokens, tok.CacheCreationTokens,
ev.InputPrice, ev.OutputPrice, ev.CacheReadPrice)
}
g.usageRec.Record(ev)
}
// conversionTarget 决定客户端协议在渠道上的处理方式:
// 渠道声明支持该协议则直通;否则转为其首选支持协议(chat > messages > responses)。
func (g *Gateway) conversionTarget(ch *store.Channel, clientProto string) string {
formats := ch.FormatsEffective()
for _, f := range formats {
if f == clientProto {
return clientProto
}
}
for _, p := range []string{convert.ProtoChat, convert.ProtoMessages, convert.ProtoResponses} {
for _, f := range formats {
if f == p {
return p
}
}
}
return ""
} }
func (g *Gateway) getUpstreamPath(protocol string) string { func (g *Gateway) getUpstreamPath(protocol string) string {
@@ -237,116 +441,198 @@ func (g *Gateway) setHeaders(req *http.Request, ch *store.Channel, apiKey string
} }
} }
func (g *Gateway) convertRequest(body []byte, from, to string) ([]byte, error) {
switch {
case from == "chat" && to == "messages":
var req convert.ChatCompletionRequest
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
msgReq, err := convert.ChatToMessages(&req)
if err != nil {
return nil, err
}
return json.Marshal(msgReq)
case from == "chat" && to == "responses": // streamResponse 流式响应:按 \n\n 分块零缓冲转发;跨协议时逐行转换。
var req convert.ChatCompletionRequest // 返回从上游 SSE usage 块累计的 token 用量(按上游协议解析)。
if err := json.Unmarshal(body, &req); err != nil { // capture 非 nil 时把上游原始行累积进去(原始响应记录)。
return nil, err // 上游部分实现(如 OpenRouter)在超时时返回 HTTP 200 但 SSE data 内带
} // error 字段;检测到则返回错误码,供 Dispatch 按失败记账。
respReq, err := convert.ChatToResponses(&req) func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, clientProto, upstreamProto string, capture *strings.Builder) (convert.TokenUsage, string) {
if err != nil { w := c.Writer
return nil, err
}
return json.Marshal(respReq)
case from == "messages" && to == "chat":
var req convert.MessagesRequest
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
// Messages -> Chat: we need to construct a ChatCompletionRequest
chatReq := &convert.ChatCompletionRequest{
Model: req.Model,
}
for _, m := range req.Messages {
chatReq.Messages = append(chatReq.Messages, m)
}
if req.Temperature != nil {
chatReq.Temperature = req.Temperature
}
if req.TopP != nil {
chatReq.TopP = req.TopP
}
chatReq.Tools = req.Tools
chatReq.Stream = req.Stream
return json.Marshal(chatReq)
case from == "responses" && to == "chat":
var req convert.ResponsesRequest
if err := json.Unmarshal(body, &req); err != nil {
return nil, err
}
chatReq := &convert.ChatCompletionRequest{
Model: req.Model,
}
for _, item := range req.Input {
chatReq.Messages = append(chatReq.Messages, convert.Message{
Role: item.Role,
Content: item.Content,
})
}
chatReq.Tools = req.Tools
chatReq.Stream = req.Stream
return json.Marshal(chatReq)
default:
return body, nil
}
}
func (g *Gateway) streamResponse(c *gin.Context, resp *http.Response, protocol string, ch *store.Channel) {
c.Header("Content-Type", "text/event-stream") c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache") c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive") c.Header("Connection", "keep-alive")
c.Status(http.StatusOK) c.Status(http.StatusOK)
writer := convert.NewSSEWriter(c.Writer) flusher, _ := w.(http.Flusher)
parser := convert.NewSSEParser(resp.Body)
// 跨协议时按行转换;同协议直通(lineConv 为 nil)。
var lineConv func([]byte) []byte
if upstreamProto != clientProto {
lineConv = convert.NewStreamTransformer(upstreamProto, clientProto)
}
// 上游原始行按 \n\n 分块,避免把 data 行内的转义换行当成事件边界。
// 同时喂入用量累计器(usage 块可能出现在任一事件)。
r := bufio.NewReaderSize(resp.Body, 32*1024)
accum := convert.NewStreamUsageAccum()
errCode := ""
for { for {
event, err := parser.ReadEvent() buf := []byte{}
if err != nil { for {
line, err := r.ReadSlice('\n')
if err == bufio.ErrBufferFull {
buf = append(buf, line...)
continue
}
buf = append(buf, line...)
if err == io.EOF { if err == io.EOF {
break if len(buf) == 0 {
return accum.Usage(), errCode
} }
log.Printf("Stream parse error: %v", err) if !bytes.HasSuffix(buf, []byte("\n")) {
break buf = append(buf, '\n')
} }
} else if err != nil {
if event.Event == "error" { log.Printf("stream read error: %v", err)
log.Printf("Upstream stream error: %s", event.Data) return accum.Usage(), errCode
break
} }
if len(buf) >= 2 && bytes.HasSuffix(buf, []byte("\n\n")) {
// Write raw SSE event based on protocol
if err := writer.WriteEvent("chat CompletionChunk", event.Data); err != nil {
break break
} }
} }
writer.WriteDone() // 原始响应捕获(仅管理员+开关开启时启用)。
if capture != nil {
capture.Write(buf)
}
// 先解析用量(data: {...} 行),再决定转发内容。
for _, data := range sseDataPayloads(buf) {
accum.Feed(data, upstreamProto)
if errCode == "" && streamChunkHasError(data) {
errCode = "upstream_stream_error"
}
}
out := buf
if lineConv != nil {
out = lineConv(buf)
}
if len(out) == 0 {
continue
}
if _, err := w.Write(out); err != nil {
return accum.Usage(), errCode // 客户端已断开
}
if flusher != nil {
flusher.Flush()
}
// 流结束标记:chat/messages 上游以 data: [DONE] 收尾。部分上游(keep-alive)
// 发完 [DONE] 后不关连接,继续读会阻塞到超时;据此主动收尾。
// responses 协议没有 [DONE],以 response.completed 事件收尾。
if streamTerminated(buf, upstreamProto) {
return accum.Usage(), errCode
}
}
} }
func (g *Gateway) bufferResponse(c *gin.Context, resp *http.Response, protocol string, ch *store.Channel) { // streamChunkHasError 判断一块 SSE data 载荷是否带 error 字段(OpenRouter 超时等)。
func streamChunkHasError(data []byte) bool {
var m map[string]any
if json.Unmarshal(data, &m) != nil {
return false
}
if _, ok := m["error"]; ok {
return true
}
// responses 协议错误事件可能形如 {"type":"error",...}
return m["type"] == "error"
}
// streamTerminated 判断一块 SSE 是否为上游流的结束事件。
func streamTerminated(chunk []byte, proto string) bool {
switch proto {
case convert.ProtoChat:
// chat 上游以 data: [DONE] 收尾;keep-alive 上游发完不关连接。
return bytes.Contains(chunk, []byte("data: [DONE]"))
case convert.ProtoMessages:
// messages 上游以 message_stop 事件结束(无 [DONE])。
return bytes.Contains(chunk, []byte(`"type":"message_stop"`)) ||
bytes.Contains(chunk, []byte(`"type": "message_stop"`)) ||
bytes.Contains(chunk, []byte("data: [DONE]"))
case convert.ProtoResponses:
return bytes.Contains(chunk, []byte(`"response.completed"`)) ||
bytes.Contains(chunk, []byte(`"type":"response.completed"`))
}
return false
}
// sseDataPayloads 从一块 SSE(一个完整事件,\n\n 结尾)中取出所有 data 行的原始载荷。
func sseDataPayloads(chunk []byte) [][]byte {
var out [][]byte
for _, line := range bytes.Split(chunk, []byte("\n")) {
line = bytes.TrimSuffix(line, []byte("\r"))
if !bytes.HasPrefix(line, []byte("data:")) {
continue
}
payload := bytes.TrimPrefix(line, []byte("data:"))
payload = bytes.TrimPrefix(payload, []byte(" "))
if len(payload) == 0 || bytes.Equal(payload, []byte("[DONE]")) {
continue
}
out = append(out, payload)
}
return out
}
// bufferResponse 非流式响应:整体读取、可选转换后写回。
// 返回 (用量, 错误码, 是否错误)。部分上游(如 OpenRouter)在超时时返回
// HTTP 200 但 JSON 内含 error 字段,需要识别并让调用方按失败处理。
func (g *Gateway) bufferResponse(c *gin.Context, resp *http.Response, clientProto, upstreamProto string, capture *strings.Builder) (convert.TokenUsage, string, bool) {
body, err := io.ReadAll(resp.Body) body, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
g.writeError(c, http.StatusBadGateway, "failed to read response") g.writeError(c, http.StatusBadGateway, "failed to read response")
return return convert.TokenUsage{}, "", false
} }
c.Data(resp.StatusCode, "application/json", body) // 原始响应捕获(仅管理员+开关开启时启用)。
if capture != nil {
capture.Write(body)
}
// 用量从上游原始响应体提取(先于转换,转换会改字段名)。
tok, _ := convert.ExtractUsageJSON(body, upstreamProto)
// HTTP 200 但带 error 字段(OpenRouter 超时 504 等):识别并转失败。
errCode, isErr := bodyHasError(body)
out := body
if upstreamProto != clientProto {
if converted, cerr := convert.ConvertResponse(body, upstreamProto, clientProto); cerr == nil {
out = converted
} else {
// 转换失败时至少剥掉非 JSON 前缀,让客户端能解析出正文
out = convert.CleanJSON(body)
}
} else {
// 直通:部分上游(如 OpenRouter)的 non-stream 响应在 JSON 前夹带空白/注释
out = convert.CleanJSON(body)
}
// 上游错误体:用 502 返回,让客户端感知失败(不伪装成 200)。
if isErr {
c.Data(http.StatusBadGateway, "application/json", out)
return tok, errCode, true
}
c.Data(resp.StatusCode, "application/json", out)
return tok, errCode, false
}
// bodyHasError 判断 JSON 响应体是否带 error 字段(openai 风格 {"error":{...}} 或
// anthropic 风格 {"type":"error",...})。返回 (错误码, 是否错误)。找不到 JSON 返回 ("", false)。
func bodyHasError(body []byte) (string, bool) {
var m map[string]any
if json.Unmarshal(bytes.TrimSpace(body), &m) != nil {
return "", false
}
if _, ok := m["error"]; ok {
return "upstream_error", true
}
if m["type"] == "error" {
return "upstream_error", true
}
return "", false
} }
func (g *Gateway) writeError(c *gin.Context, status int, message string) { func (g *Gateway) writeError(c *gin.Context, status int, message string) {
+22 -2
View File
@@ -2,6 +2,8 @@ package proxy
import ( import (
"net/http" "net/http"
"opencatd-open/internal/dao"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -41,9 +43,27 @@ func (g *Gateway) HandleResponses(c *gin.Context) {
// HandleModels handles GET /v1/models // HandleModels handles GET /v1/models
func (g *Gateway) HandleModels(c *gin.Context) { func (g *Gateway) HandleModels(c *gin.Context) {
// TODO: Return list of available models based on enabled channels modelDAO := dao.NewModelDAO(g.db)
models, _, err := modelDAO.List(1000, 0)
if err != nil {
g.writeError(c, http.StatusInternalServerError, "failed to list models")
return
}
now := time.Now().Unix()
data := make([]gin.H, 0, len(models))
for _, m := range models {
if !m.Enabled {
continue
}
data = append(data, gin.H{
"id": m.Name,
"object": "model",
"created": now,
"owned_by": "opencatd-open",
})
}
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"object": "list", "object": "list",
"data": []interface{}{}, "data": data,
}) })
} }
+7 -1
View File
@@ -2,6 +2,7 @@ package service
import ( import (
"context" "context"
"fmt"
"opencatd-open/internal/channel" "opencatd-open/internal/channel"
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/store" "opencatd-open/internal/store"
@@ -55,7 +56,12 @@ func (s *ChannelServiceImpl) GetAPIKey(ctx context.Context, channelID uint64) (s
// SelectForModel selects the best channel for a model // SelectForModel selects the best channel for a model
func (s *ChannelServiceImpl) SelectForModel(ctx context.Context, modelName string) (*store.Channel, error) { func (s *ChannelServiceImpl) SelectForModel(ctx context.Context, modelName string) (*store.Channel, error) {
return s.channelSvc.SelectChannel(ctx, modelName) cands := s.channelSvc.Candidates(modelName)
picked := s.channelSvc.Pick(cands)
if picked == nil {
return nil, fmt.Errorf("no enabled channels for model: %s", modelName)
}
return picked.Channel, nil
} }
// BindModels binds models to a channel // BindModels binds models to a channel
-203
View File
@@ -1,203 +0,0 @@
package service
import (
"encoding/base64"
"fmt"
"net/http"
"opencatd-open/internal/store"
"opencatd-open/pkg/config"
"strconv"
"strings"
"time"
"github.com/go-webauthn/webauthn/protocol"
"github.com/go-webauthn/webauthn/webauthn"
"gorm.io/gorm"
)
type WebAuthnUser struct {
User *store.User
Credentials []webauthn.Credential
}
func (u *WebAuthnUser) WebAuthnID() []byte {
return []byte(strconv.FormatUint(u.User.ID, 10))
}
func (u *WebAuthnUser) WebAuthnName() string {
return u.User.Username
}
func (u *WebAuthnUser) WebAuthnDisplayName() string {
return u.User.Username
}
func (u *WebAuthnUser) WebAuthnCredentials() []webauthn.Credential {
return u.Credentials
}
func (u *WebAuthnUser) WebAuthnCredentialDescriptors() (descriptors []protocol.CredentialDescriptor) {
credentials := u.WebAuthnCredentials()
descriptors = make([]protocol.CredentialDescriptor, len(credentials))
for i, credential := range credentials {
descriptors[i] = credential.Descriptor()
}
return descriptors
}
type WebAuthnService struct {
cfg *config.Config
DB *gorm.DB
WebAuthn *webauthn.WebAuthn
}
func NewWebAuthnService(cfg *config.Config, db *gorm.DB) (*WebAuthnService, error) {
wconfig := &webauthn.Config{
RPDisplayName: cfg.AppName,
RPID: cfg.RPID,
RPOrigins: cfg.RPOrigins,
AuthenticatorSelection: protocol.AuthenticatorSelection{
RequireResidentKey: protocol.ResidentKeyRequired(),
ResidentKey: protocol.ResidentKeyRequirementRequired,
UserVerification: protocol.VerificationPreferred,
},
}
wa, err := webauthn.New(wconfig)
if err != nil {
return nil, err
}
return &WebAuthnService{
cfg: cfg,
DB: db,
WebAuthn: wa,
}, nil
}
func (s *WebAuthnService) GetUserWithCredentials(userID uint64) (*WebAuthnUser, error) {
var user store.User
if err := s.DB.First(&user, userID).Error; err != nil {
return nil, err
}
var passkeys []store.Passkey
if err := s.DB.Where("user_id = ?", userID).Find(&passkeys).Error; err != nil {
return nil, err
}
credentials := make([]webauthn.Credential, len(passkeys))
for i, pk := range passkeys {
credentialIDBytes, err := base64.StdEncoding.DecodeString(pk.CredentialID)
if err != nil {
return nil, fmt.Errorf("failed to decode CredentialID: %w", err)
}
publicKeyBytes, err := base64.StdEncoding.DecodeString(pk.PublicKey)
if err != nil {
return nil, fmt.Errorf("failed to decode PublicKey: %w", err)
}
aaguidBytes, err := base64.StdEncoding.DecodeString(pk.AAGUID)
if err != nil {
return nil, fmt.Errorf("failed to decode AAGUID: %w", err)
}
var transport []protocol.AuthenticatorTransport
if pk.Transport != "" {
transport = []protocol.AuthenticatorTransport{protocol.AuthenticatorTransport(pk.Transport)}
}
credentials[i] = webauthn.Credential{
ID: credentialIDBytes,
PublicKey: publicKeyBytes,
AttestationType: pk.AttestationType,
Transport: transport,
Flags: webauthn.CredentialFlags{
UserPresent: true,
UserVerified: true,
BackupEligible: pk.BackupEligible,
BackupState: pk.BackupState,
},
Authenticator: webauthn.Authenticator{
AAGUID: aaguidBytes,
SignCount: uint32(pk.SignCount),
},
}
}
return &WebAuthnUser{
User: &user,
Credentials: credentials,
}, nil
}
func (s *WebAuthnService) BeginRegistration(userID uint64) (*protocol.CredentialCreation, error) {
user, err := s.GetUserWithCredentials(userID)
if err != nil {
return nil, err
}
options, _, err := s.WebAuthn.BeginRegistration(user)
if err != nil {
return nil, err
}
return options, nil
}
func (s *WebAuthnService) FinishRegistration(userID uint64, response *http.Request, deviceName string) (*store.Passkey, error) {
user, err := s.GetUserWithCredentials(userID)
if err != nil {
return nil, err
}
credential, err := s.WebAuthn.FinishRegistration(user, webauthn.SessionData{}, response)
if err != nil {
return nil, err
}
var transport string
if len(credential.Transport) > 0 {
transport = string(credential.Transport[0])
}
passkey := &store.Passkey{
UserID: userID,
CredentialID: base64.StdEncoding.EncodeToString(credential.ID),
PublicKey: base64.StdEncoding.EncodeToString(credential.PublicKey),
AttestationType: string(credential.AttestationType),
AAGUID: base64.StdEncoding.EncodeToString(credential.Authenticator.AAGUID),
SignCount: uint64(credential.Authenticator.SignCount),
Name: deviceName,
DeviceType: strings.TrimSpace(fmt.Sprintf("%s", deviceName)),
LastUsedAt: time.Now().Unix(),
BackupEligible: credential.Flags.BackupEligible,
BackupState: credential.Flags.BackupState,
Transport: transport,
}
if err := s.DB.Create(passkey).Error; err != nil {
return nil, err
}
return passkey, nil
}
func (s *WebAuthnService) BeginLogin() (*protocol.CredentialAssertion, error) {
options, _, err := s.WebAuthn.BeginDiscoverableLogin()
if err != nil {
return nil, err
}
return options, nil
}
func (s *WebAuthnService) ListPasskeys(userID uint64) ([]store.Passkey, error) {
var passkeys []store.Passkey
if err := s.DB.Where("user_id = ?", userID).Find(&passkeys).Error; err != nil {
return nil, err
}
return passkeys, nil
}
func (s *WebAuthnService) DeletePasskey(userID uint64, passkeyID uint64) error {
return s.DB.Where("id = ? AND user_id = ?", passkeyID, userID).Delete(&store.Passkey{}).Error
}
+19 -5
View File
@@ -3,6 +3,8 @@ package store
import ( import (
"fmt" "fmt"
"log" "log"
"os"
"path/filepath"
"opencatd-open/pkg/config" "opencatd-open/pkg/config"
_ "github.com/lib/pq" _ "github.com/lib/pq"
@@ -15,11 +17,17 @@ import (
var DB *gorm.DB var DB *gorm.DB
func InitDB(cfg *config.Config) (*gorm.DB, error) { func InitDB(cfg *config.Config) (*gorm.DB, error) {
var dialector gorm.Dialector var (
dialector gorm.Dialector
err error
)
switch cfg.DB_Type { switch cfg.DB_Type {
case "sqlite": case "sqlite":
dialector = sqliteDialector(cfg.DSN) dialector, err = sqliteDialector(cfg.DSN)
if err != nil {
return nil, err
}
case "postgres": case "postgres":
dialector = postgresDialector(cfg.DSN) dialector = postgresDialector(cfg.DSN)
case "mysql": case "mysql":
@@ -48,11 +56,17 @@ func InitDB(cfg *config.Config) (*gorm.DB, error) {
return db, nil return db, nil
} }
func sqliteDialector(dsn string) gorm.Dialector { func sqliteDialector(dsn string) (gorm.Dialector, error) {
if dsn == "" { if dsn == "" {
dsn = "opencatd.db" dsn = "db/openteam.db"
} }
return gormlite.Open(dsn) // sqlite 不会自动创建上级目录,先确保它存在(与 docker-compose 挂载的 /app/db 对应)
if dir := filepath.Dir(dsn); dir != "." && dir != string(filepath.Separator) {
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("failed to create database directory %s: %w", dir, err)
}
}
return gormlite.Open(dsn), nil
} }
func postgresDialector(dsn string) gorm.Dialector { func postgresDialector(dsn string) gorm.Dialector {
+7 -11
View File
@@ -79,7 +79,9 @@ type Channel struct {
BaseURL string `gorm:"size:255;not null" json:"base_url"` BaseURL string `gorm:"size:255;not null" json:"base_url"`
BaseURLs map[string]string `gorm:"type:jsonb;serializer:json" json:"base_urls,omitempty"` BaseURLs map[string]string `gorm:"type:jsonb;serializer:json" json:"base_urls,omitempty"`
APIKeyEnc string `gorm:"size:1024;not null" json:"-"` APIKeyEnc string `gorm:"size:1024;not null" json:"-"`
Weight int `gorm:"not null;default:1" json:"weight"` // Weight 为 0 表示不参与加权随机选择(探活/回退语义),因此不能加 gorm
// default 标签 —— 零值字段会被 default 值覆盖,导致 0 被静默改写为 1。
Weight int `gorm:"not null" json:"weight"`
Priority int `gorm:"not null;default:0" json:"priority"` Priority int `gorm:"not null;default:0" json:"priority"`
TimeoutMS int `gorm:"not null;default:120000" json:"timeout_ms"` TimeoutMS int `gorm:"not null;default:120000" json:"timeout_ms"`
MaxConcurrency int `gorm:"not null;default:16" json:"max_concurrency"` MaxConcurrency int `gorm:"not null;default:16" json:"max_concurrency"`
@@ -172,6 +174,8 @@ type UsageLog struct {
LatencyMS int `json:"latency_ms"` LatencyMS int `json:"latency_ms"`
Status string `gorm:"size:16;not null" json:"status"` Status string `gorm:"size:16;not null" json:"status"`
ErrorCode *string `json:"error_code,omitempty"` ErrorCode *string `json:"error_code,omitempty"`
RawRequest string `gorm:"type:text" json:"raw_request,omitempty"` // 客户端原始请求体(未转换;仅管理员+开关开启时记录)
RawResponse string `gorm:"type:text" json:"raw_response,omitempty"` // 上游原始响应(未转换;流式为全部 SSE 事件)
CreatedAt time.Time `gorm:"index" json:"created_at"` CreatedAt time.Time `gorm:"index" json:"created_at"`
} }
@@ -193,16 +197,8 @@ type Passkey struct {
ID uint64 `gorm:"primaryKey;autoIncrement" json:"id"` ID uint64 `gorm:"primaryKey;autoIncrement" json:"id"`
UserID uint64 `gorm:"index;not null" json:"user_id"` UserID uint64 `gorm:"index;not null" json:"user_id"`
Name string `gorm:"size:64" json:"name"` Name string `gorm:"size:64" json:"name"`
CredentialID string `gorm:"size:255;not null" json:"-"` CredentialID []byte `gorm:"size:255;not null" json:"-"`
PublicKey string `gorm:"size:512;not null" json:"-"` Credential []byte `gorm:"type:blob;not null" json:"-"`
AttestationType string `gorm:"size:64" json:"-"`
AAGUID string `gorm:"size:64" json:"-"`
SignCount uint64 `json:"-"`
DeviceType string `gorm:"size:255" json:"device_type,omitempty"`
LastUsedAt int64 `json:"last_used_at,omitempty"`
BackupEligible bool `json:"-"`
BackupState bool `json:"-"`
Transport string `gorm:"size:32" json:"-"`
CreatedAt time.Time `json:"created_at"` CreatedAt time.Time `json:"created_at"`
} }
+65 -1
View File
@@ -2,6 +2,7 @@ package usage
import ( import (
"context" "context"
"fmt"
"log" "log"
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/store" "opencatd-open/internal/store"
@@ -17,10 +18,22 @@ type Event struct {
PromptTokens int PromptTokens int
CompletionTokens int CompletionTokens int
CacheReadTokens int CacheReadTokens int
CacheCreationTokens int
Cost float64 Cost float64
IsError bool IsError bool
IsCanceled bool IsCanceled bool
RequestID string RequestID string
KeyID uint64
Protocol string
ErrorCode string
LatencyMS int
InputPrice float64
OutputPrice float64
CacheReadPrice float64
TraceID string // TraceID for distributed tracing
ModelID uint64 // Model ID from channel-model binding
RawRequest string // 客户端原始请求体(仅管理员+开关开启时记录)
RawResponse string // 上游原始响应(未转换;流式为全部 SSE 事件)
} }
// Recorder handles async usage recording // Recorder handles async usage recording
@@ -117,16 +130,33 @@ func (r *Recorder) flush(events []Event) {
status = store.UsageStatusCanceled status = store.UsageStatusCanceled
} }
var errCode *string
if e.ErrorCode != "" {
errCode = &e.ErrorCode
}
log := &store.UsageLog{ log := &store.UsageLog{
UserID: e.UserID, UserID: e.UserID,
ModelName: e.ModelName, KeyID: e.KeyID,
ChannelID: e.ChannelID, ChannelID: e.ChannelID,
ModelID: e.ModelID,
ModelName: e.ModelName,
Protocol: e.Protocol,
InputTokens: int64(e.PromptTokens), InputTokens: int64(e.PromptTokens),
OutputTokens: int64(e.CompletionTokens), OutputTokens: int64(e.CompletionTokens),
CacheReadTokens: int64(e.CacheReadTokens), CacheReadTokens: int64(e.CacheReadTokens),
CacheCreationTokens: int64(e.CacheCreationTokens),
InputPrice: e.InputPrice,
OutputPrice: e.OutputPrice,
CacheReadPrice: e.CacheReadPrice,
Cost: e.Cost, Cost: e.Cost,
LatencyMS: e.LatencyMS,
Status: status, Status: status,
ErrorCode: errCode,
RequestID: e.RequestID, RequestID: e.RequestID,
TraceID: e.TraceID,
RawRequest: e.RawRequest,
RawResponse: e.RawResponse,
} }
logs = append(logs, log) logs = append(logs, log)
} }
@@ -136,5 +166,39 @@ func (r *Recorder) flush(events []Event) {
log.Printf("Failed to batch create usage logs: %v", err) log.Printf("Failed to batch create usage logs: %v", err)
} }
// Daily rollup for success and canceled requests
dailyMap := make(map[string]*store.UsageDaily)
for _, e := range events {
if e.IsError {
continue
}
date := time.Now().Format("2006-01-02")
key := fmt.Sprintf("%d:%d:%s", e.UserID, e.ModelID, date)
d := dailyMap[key]
if d == nil {
d = &store.UsageDaily{
UserID: e.UserID,
ModelID: e.ModelID,
Date: date,
Requests: 0,
InputTokens: 0,
OutputTokens: 0,
CacheReadTokens: 0,
Cost: 0,
}
dailyMap[key] = d
}
d.Requests++
d.InputTokens += int64(e.PromptTokens)
d.OutputTokens += int64(e.CompletionTokens)
d.CacheReadTokens += int64(e.CacheReadTokens)
d.Cost += e.Cost
}
for _, d := range dailyMap {
if err := r.dailyDAO.UpsertDailyUsage(context.Background(), d); err != nil {
log.Printf("Failed to upsert daily usage: %v", err)
}
}
log.Printf("Flushed %d usage logs", len(logs)) log.Printf("Flushed %d usage logs", len(logs))
} }
+19
View File
@@ -48,6 +48,25 @@ func Auth(db *gorm.DB) gin.HandlerFunc {
func CheckRole(role string) gin.HandlerFunc { func CheckRole(role string) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
userRole, _ := c.Get("user_role")
if roleStr, ok := userRole.(string); !ok || roleStr != role {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "权限不足"})
return
}
c.Next()
}
}
// AdminOnly 管理后台中间件:要求 user_role 为 admin。
// 由 middleware.Auth 先行设置 user_role;缺失时拒绝。
func AdminOnly() gin.HandlerFunc {
return func(c *gin.Context) {
role, _ := c.Get("user_role")
roleStr, _ := role.(string)
if roleStr != store.RoleAdmin {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "需要管理员权限"})
return
}
c.Next() c.Next()
} }
} }
+39 -33
View File
@@ -9,59 +9,65 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
// keyPrefixLen 是 key_prefix 列的截断长度,必须与 api.go:459 的 keyValue[:12] 一致。
// 真实 key 为 sk-ot- + 48 位 hex(54 字符),故 12 位足够唯一。
const keyPrefixLen = 12
func AuthLLM(db *gorm.DB) gin.HandlerFunc { func AuthLLM(db *gorm.DB) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
authToken := c.GetHeader("Authorization") key := extractAPIKey(c.GetHeader("Authorization"))
if authToken == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ // 区分「没传」和「传了但不对」,便于排查客户端配置。
"error": map[string]interface{}{ if strings.TrimSpace(c.GetHeader("Authorization")) == "" {
"message": "未提供认证信息", unauthorized(c, "未提供认证信息")
"type": "invalid_request_error", return
}, }
}) // 长度不足时直接拒绝:避免下方 authToken[:12] 越界 panic 打崩进程。
if len(key) < keyPrefixLen {
unauthorized(c, "无效的API密钥")
return return
} }
// Extract API key from Bearer token
if len(authToken) > 7 {
authToken = authToken[7:]
}
// Find API key by prefix
var apiKey store.APIKey var apiKey store.APIKey
if err := db.Where("key_prefix = ? AND status = ?", authToken[:8], store.KeyStatusActive).First(&apiKey).Error; err != nil { if err := db.Where("key_prefix = ? AND status = ?", key[:keyPrefixLen], store.KeyStatusActive).First(&apiKey).Error; err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ unauthorized(c, "无效的API密钥")
"error": map[string]interface{}{
"message": "无效的API密钥",
"type": "invalid_request_error",
},
})
return return
} }
// Verify full key hash // Verify full key hash
keyHash := store.HashAPIKey(authToken) if apiKey.KeyHash != store.HashAPIKey(key) {
if apiKey.KeyHash != keyHash { unauthorized(c, "无效的API密钥")
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{ return
"error": map[string]interface{}{ }
"message": "无效的API密钥",
"type": "invalid_request_error", // 附带用户角色(判断是否管理员,供原始请求/响应记录等管理能力使用)。
}, var user store.User
}) if err := db.First(&user, apiKey.UserID).Error; err != nil {
unauthorized(c, "无效的API密钥")
return return
} }
c.Set("api_key", &apiKey) c.Set("api_key", &apiKey)
c.Set("user_id", apiKey.UserID) c.Set("user_id", apiKey.UserID)
c.Set("user_role", user.Role)
c.Next() c.Next()
} }
} }
// extractAPIKey extracts the API key from the Authorization header // extractAPIKey 从 Authorization 头取 Bearer token,兼容无 "Bearer " 前缀的直传。
func extractAPIKey(c *gin.Context) string { func extractAPIKey(auth string) string {
auth := c.GetHeader("Authorization") auth = strings.TrimSpace(auth)
if strings.HasPrefix(auth, "Bearer ") { if strings.HasPrefix(auth, "Bearer ") {
return auth[7:] return strings.TrimSpace(auth[len("Bearer "):])
} }
return auth return auth
} }
func unauthorized(c *gin.Context, message string) {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
"error": map[string]interface{}{
"message": message,
"type": "invalid_request_error",
},
})
}
+76 -4
View File
@@ -10,6 +10,7 @@ import (
"opencatd-open/internal/api" "opencatd-open/internal/api"
"opencatd-open/internal/channel" "opencatd-open/internal/channel"
"opencatd-open/internal/dao" "opencatd-open/internal/dao"
"opencatd-open/internal/passkey"
"opencatd-open/internal/proxy" "opencatd-open/internal/proxy"
"opencatd-open/internal/usage" "opencatd-open/internal/usage"
"opencatd-open/middleware" "opencatd-open/middleware"
@@ -21,6 +22,7 @@ import (
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/redis/go-redis/v9"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -50,7 +52,7 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
// Initialize health checker and start periodic checks // Initialize health checker and start periodic checks
healthChecker := channel.NewHealthChecker(channelDAO, channelSvc) healthChecker := channel.NewHealthChecker(channelDAO, channelSvc)
go healthChecker.StartPeriodicCheck(ctx, 5*time.Minute) go healthChecker.StartPeriodicCheck(ctx)
// Initialize usage recorder and start background worker // Initialize usage recorder and start background worker
usageRecorder := usage.NewRecorder(usageDAO, dailyDAO) usageRecorder := usage.NewRecorder(usageDAO, dailyDAO)
@@ -60,9 +62,29 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
// Initialize gateway // Initialize gateway
gateway := proxy.NewGateway(ctx, cfg, db, &wg, userDAO, apiKeyDAO, usageDAO, dailyDAO) gateway := proxy.NewGateway(ctx, cfg, db, &wg, userDAO, apiKeyDAO, usageDAO, dailyDAO)
gateway.SetChannelService(channelSvc) gateway.SetChannelService(channelSvc)
gateway.SetUsageRecorder(usageRecorder)
// Initialize passkey service
var rdb *redis.Client
if cfg.RedisHost != "" {
rdb = redis.NewClient(&redis.Options{
Addr: fmt.Sprintf("%s:%d", cfg.RedisHost, cfg.RedisPort),
Password: cfg.RedisPassword,
DB: cfg.RedisDB,
})
}
passkeySvc, err := passkey.New(db, passkey.Config{
RPID: cfg.RPID,
Origins: cfg.RPOrigins,
Name: cfg.AppName,
Redis: rdb,
})
if err != nil {
log.Fatalf("Failed to initialize passkey service: %v", err)
}
// Initialize API handler // Initialize API handler
apiHandler := api.NewHandler(db) apiHandler := api.NewHandler(db, passkeySvc)
r := gin.Default() r := gin.Default()
r.Use(middleware.CORS()) r.Use(middleware.CORS())
@@ -72,6 +94,8 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
{ {
public.POST("/register", apiHandler.Register) public.POST("/register", apiHandler.Register)
public.POST("/login", apiHandler.Login) public.POST("/login", apiHandler.Login)
public.POST("/passkey/begin", apiHandler.PasskeyLoginBegin)
public.POST("/passkey/finish", apiHandler.PasskeyLoginComplete)
} }
// API routes (authenticated) // API routes (authenticated)
@@ -83,6 +107,12 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
apiGroup.POST("/profile/update", apiHandler.UpdateProfile) apiGroup.POST("/profile/update", apiHandler.UpdateProfile)
apiGroup.POST("/profile/update/password", apiHandler.UpdatePassword) apiGroup.POST("/profile/update/password", apiHandler.UpdatePassword)
// Passkey management
apiGroup.POST("/webauthn/register/begin", apiHandler.PasskeyRegisterBegin)
apiGroup.POST("/webauthn/register/complete", apiHandler.PasskeyRegisterComplete)
apiGroup.GET("/webauthn/passkeys", apiHandler.PasskeyList)
apiGroup.DELETE("/webauthn/passkeys/:id", apiHandler.PasskeyDelete)
// User management (admin) // User management (admin)
apiGroup.GET("/users", apiHandler.ListUsers) apiGroup.GET("/users", apiHandler.ListUsers)
apiGroup.GET("/users/:id", apiHandler.GetUser) apiGroup.GET("/users/:id", apiHandler.GetUser)
@@ -99,7 +129,7 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
apiGroup.DELETE("/keys/:id", apiHandler.DeleteApiKey) apiGroup.DELETE("/keys/:id", apiHandler.DeleteApiKey)
apiGroup.POST("/keys/batch/:option", apiHandler.BatchApiKeys) apiGroup.POST("/keys/batch/:option", apiHandler.BatchApiKeys)
// Channel management // Channel management (legacy endpoints)
apiGroup.GET("/channels", apiHandler.ListChannels) apiGroup.GET("/channels", apiHandler.ListChannels)
apiGroup.POST("/channels", apiHandler.CreateChannel) apiGroup.POST("/channels", apiHandler.CreateChannel)
apiGroup.PUT("/channels/:id", apiHandler.UpdateChannel) apiGroup.PUT("/channels/:id", apiHandler.UpdateChannel)
@@ -107,11 +137,53 @@ func SetRouter(cfg *config.Config, db *gorm.DB, web *embed.FS) {
apiGroup.GET("/channels/:id/models", apiHandler.GetChannelModels) apiGroup.GET("/channels/:id/models", apiHandler.GetChannelModels)
apiGroup.POST("/channels/:id/models", apiHandler.BindChannelModels) apiGroup.POST("/channels/:id/models", apiHandler.BindChannelModels)
// Model management // Model management (legacy endpoints)
apiGroup.GET("/models", apiHandler.ListModels) apiGroup.GET("/models", apiHandler.ListModels)
apiGroup.POST("/models", apiHandler.CreateModel) apiGroup.POST("/models", apiHandler.CreateModel)
apiGroup.PUT("/models/:id", apiHandler.UpdateModel) apiGroup.PUT("/models/:id", apiHandler.UpdateModel)
apiGroup.DELETE("/models/:id", apiHandler.DeleteModel) apiGroup.DELETE("/models/:id", apiHandler.DeleteModel)
// 用户自身用量统计
apiGroup.GET("/usage/stats", apiHandler.MyUsageStats)
apiGroup.GET("/usage/monthly", apiHandler.MyUsageMonthly)
apiGroup.GET("/usage/logs", apiHandler.MyUsageLogs)
}
// Admin API (requires admin role)
adminGroup := r.Group("/api/admin", middleware.Auth(db), middleware.AdminOnly())
{
// Admin channel management (enhanced)
adminGroup.GET("/channels", apiHandler.AdminChannels)
adminGroup.POST("/channels", apiHandler.AdminCreateChannel)
adminGroup.PUT("/channels/:id", apiHandler.AdminUpdateChannel)
adminGroup.DELETE("/channels/:id", apiHandler.AdminDeleteChannel)
adminGroup.POST("/channels/:id/test", apiHandler.AdminTestChannel)
adminGroup.GET("/channels/:id/models/remote", apiHandler.AdminChannelRemoteModels)
adminGroup.GET("/channels/:id/models", apiHandler.AdminChannelModels)
adminGroup.POST("/channels/:id/models", apiHandler.AdminChannelAddModel)
adminGroup.PATCH("/channels/:id/models/:bid", apiHandler.AdminChannelUpdateModel)
adminGroup.DELETE("/channels/:id/models/:bid", apiHandler.AdminChannelDeleteModel)
// Admin model management (enhanced)
adminGroup.GET("/models", apiHandler.AdminModels)
adminGroup.DELETE("/models/unused", apiHandler.AdminDeleteUnusedModels)
adminGroup.POST("/models", apiHandler.AdminCreateModel)
adminGroup.PUT("/models/:id", apiHandler.AdminUpdateModel)
adminGroup.DELETE("/models/:id", apiHandler.AdminDeleteModel)
adminGroup.POST("/models/:id/bindings", apiHandler.AdminCreateModelBinding)
adminGroup.DELETE("/models/:id/bindings/:bid", apiHandler.AdminDeleteModelBinding)
// Admin system config
adminGroup.GET("/config", apiHandler.AdminGetConfig)
adminGroup.PUT("/config", apiHandler.AdminUpdateConfig)
adminGroup.GET("/config/registration", apiHandler.AdminGetRegistration)
adminGroup.PUT("/config/registration", apiHandler.AdminUpdateRegistration)
adminGroup.GET("/config/password-login", apiHandler.AdminGetPasswordLogin)
adminGroup.PUT("/config/password-login", apiHandler.AdminUpdatePasswordLogin)
// Admin usage
adminGroup.GET("/usage/logs", apiHandler.AdminUsageLogs)
adminGroup.GET("/usage/summary", apiHandler.AdminUsageSummary)
} }
// LLM proxy routes // LLM proxy routes
+2 -2
View File
@@ -18,12 +18,12 @@ ARG TARGETARCH
RUN apk --no-cache add make upx RUN apk --no-cache add make upx
WORKDIR /build WORKDIR /build
COPY . . COPY . .
COPY --from=frontend /frontend-build/dist /build/cmd/openteam/dist COPY --from=frontend /frontend-build/dist /build/backend/cmd/openteam/dist
ENV GO111MODULE=on \ ENV GO111MODULE=on \
CGO_ENABLED=0 \ CGO_ENABLED=0 \
GOOS=$TARGETOS \ GOOS=$TARGETOS \
GOARCH=$TARGETARCH GOARCH=$TARGETARCH
RUN make build RUN make build-backend
FROM alpine:latest AS runner FROM alpine:latest AS runner
# 设置alpine 时间为上海时间 # 设置alpine 时间为上海时间
+2 -2
View File
@@ -21,13 +21,13 @@ RUN sed -i 's/dl-cdn.alpinelinux.org/mirrors.aliyun.com/g' /etc/apk/repositories
&& apk --no-cache add make upx && apk --no-cache add make upx
WORKDIR /build WORKDIR /build
COPY . . COPY . .
COPY --from=frontend /frontend-build/dist /build/cmd/openteam/dist COPY --from=frontend /frontend-build/dist /build/backend/cmd/openteam/dist
ENV GO111MODULE=on \ ENV GO111MODULE=on \
GOPROXY=https://goproxy.cn,direct \ GOPROXY=https://goproxy.cn,direct \
CGO_ENABLED=0 \ CGO_ENABLED=0 \
GOOS=$TARGETOS \ GOOS=$TARGETOS \
GOARCH=$TARGETARCH GOARCH=$TARGETARCH
RUN make build RUN make build-backend
FROM alpine:latest AS runner FROM alpine:latest AS runner
# 设置alpine 时间为上海时间 # 设置alpine 时间为上海时间
+276
View File
@@ -0,0 +1,276 @@
# 网关调用流程示意图
> 对应实现:`backend/router/setRouter.go`(路由注册)、`backend/middleware/auth_llm.go`(鉴权)、
> `backend/internal/proxy/{gateway.go,handlers.go,convert/*}`(网关与三协议互转)、
> `backend/internal/channel/{channel.go,health.go}`(渠道路由与健康)、
> `backend/internal/usage/recorder.go`(用量异步落库)。
## 0. 总览
```
┌────────────────────────────────────────────────┐
│ Gin Router (/v1) │
│ ┌──────────────────────────────────────────┐ │
客户端 ───────────▶│ │ middleware.AuthLLM (密钥鉴权, 401拦截) │ │
Bearer sk-ot-… │ └──────────────────────────────────────────┘ │
│ ┌───────┬────────┬─────────┬─────────┐ │
│ │ chat │messages│responses│ models │ │
│ │Handle │Handle │ Handle │ Handle │ │
│ │Chat │Messages│Responses│ Models │ │
│ └───┬───┴───┬────┴────┬────┴─────────┘ │
│ └───────┴────┬────┘ │
│ ParseRequest │
│ (model/stream) │
│ │ │
│ Dispatch ◀── Candidates/Pick │
│ (转换+故障转移+用量记录) │
└───────────────────┼──────────────────────────────┘
│
▼
上游 /v1/* (chat|messages|responses)
```
## 1. 请求入口与鉴权
```mermaid
sequenceDiagram
autonumber
participant C as 客户端
participant R as Gin /v1 路由
participant A as AuthLLM
participant DB as SQLite(APIKey)
participant H as HandleXxx
C->>R: POST /v1/chat/completions 等
Note over R: /v1 组挂 middleware.AuthLLM
R->>A: 进入中间件
A->>A: 提取 Bearer token(兼容无前缀直传)
alt 未携带 Authorization
A-->>C: 401 「未提供认证信息」
else token 长度 < 12 或 prefix 不匹配
A-->>C: 401 「无效的API密钥」
else prefix 命中
A->>DB: SELECT * WHERE key_prefix=? AND status=active
A->>A: sha256(token) == KeyHash ?
alt 哈希不一致
A-->>C: 401 「无效的API密钥」
else 校验通过
A->>H: c.Set(api_key, user_id) → 放行
end
end
```
> 关键点:`key_prefix` 取 `sk-ot-` 后前 12 位(`api.go` 创建密钥时 `keyValue[:12]`),
> `auth_llm.go` 用同一常量 `keyPrefixLen=12` 切片,避免越界 panic。
## 2. 三协议主调用流程
```mermaid
sequenceDiagram
autonumber
participant C as 客户端
participant H as HandleChat/Messages/Responses
participant P as ParseRequest
participant G as Dispatch
participant S as ChannelService
participant U as usage.Recorder
participant UP as 上游(OpenRouter等)
C->>H: 请求体 (model, stream, messages/input…)
H->>P: ParseRequest(protocol)
P->>P: 读 body → 解析 model / stream
P-->>H: Request{Model, Stream, Protocol, Body, …}
H->>G: Dispatch(req)
G->>S: Candidates(req.Model)
S-->>G: []Candidate{Channel, Binding?}
G->>S: FilterHealthy(cands)
G->>S: Pick(cands) → 加权随机选定一个候选
Note over G,S: 绑定优先(携带 upstream_model 映射);<br/>无绑定回退到权重最低的健康备用渠道
loop 故障转移(候选耗尽前)
G->>G: conversionTarget(ch, proto) → 渠道首选协议
alt 客户端协议 ≠ 渠道协议
G->>G: ConvertRequest(body, from, to) 转换请求体
Note over G: 走 convert 包(chat/messages/responses 互转)
end
alt 有 Binding.UpstreamModel
G->>G: rewriteModel(body, upstreamModel) 别名映射
end
G->>UP: POST {base}/v1/{path} (按渠道协议拼 URL/头)
alt 连接失败 或 429/5xx
S->>S: RecordFailure(ch) → 连续2次 degraded 熔断
G->>U: Record(error 事件, error_code)
Note over G: continue → 换下一个候选渠道
else 4xx
G-->>C: 透传上游错误体 (不重试)
else 2xx
S->>S: RecordSuccess(ch)
alt stream=true
G->>G: streamResponse → 逐块转发 + 累计 usage
else
G->>G: bufferResponse → 整体转发 + 提取 usage
end
G->>G: recordUsage (按模型定价计算 cost)
G->>U: Record(成功事件, tokens, cost)
G-->>C: 响应
end
end
Note over G,C: 全部候选失败 → 502/503
```
> **故障转移规则**(对齐参考实现 `doProxy`):
> - 连接错误、429、5xx → 可重试,换下一个候选;
> - 4xx(如 400 参数错误)→ 透传上游错误体,不重试;
> - 无可用渠道(全部不健康/无绑定且无备用)→ 502/503 + `error_code=no_channel`。
## 3. 路由选择细节
```mermaid
flowchart TD
A[客户端 model 名] --> B{存在启用模型行?}
B -- 是 --> C{有绑定且渠道健康?}
C -- 是 --> D[候选 = 绑定该模型的渠道<br/>排序 priority ASC, weight DESC, id ASC]
C -- 否 --> E
B -- 否 --> E[候选 = 全部启用渠道<br/>取权重最低的健康备用渠道]
D --> F[FilterHealthy 内存熔断过滤]
E --> F
F --> G[Pick 加权随机选中一个]
G --> H[Dispatch 开始尝试]
H --> I{尝试成功?}
I -- 失败可重试 --> J[RecordFailure + 换下一个]
J --> H
I -- 成功 --> K[RecordSuccess + 响应 + 记账]
J -. 全部耗尽 .-> L[502/503]
```
## 4. 跨协议转换(client ↔ 渠道原生协议)
```mermaid
flowchart LR
subgraph 客户端协议
CHAT[/"chat<br/>chat/completions"/]
MSG[/"messages<br/>(Anthropic)"/]
RESP[/"responses<br/>(OpenAI)"/]
end
subgraph 中间模型
MID["Chat 形状<br/>(标准中间模型)"]
end
subgraph 渠道协议
UCHAT[/"chat"/]
UMSG[/"messages"/]
URESP[/"responses"/]
end
CHAT -->|直通| UCHAT
MSG -->|messagesToChat| MID -->|chatToMessages| UMSG
MSG -->|messagesToChat| MID -->|chatToResponses| URESP
RESP -->|responsesToChat| MID -->|chatToResponses| URESP
RESP -->|responsesToChat| MID -->|chatToMessages| UMSG
```
> 转换入口:`convert.ConvertRequest`(请求体)、`convert.ConvertResponse`(非流式响应)、
> `convert.NewStreamTransformer`(流式 SSE 逐行转换)。跨两跳时经 Chat 中转(如
> `responses→messages` = `responsesToChatReq` + `chatToMessagesReq`)。
## 5. 流式 / 非流式响应与用量提取
```mermaid
sequenceDiagram
autonumber
participant G as Dispatch
participant S as streamResponse
participant B as bufferResponse
participant ACC as StreamUsageAccum
participant W as 客户端 Writer
participant UP as 上游
alt stream=true
G->>S: streamResponse(resp, clientProto, upstreamProto)
S->>UP: 按 \n\n 读块 (bufio)
loop 每个 SSE 块
S->>ACC: sseDataPayloads(chunk) → Feed(data, upstreamProto)
Note over ACC: 逐协议累计 usage 字段
alt 跨协议
S->>S: NewStreamTransformer(upstream→client).line(chunk)
end
S->>W: 写块 + Flush
alt 遇到流结束标记
Note over S: chat: data:[DONE]<br/>messages: message_stop<br/>responses: response.completed
S-->>G: 返回累计 TokenUsage
end
end
else stream=false
G->>B: bufferResponse(resp, clientProto, upstreamProto)
B->>B: io.ReadAll
B->>B: ExtractUsageJSON(body, upstreamProto)
alt 跨协议
B->>B: ConvertResponse(body, upstream→client)
else 直通
B->>B: CleanJSON(body) 去空白/SSE注释前缀
end
B-->>G: 返回 TokenUsage
end
G->>G: recordUsage(req, cand, ch, ev, tok)
Note over G: 定价 cost = (非缓存输入×输入价 + 缓存读×缓存价<br/>+ 缓存写×输出价 + 输出×输出价) / 1e6
G->>G: usageRec.Record(Event)
```
## 6. 用量异步落库
```mermaid
sequenceDiagram
autonumber
participant G as Gateway
participant R as usage.Recorder
participant U as UsageDAO
participant D as DailyUsageDAO
participant DB as SQLite
G->>R: Record(Event) 每次请求(成功/错误/取消)
Note over R: 缓冲 channel (10000), 每 5s 或满 100 条 flush
R->>U: BatchCreate(UsageLog[])
R->>D: UpsertDailyUsage(UsageDaily) 按(user_id,model_id,date)增量累加
U->>DB: INSERT usage_logs
D->>DB: ON CONFLICT 累加 requests/input/output/cache/cost
```
> `usage_dailies` 用 `gorm.Expr("requests + ?")` 增量累加而非覆盖,保证多次 flush 不互相清零。
## 7. 渠道健康与熔断
```mermaid
flowchart TD
A[请求失败] --> B[RecordFailure: consecutive++]
B --> C{consecutive >= 2?}
C -- 是 --> D[status=degraded + 5min cooldown]
C -- 否 --> E[仅计数]
D --> F{后续请求 Candidates}
F --> G{FilterHealthy 该渠道}
G -- degraded/cooldown 未过期 --> H[排除, 走其他渠道/备用]
G -- healthy --> I[参与选择]
D -. 冷却过期 .-> J[复位 healthy]
J --> I
K[健康检查周期探测成功] --> L[RecordSuccess: 复位 healthy]
```
> 说明:`health.go` 的 `StartPeriodicCheck`(默认 5min)会探测各渠道 `/models`,
> 成功调 `RecordSuccess` 复位;失败调 `RecordFailure` 进入熔断计数。
## 8. 关键代码锚点
| 环节 | 位置 |
|---|---|
| /v1 路由注册 + AuthLLM | `router/setRouter.go:146` |
| 密钥鉴权 | `middleware/auth_llm.go` |
| 请求解析 | `proxy/gateway.go:104 ParseRequest` |
| 主调度 + 故障转移 | `proxy/gateway.go:157 Dispatch` |
| 流式转发 + 结束检测 | `proxy/gateway.go:388 streamResponse` |
| 非流式转发 | `proxy/gateway.go:496 bufferResponse` |
| 用量记账 | `proxy/gateway.go:302 recordUsage` |
| 候选构建 | `channel/channel.go:51 Candidates` |
| 内存健康过滤 | `channel/channel.go:252 FilterHealthy` |
| 加权选择 | `channel/channel.go:222 Pick` |
| 失败熔断 | `channel/channel.go:139 RecordFailure` |
| 三协议互转 | `proxy/convert/{convert.go,json_chat.go,json_responses.go,stream_transform.go}` |
| 用量异步落库 | `usage/recorder.go:114 flush` |
+7 -1
View File
@@ -2,6 +2,7 @@
import axios from 'axios' import axios from 'axios'
import type { AxiosError, InternalAxiosRequestConfig } from 'axios' import type { AxiosError, InternalAxiosRequestConfig } from 'axios'
import { useAuthStore } from '@/stores/auth' import { useAuthStore } from '@/stores/auth'
import router from '@/router'
const baseURL = import.meta.env.VITE_API_BASE_URL || '/api' const baseURL = import.meta.env.VITE_API_BASE_URL || '/api'
if (import.meta.env.DEV) { // Vite 的方式判断开发环境 if (import.meta.env.DEV) { // Vite 的方式判断开发环境
@@ -49,7 +50,12 @@ service.interceptors.response.use(
if (error.response && error.response.status === 401) { if (error.response && error.response.status === 401) {
const authStore = useAuthStore(); const authStore = useAuthStore();
authStore.clear(); authStore.clear();
window.location.href = '/login'; // 守卫校验期间(尚未进入受保护路由)由守卫负责跳登录;
// 这里只处理已登录状态下 token 失效的情况,且不再用 location.href 硬刷新
const current = router.currentRoute.value;
if (current.matched.some(record => record.meta.requiresAuth)) {
router.push({ path: '/login', query: { redirect: current.fullPath } });
}
} }
return Promise.reject(error); return Promise.reject(error);
} }
@@ -141,8 +141,8 @@ const rightIcons = [
{ id: 'zhipu', label: 'Zhipu', img: LobeIcon('zhipu'), color: '#4268fa' }, { id: 'zhipu', label: 'Zhipu', img: LobeIcon('zhipu'), color: '#4268fa' },
{ id: 'qwen', label: 'Qwen', img: LobeIcon('qwen'), color: '#615ced' }, { id: 'qwen', label: 'Qwen', img: LobeIcon('qwen'), color: '#615ced' },
{ id: 'deepseek', label: 'DeepSeek', img: LobeIcon('deepseek'), color: '#4d6bfe' }, { id: 'deepseek', label: 'DeepSeek', img: LobeIcon('deepseek'), color: '#4d6bfe' },
{ id: 'moonshot', label: 'Moonshot', img: LobeIcon('moonshot'), color: '#000' }, { id: 'moonshot', label: 'Moonshot', img: LobeIcon('moonshot'), color: '#666' },
{ id: 'minimax', label: 'MiniMax', img: LobeIcon('minimax'), color: '#000' }, { id: 'minimax', label: 'MiniMax', img: LobeIcon('minimax'), color: '#F23F5D' },
{ id: 'bedrock', label: 'Bedrock', img: LobeIcon('bedrock'), color: '#ff9900' }, { id: 'bedrock', label: 'Bedrock', img: LobeIcon('bedrock'), color: '#ff9900' },
{ id: 'azure', label: 'Azure', img: LobeIcon('azure'), color: '#0078d4' }, { id: 'azure', label: 'Azure', img: LobeIcon('azure'), color: '#0078d4' },
{ id: 'volcengine', label: 'Volcengine', img: LobeIcon('volcengine'), color: '#325ab4' }, { id: 'volcengine', label: 'Volcengine', img: LobeIcon('volcengine'), color: '#325ab4' },
+2 -2
View File
@@ -47,9 +47,9 @@ const iconForType = (type: ToastType) => {
const typeClasses = (type: ToastType) => { const typeClasses = (type: ToastType) => {
switch (type) { switch (type) {
case 'success': case 'success':
return 'border-success/20 bg-success/10 text-success-content dark:border-success/30 dark:bg-success/15'; return 'border-success bg-success/15 text-success';
case 'error': case 'error':
return 'border-error/20 bg-error/10 text-error-content dark:border-error/30 dark:bg-error/15'; return 'border-error bg-error/15 text-error';
default: default:
return 'border-base-300 bg-base-100 text-base-content'; return 'border-base-300 bg-base-100 text-base-content';
} }
+30
View File
@@ -0,0 +1,30 @@
<script setup lang="ts">
withDefaults(defineProps<{ variant?: 'neutral' | 'ok' | 'warn' | 'err' | 'accent' }>(), {
variant: 'neutral',
})
</script>
<template>
<span
class="inline-flex items-center gap-1.5 rounded-full px-2 py-0.5 text-[11px] leading-5"
:class="{
neutral: 'bg-base-200 text-base-content',
ok: 'bg-success/10 text-success',
warn: 'bg-warning/10 text-warning',
err: 'bg-error/10 text-error',
accent: 'bg-primary/10 text-primary',
}[variant]"
>
<span
v-if="variant !== 'neutral'"
class="size-1.5 rounded-full"
:class="{
ok: 'bg-success',
warn: 'bg-warning',
err: 'bg-error',
accent: 'bg-primary',
}[variant]"
/>
<slot />
</span>
</template>
+27
View File
@@ -0,0 +1,27 @@
<script setup lang="ts">
withDefaults(
defineProps<{
variant?: 'primary' | 'ghost' | 'danger'
size?: 'sm' | 'md'
loading?: boolean
disabled?: boolean
}>(),
{ variant: 'primary', size: 'md', loading: false, disabled: false },
)
</script>
<template>
<button
:disabled="disabled || loading"
class="inline-flex items-center justify-center gap-2 rounded-md font-medium transition-[transform,background-color,border-color,color] duration-150 active:scale-[0.98] disabled:pointer-events-none disabled:opacity-50 select-none"
:class="[
size === 'sm' ? 'h-8 px-3 text-xs' : 'h-10 px-4 text-sm',
variant === 'primary' && 'bg-primary text-primary-content hover:bg-primary/90',
variant === 'ghost' && 'border border-base-300/60 text-base-content hover:bg-base-200/50',
variant === 'danger' && 'border border-error text-error hover:bg-error/10',
]"
>
<span v-if="loading" class="size-3.5 animate-spin rounded-full border-2 border-current border-t-transparent" />
<slot />
</button>
</template>
+36
View File
@@ -0,0 +1,36 @@
<script setup lang="ts">
withDefaults(
defineProps<{
label?: string
modelValue?: string | number
type?: string
placeholder?: string
hint?: string
error?: string
autocomplete?: string
disabled?: boolean
maxlength?: number
}>(),
{ type: 'text', modelValue: '', disabled: false },
)
const emit = defineEmits<{ 'update:modelValue': [string | number] }>()
</script>
<template>
<label class="block">
<span v-if="label" class="mb-1.5 block text-xs font-medium text-base-content/50">{{ label }}</span>
<input
:type="type"
:value="modelValue"
:placeholder="placeholder"
:autocomplete="autocomplete"
:disabled="disabled"
:maxlength="maxlength"
class="h-10 w-full rounded-md border border-base-300/60 bg-base-100 px-3 text-sm text-base-content placeholder-base-content/40 outline-none transition focus:border-primary focus:ring-2 focus:ring-primary disabled:cursor-not-allowed disabled:opacity-50"
:class="error && 'border-error focus:border-error focus:ring-error'"
@input="emit('update:modelValue', ($event.target as HTMLInputElement).value as string | number)"
/>
<span v-if="hint && !error" class="mt-1.5 block text-xs text-base-content/50">{{ hint }}</span>
<span v-if="error" class="mt-1.5 block text-xs text-error">{{ error }}</span>
</label>
</template>
+84
View File
@@ -0,0 +1,84 @@
<script setup lang="ts">
import { nextTick, onMounted, onUnmounted, ref, watch } from 'vue'
import { X } from '@lucide/vue'
const props = withDefaults(
defineProps<{
open: boolean
title?: string
width?: string
}>(),
{ width: 'max-w-md' },
)
const emit = defineEmits<{ close: [] }>()
const panel = ref<HTMLElement | null>(null)
function onKey(e: KeyboardEvent) {
if (e.key === 'Escape' && props.open) emit('close')
}
onMounted(() => window.addEventListener('keydown', onKey))
onUnmounted(() => window.removeEventListener('keydown', onKey))
watch(
() => props.open,
async (open) => {
document.body.style.overflow = open ? 'hidden' : ''
if (open) {
await nextTick()
panel.value?.focus()
}
},
)
onUnmounted(() => {
document.body.style.overflow = ''
})
</script>
<template>
<Teleport to="body">
<Transition
enter-active-class="transition-opacity duration-150"
enter-from-class="opacity-0"
leave-active-class="transition-opacity duration-150"
leave-to-class="opacity-0"
>
<div
v-if="open"
class="fixed inset-0 z-50 flex items-start justify-center overflow-y-auto bg-black/60 p-4 pt-[12vh] backdrop-blur-sm"
@mousedown.self="emit('close')"
>
<Transition
enter-active-class="transition-transform duration-150"
enter-from-class="scale-[0.97] opacity-0"
leave-active-class="transition-transform duration-150"
leave-to-class="scale-[0.97] opacity-0"
>
<div
v-if="open"
ref="panel"
role="dialog"
aria-modal="true"
:aria-label="title || '对话框'"
tabindex="-1"
class="card w-full bg-base-100 shadow-xl outline-none"
:class="width"
>
<div class="flex items-center justify-between border-b border-base-300/60 px-5 py-3.5">
<h3 class="text-sm font-semibold text-base-content">{{ title }}</h3>
<button class="rounded-md p-1 text-base-content/40 hover:bg-base-200 hover:text-base-content" aria-label="关闭" @click="emit('close')">
<X :size="16" />
</button>
</div>
<div class="px-5 py-4">
<slot />
</div>
<div v-if="$slots.footer" class="flex justify-end gap-2 border-t border-base-300/60 px-5 py-3.5">
<slot name="footer" />
</div>
</div>
</Transition>
</div>
</Transition>
</Teleport>
</template>
+5 -1
View File
@@ -1,6 +1,10 @@
<!-- src/layouts/DashboardLayout.vue --> <!-- src/layouts/DashboardLayout.vue -->
<template> <template>
<div class="min-h-screen bg-base-200"> <!-- 用户信息就绪前不渲染后台内容,避免未授权内容闪现 -->
<div v-if="!authStore.user" class="flex min-h-screen items-center justify-center bg-base-200">
<span class="loading loading-spinner loading-lg text-base-content/30"></span>
</div>
<div v-else class="min-h-screen bg-base-200">
<div class="drawer" :class="{ 'lg:drawer-open': isLargeSidebarOpen }"> <div class="drawer" :class="{ 'lg:drawer-open': isLargeSidebarOpen }">
<input id="ot-drawer" type="checkbox" class="drawer-toggle" /> <input id="ot-drawer" type="checkbox" class="drawer-toggle" />
+27
View File
@@ -0,0 +1,27 @@
// 协议格式显示名与选项
export const PROTOCOL_NAMES: Record<string, string> = {
chat: 'OpenAI Chat Completions',
responses: 'OpenAI Responses API',
messages: 'Anthropic Messages',
}
// 渠道表格用的短标识
export const PROTOCOL_SHORT: Record<string, string> = {
chat: 'chat/completions',
responses: 'responses',
messages: 'messages',
}
export function protocolShort(p: string): string {
return PROTOCOL_SHORT[p] ?? p
}
export const PROTOCOL_OPTIONS: { value: string; label: string }[] = [
{ value: 'chat', label: 'OpenAI Chat Completions' },
{ value: 'responses', label: 'OpenAI Responses API' },
{ value: 'messages', label: 'Anthropic Messages' },
]
export function protocolName(p: string): string {
return PROTOCOL_NAMES[p] ?? p
}
+1 -1
View File
@@ -10,6 +10,6 @@ const pinia = createPinia()
const app = createApp(App) const app = createApp(App)
app.provide('request', request) app.provide('request', request)
app.use(pinia) // 必须先于 router:路由守卫里会用到 auth store
app.use(router) app.use(router)
app.use(pinia)
app.mount('#app') app.mount('#app')
+28 -6
View File
@@ -1,19 +1,41 @@
import { createRouter, createWebHistory } from 'vue-router' import { createRouter, createWebHistory } from 'vue-router'
import { routes } from '@/utils/router_menu' import { routes } from '@/utils/router_menu'
import { useAuthStore } from '@/stores/auth'
const router = createRouter({ const router = createRouter({
history: createWebHistory(), history: createWebHistory(),
routes, routes,
}) })
router.beforeEach((to, from, next) => { // 受保护页面必须先通过服务端校验才渲染:
const isAuthenticated = localStorage.getItem('token') // 本地 token 存在不代表有效(可能已过期/被重置),若只查 localStorage,
// 页面会先渲染约 1 秒、等 /profile 返回 401 后才被踢回登录页。
router.beforeEach(async (to) => {
const requiresAuth = to.matched.some(record => record.meta.requiresAuth) const requiresAuth = to.matched.some(record => record.meta.requiresAuth)
if (requiresAuth && !isAuthenticated) { if (!requiresAuth) return true
next('/login')
} else { const authStore = useAuthStore()
next()
if (!authStore.token) {
return { path: '/login', query: { redirect: to.fullPath } }
} }
// 有 token 但还没加载用户信息时,先向服务端确认身份,失败则不得进入
if (!authStore.user) {
try {
await authStore.getProfile()
} catch {
authStore.clear()
return { path: '/login', query: { redirect: to.fullPath } }
}
}
// 管理后台仅对 role >= 10 开放
if (to.matched.some(record => record.meta.requiresAdmin) && (authStore.user?.role ?? 0) < 10) {
return '/dashboard/overview'
}
return true
}) })
export default router export default router
+101
View File
@@ -8,6 +8,8 @@ export type Channel = {
name: string name: string
provider: string provider: string
base_url: string base_url: string
base_urls?: Record<string, string>
api_key_masked?: string
weight: number weight: number
priority: number priority: number
timeout_ms: number timeout_ms: number
@@ -31,6 +33,14 @@ export type NewChannelPayload = {
formats?: string[] formats?: string[]
} }
export type ChannelModelBinding = {
id: number
model_id: number
model_name: string
upstream_model: string
weight: number
}
export const useChannelStore = defineStore('channel', () => { export const useChannelStore = defineStore('channel', () => {
const loading = ref(false); const loading = ref(false);
const error = ref<string | null>(null); const error = ref<string | null>(null);
@@ -125,6 +135,91 @@ export const useChannelStore = defineStore('channel', () => {
} }
}; };
// Admin API methods
const testChannel = async (id: number | string) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.post(`/admin/channels/${id}/test`);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to test channel';
throw err;
} finally {
loading.value = false;
}
};
const fetchRemoteModels = async (id: number | string) => {
loading.value = true;
error.value = null;
try {
const response = await request.get(`/admin/channels/${id}/models/remote`);
return response.data.data ?? [];
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to fetch remote models';
throw err;
} finally {
loading.value = false;
}
};
const fetchChannelModels = async (id: number | string) => {
loading.value = true;
error.value = null;
try {
const response = await request.get(`/admin/channels/${id}/models`);
return response.data.data ?? [];
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to fetch channel models';
throw err;
} finally {
loading.value = false;
}
};
const addChannelModel = async (id: number | string, data: { model_id: number; upstream_model: string; weight?: number }) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.post(`/admin/channels/${id}/models`, data);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to add model';
throw err;
} finally {
loading.value = false;
}
};
const updateChannelModel = async (channelId: number | string, bindingId: number | string, data: { upstream_model?: string; weight?: number }) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.patch(`/admin/channels/${channelId}/models/${bindingId}`, data);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to update binding';
throw err;
} finally {
loading.value = false;
}
};
const deleteChannelModel = async (channelId: number | string, bindingId: number | string) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.delete(`/admin/channels/${channelId}/models/${bindingId}`);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to delete binding';
throw err;
} finally {
loading.value = false;
}
};
return { return {
loading, error, loading, error,
channel, channels, totalChannels, channel, channels, totalChannels,
@@ -134,5 +229,11 @@ export const useChannelStore = defineStore('channel', () => {
updateChannel, updateChannel,
deleteChannel, deleteChannel,
batchChannels, batchChannels,
testChannel,
fetchRemoteModels,
fetchChannelModels,
addChannelModel,
updateChannelModel,
deleteChannelModel,
}; };
}); });
+171
View File
@@ -0,0 +1,171 @@
import { defineStore } from 'pinia';
import { ref } from 'vue';
import type { AxiosResponse } from 'axios';
import request from '@/api/client';
export type Model = {
id: number
name: string
display_name?: string
input_price: number
output_price: number
cache_read_price: number
enabled: boolean
sort: number
channels?: ModelBinding[]
used?: boolean
needs_pricing?: boolean
denied?: boolean
created_at?: string
updated_at?: string
[key: string]: unknown
}
export type ModelBinding = {
id: number
channel_id: number
channel_name: string
upstream_model: string
weight: number
}
export type NewModelPayload = {
name: string
display_name?: string
input_price?: number
output_price?: number
cache_read_price?: number
sort?: number
enabled?: boolean
}
export type ModelSummary = {
total: number
unpriced: number
missing: OrphanBinding[]
denied_count: number
}
export type OrphanBinding = {
channel: string
model_id: number
upstream_model: string
}
export const useModelStore = defineStore('model', () => {
const loading = ref(false);
const error = ref<string | null>(null);
const models = ref<Model[]>([]);
const summary = ref<ModelSummary | null>(null);
const fetchModels = async () => {
loading.value = true;
error.value = null;
try {
const response = await request.get('/admin/models');
models.value = response.data.data ?? [];
summary.value = response.data.summary ?? null;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to fetch models';
throw err;
} finally {
loading.value = false;
}
};
const createModel = async (data: NewModelPayload) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.post('/admin/models', data);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to create model';
throw err;
} finally {
loading.value = false;
}
};
const updateModel = async (id: number | string, data: Partial<Model>) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.put(`/admin/models/${id}`, data);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to update model';
throw err;
} finally {
loading.value = false;
}
};
const deleteModel = async (id: number | string) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.delete(`/admin/models/${id}`);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to delete model';
throw err;
} finally {
loading.value = false;
}
};
const deleteUnusedModels = async () => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.delete('/admin/models/unused');
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to delete unused models';
throw err;
} finally {
loading.value = false;
}
};
const createModelBinding = async (modelId: number | string, data: { channel_id: number; upstream_model: string; weight?: number }) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.post(`/admin/models/${modelId}/bindings`, data);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to create binding';
throw err;
} finally {
loading.value = false;
}
};
const deleteModelBinding = async (modelId: number | string, bindingId: number | string) => {
loading.value = true;
error.value = null;
try {
const response: AxiosResponse = await request.delete(`/admin/models/${modelId}/bindings/${bindingId}`);
return response;
} catch (err: any) {
error.value = err.response?.data?.error || 'Failed to delete binding';
throw err;
} finally {
loading.value = false;
}
};
return {
loading, error,
models, summary,
fetchModels,
createModel,
updateModel,
deleteModel,
deleteUnusedModels,
createModelBinding,
deleteModelBinding,
};
});
+114
View File
@@ -0,0 +1,114 @@
// src/stores/usage.ts
import { defineStore } from 'pinia'
import { ref } from 'vue'
import request from '@/api/client'
import type { UsageStatsData, UsageLogItem, AdminUsageSummary, MonthlyUsageData } from '@/types'
export const useUsageStore = defineStore('usage', () => {
const loading = ref(false)
const error = ref<string | null>(null)
// 普通用户:每日统计
const stats = ref<UsageStatsData | null>(null)
// 普通用户:年度按月统计(含按模型分解)
const monthly = ref<MonthlyUsageData | null>(null)
// 普通用户:自身明细
const myLogs = ref<UsageLogItem[]>([])
const myLogsTotal = ref(0)
// 管理后台:全量明细
const adminLogs = ref<UsageLogItem[]>([])
const adminLogsTotal = ref(0)
const adminSummary = ref<AdminUsageSummary | null>(null)
async function fetchStats(days = 30) {
loading.value = true
error.value = null
try {
const res = await request.get('/usage/stats', { params: { days } })
stats.value = res.data?.data ?? null
} catch (err: any) {
error.value = err.response?.data?.error || '获取用量统计失败'
throw err
} finally {
loading.value = false
}
}
async function fetchMonthly(year?: number) {
loading.value = true
error.value = null
try {
const res = await request.get('/usage/monthly', { params: year ? { year } : {} })
monthly.value = res.data?.data ?? null
} catch (err: any) {
error.value = err.response?.data?.error || '获取月度统计失败'
throw err
} finally {
loading.value = false
}
}
async function fetchMyLogs(pageSize = 20, page = 1) {
loading.value = true
error.value = null
try {
const res = await request.get('/usage/logs', { params: { pageSize, page } })
myLogs.value = res.data?.data ?? []
myLogsTotal.value = res.data?.total ?? 0
} catch (err: any) {
error.value = err.response?.data?.error || '获取用量明细失败'
throw err
} finally {
loading.value = false
}
}
async function fetchAdminLogs(params: Record<string, any> = {}) {
loading.value = true
error.value = null
try {
const res = await request.get('/admin/usage/logs', { params })
adminLogs.value = res.data?.data ?? []
adminLogsTotal.value = res.data?.total ?? 0
} catch (err: any) {
error.value = err.response?.data?.error || '获取用量明细失败'
throw err
} finally {
loading.value = false
}
}
async function fetchAdminSummary(params: Record<string, any> = {}) {
loading.value = true
error.value = null
try {
const res = await request.get('/admin/usage/summary', { params })
adminSummary.value = res.data?.data ?? null
} catch (err: any) {
error.value = err.response?.data?.error || '获取用量汇总失败'
throw err
} finally {
loading.value = false
}
}
return {
loading,
error,
stats,
monthly,
myLogs,
myLogsTotal,
adminLogs,
adminLogsTotal,
adminSummary,
fetchStats,
fetchMonthly,
fetchMyLogs,
fetchAdminLogs,
fetchAdminSummary,
}
})
+21 -23
View File
@@ -16,26 +16,21 @@ export const useWebAuthStore = defineStore("webauth", () => {
const loading = ref(false); const loading = ref(false);
const error = ref<string | null>(null); const error = ref<string | null>(null);
const addPasskey = async () => { const addPasskey = async (name?: string) => {
error.value = ""; error.value = "";
loading.value = true; loading.value = true;
try { try {
// 1. 从后端获取注册选项 (Creation Options) // 1. 从后端获取注册选项 (Creation Options)
const res = await request.get("/profile/passkey"); const res = await request.post("/webauthn/register/begin", {});
// console.log("begin:", res.data.data.publicKey); const { creation, challenge } = res.data.data;
const options = res.data.data.publicKey;
// 调用 Web Authentication API 进行注册 // 调用 Web Authentication API 进行注册
// const credential = await navigator.credentials.create(options);
// console.log("credential:", credential);
let attestation; let attestation;
try { try {
// Pass 'undefined' as the second argument if you are not using an AbortSignal // Pass 'undefined' as the second argument if you are not using an AbortSignal
attestation = await startRegistration({ optionsJSON: options }); attestation = await startRegistration({ optionsJSON: creation });
// console.log("WebAuthn 注册结果 (Attestation):", JSON.stringify(attestation));
error.value = null; error.value = null;
} catch (regError: any) { } catch (regError: any) {
// console.log("WebAuthn 注册失败或取消:", regError);
if (regError.name === "NotAllowedError") { if (regError.name === "NotAllowedError") {
error.value = "Passkey 操作被取消或不允许。"; error.value = "Passkey 操作被取消或不允许。";
} else { } else {
@@ -45,8 +40,11 @@ export const useWebAuthStore = defineStore("webauth", () => {
} }
// 3. 将注册结果 (Attestation) 发送到后端进行验证和保存 // 3. 将注册结果 (Attestation) 发送到后端进行验证和保存
const res2: AxiosResponse = await request.post("/profile/passkey", attestation); const res2: AxiosResponse = await request.post("/webauthn/register/complete", {
// console.log("end:", res2); challenge,
name: name || "passkey",
credential: attestation,
});
return res2; return res2;
} catch (err: any) { } catch (err: any) {
error.value = err.response?.data?.error || "添加 Passkey 失败,请稍后重试。"; error.value = err.response?.data?.error || "添加 Passkey 失败,请稍后重试。";
@@ -56,20 +54,18 @@ export const useWebAuthStore = defineStore("webauth", () => {
} }
}; };
const loginPasskey = async () => { const loginPasskey = async (username?: string) => {
error.value = null; error.value = null;
loading.value = true; loading.value = true;
try { try {
// 1. 从后端获取登录选项 (Assertion Options) // 1. 从后端获取登录选项 (Assertion Options)
const res = await request.get("/auth/passkey/begin"); const res = await request.post("/auth/passkey/begin", { username });
// console.log("login begin:", res.data); const { assertion, challenge, user_id } = res.data.data;
const options = res.data.data.publicKey;
// 2. 调用 Web Authentication API 进行认证 // 2. 调用 Web Authentication API 进行认证
let assertion; let credential;
try { try {
assertion = await startAuthentication({ optionsJSON: options }); credential = await startAuthentication({ optionsJSON: assertion });
// console.log("WebAuthn 认证结果 (Assertion):", JSON.stringify(assertion));
} catch (loginError: any) { } catch (loginError: any) {
if (loginError.name === "NotAllowedError") { if (loginError.name === "NotAllowedError") {
error.value = "Passkey 登录被取消或不允许。"; error.value = "Passkey 登录被取消或不允许。";
@@ -80,8 +76,11 @@ export const useWebAuthStore = defineStore("webauth", () => {
} }
// 3. 将认证结果 (Assertion) 发送到后端进行验证并获取 Token // 3. 将认证结果 (Assertion) 发送到后端进行验证并获取 Token
const challenge = options.challenge; // 从 begin 接口返回的 options 中获取 challenge const res2: AxiosResponse = await request.post("/auth/passkey/finish", {
const res2: AxiosResponse = await request.post(`/auth/passkey/finish?challenge=${challenge}`, assertion); challenge,
credential,
user_id,
});
// 4. 处理登录成功的响应,通常包含 Token // 4. 处理登录成功的响应,通常包含 Token
if (res2.status === 200 && !!res2.data.data?.token) { if (res2.status === 200 && !!res2.data.data?.token) {
@@ -103,8 +102,7 @@ export const useWebAuthStore = defineStore("webauth", () => {
loading.value = true; loading.value = true;
error.value = null; error.value = null;
try { try {
const response = await request.get('/profile/passkeys') const response = await request.get('/webauthn/passkeys')
// console.log('getPasskeys',response.data.data)
passkeys.value = response.data.data passkeys.value = response.data.data
} catch (err: any) { } catch (err: any) {
error.value = err.response?.data?.error || '获取token列表失败'; error.value = err.response?.data?.error || '获取token列表失败';
@@ -118,7 +116,7 @@ export const useWebAuthStore = defineStore("webauth", () => {
loading.value = true; loading.value = true;
error.value = null; error.value = null;
try { try {
const response: AxiosResponse = await request.delete(`/profile/passkeys/${id}`) const response: AxiosResponse = await request.delete(`/webauthn/passkeys/${id}`)
return response return response
} catch (err: any) { } catch (err: any) {
error.value = err.response?.data?.error || `删除passkey ${id} 失败`; error.value = err.response?.data?.error || `删除passkey ${id} 失败`;
+138
View File
@@ -115,3 +115,141 @@ export type NewUserPayload = {
unlimited_quota?: boolean unlimited_quota?: boolean
language?: string language?: string
} }
// Channel 渠道管理
export interface Channel {
id: number
name: string
provider: 'openai' | 'anthropic' | 'compatible'
formats: string[] // chat | responses | messages
base_url: string
base_urls?: Record<string, string> | null
api_key_masked: string
weight: number
priority: number
timeout_ms: number
max_concurrency: number
health_status: string
enabled: boolean
created_at: string
}
export interface ChannelModelMapping {
id: number
model_id: number
model_name: string
upstream_model: string
weight: number
}
// Model 模型定价
export interface ModelBinding {
id: number
channel_id: number
channel_name: string
upstream_model: string
weight: number
}
export interface Model {
id: number
name: string
input_price: number
output_price: number
cache_read_price: number
enabled: boolean
sort: number
channels: ModelBinding[]
used?: boolean
needs_pricing?: boolean
denied?: boolean
}
export interface ModelSummary {
total: number
unpriced: number
missing: { channel: string; model_id: number; upstream_model: string }[]
denied_count: number
}
// ---- 用量统计 ----
export interface UsageDaily {
id?: number
user_id?: number
model_id?: number
date: string
requests: number
input_tokens: number
output_tokens: number
cache_read_tokens: number
cost: number
}
export interface UsageTotals {
requests: number
input_tokens: number
output_tokens: number
cache_read_tokens: number
cost: number
}
export interface UsageStatsData {
dates: string[]
daily: Record<string, UsageDaily>
totals: UsageTotals
}
// 月度按模型用量分解(柱状图分色堆叠用)
export interface MonthlyModelUsage {
model_id: number
model_name: string
requests: number
input_tokens: number
output_tokens: number
cache_read_tokens: number
cost: number
}
// 单个自然月的聚合(models 已按 token 总量降序)
export interface MonthlyUsage {
month: string // "2026-09"
requests: number
input_tokens: number
output_tokens: number
cache_read_tokens: number
cost: number
models: MonthlyModelUsage[]
}
export interface MonthlyUsageData {
year: number
months: MonthlyUsage[]
}
export interface UsageLogItem {
id: number
request_id?: string
user_id: number
channel_id: number
model_id: number
model_name: string
protocol: string
input_tokens: number
output_tokens: number
cache_read_tokens: number
cache_creation_tokens: number
cost: number
latency_ms: number
status: string
error_code?: string | null
created_at: string
username?: string
raw_request?: string
raw_response?: string
}
export interface AdminUsageSummary {
totals: UsageTotals
per_user: Record<string, { user_id: number; requests: number; input_tokens: number; output_tokens: number; cost: number }>
}
+14 -3
View File
@@ -6,6 +6,9 @@ import {
KeyRoundIcon, KeyRoundIcon,
SettingsIcon, SettingsIcon,
GlobeIcon, GlobeIcon,
BoxesIcon,
SlidersHorizontalIcon,
ChartColumnBig,
} from '@lucide/vue' } from '@lucide/vue'
export type MenuLink = { label: string; to: string; icon?: Component } export type MenuLink = { label: string; to: string; icon?: Component }
@@ -17,6 +20,7 @@ declare module 'vue-router' {
icon?: Component icon?: Component
showInSidebar?: boolean showInSidebar?: boolean
requiresAuth?: boolean requiresAuth?: boolean
requiresAdmin?: boolean
open?: boolean open?: boolean
badge?: string badge?: string
} }
@@ -38,18 +42,21 @@ export const routes: RouteRecordRaw[] = [
redirect: '/dashboard/overview', redirect: '/dashboard/overview',
children: [ children: [
{ path: 'overview', name: 'Overview', component: () => import('@/views/dashboard/Overview.vue'), meta: { title: '仪表盘' } }, { path: 'overview', name: 'Overview', component: () => import('@/views/dashboard/Overview.vue'), meta: { title: '仪表盘' } },
{ path: 'usage', name: 'UsageStats', component: () => import('@/views/dashboard/UsageStats.vue'), meta: { title: '用量统计' } },
{ path: 'apikeys', name: 'ApiKeys', component: () => import('@/views/dashboard/ApiKeys.vue'), meta: { title: 'API Keys' } }, { path: 'apikeys', name: 'ApiKeys', component: () => import('@/views/dashboard/ApiKeys.vue'), meta: { title: 'API Keys' } },
{ {
path: 'manager', path: 'manager',
name: 'Manager', name: 'Manager',
meta: { title: '管理后台' }, meta: { title: '管理后台', requiresAdmin: true },
redirect: '/dashboard/manager/users', redirect: '/dashboard/manager/users',
children: [ children: [
{ path: 'users', name: 'User', component: () => import('@/views/dashboard/User.vue'), meta: { title: '用户管理' } }, { path: 'users', name: 'User', component: () => import('@/views/dashboard/User.vue'), meta: { title: '用户管理' } },
{ path: 'users/new', name: 'UserNew', component: () => import('@/views/dashboard/UserNew.vue'), meta: { title: '新建用户' } }, { path: 'users/new', name: 'UserNew', component: () => import('@/views/dashboard/UserNew.vue'), meta: { title: '新建用户' } },
{ path: 'users/view', name: 'UserView', component: () => import('@/views/dashboard/UserView.vue'), meta: { title: '用户详情' } }, { path: 'users/view', name: 'UserView', component: () => import('@/views/dashboard/UserView.vue'), meta: { title: '用户详情' } },
{ path: 'channels', name: 'Channels', component: () => import('@/views/dashboard/Keys.vue'), meta: { title: '渠道管理' } }, { path: 'channels', name: 'Channels', component: () => import('@/views/dashboard/ChannelsView.vue'), meta: { title: '渠道管理' } },
{ path: 'channels/view', name: 'ChannelView', component: () => import('@/views/dashboard/KeyView.vue'), meta: { title: '渠道详情' } }, { path: 'models', name: 'Models', component: () => import('@/views/dashboard/Models.vue'), meta: { title: '模型定价' } },
{ path: 'usage-logs', name: 'UsageLogs', component: () => import('@/views/dashboard/UsageLogs.vue'), meta: { title: '用量明细' } },
{ path: 'config', name: 'SystemConfig', component: () => import('@/views/dashboard/SystemConfig.vue'), meta: { title: '系统配置' } },
], ],
}, },
{ {
@@ -68,6 +75,7 @@ export const routes: RouteRecordRaw[] = [
// 控制台菜单(所有登录用户) // 控制台菜单(所有登录用户)
export const consoleMenu: MenuLink[] = [ export const consoleMenu: MenuLink[] = [
{ label: '仪表盘', to: '/dashboard/overview', icon: GaugeIcon }, { label: '仪表盘', to: '/dashboard/overview', icon: GaugeIcon },
{ label: '用量统计', to: '/dashboard/usage', icon: ChartColumnBig },
{ label: 'API Keys', to: '/dashboard/apikeys', icon: KeyRoundIcon }, { label: 'API Keys', to: '/dashboard/apikeys', icon: KeyRoundIcon },
{ label: '账户设置', to: '/dashboard/settings/profile', icon: SettingsIcon }, { label: '账户设置', to: '/dashboard/settings/profile', icon: SettingsIcon },
] ]
@@ -76,4 +84,7 @@ export const consoleMenu: MenuLink[] = [
export const adminMenu: MenuLink[] = [ export const adminMenu: MenuLink[] = [
{ label: '用户管理', to: '/dashboard/manager/users', icon: UsersRoundIcon }, { label: '用户管理', to: '/dashboard/manager/users', icon: UsersRoundIcon },
{ label: '渠道管理', to: '/dashboard/manager/channels', icon: GlobeIcon }, { label: '渠道管理', to: '/dashboard/manager/channels', icon: GlobeIcon },
{ label: '模型定价', to: '/dashboard/manager/models', icon: BoxesIcon },
{ label: '用量明细', to: '/dashboard/manager/usage-logs', icon: ChartColumnBig },
{ label: '系统配置', to: '/dashboard/manager/config', icon: SlidersHorizontalIcon },
] ]
+7 -3
View File
@@ -62,17 +62,21 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref, reactive, onMounted } from 'vue' import { ref, reactive, onMounted } from 'vue'
import { useRouter } from 'vue-router' import { useRoute, useRouter } from 'vue-router'
import { CircleAlert } from '@lucide/vue' import { CircleAlert } from '@lucide/vue'
import { useAuthStore } from '@/stores/auth'; import { useAuthStore } from '@/stores/auth';
import { useWebAuthStore } from '@/stores/webauth'; import { useWebAuthStore } from '@/stores/webauth';
import { useToast } from '@/composables/toast'; import { useToast } from '@/composables/toast';
const router = useRouter() const router = useRouter()
const route = useRoute()
const authStore = useAuthStore(); const authStore = useAuthStore();
const webauthStore = useWebAuthStore(); const webauthStore = useWebAuthStore();
const { setToast } = useToast(); const { setToast } = useToast();
// 被守卫拦下时带上原始目标,登录成功后回跳
const redirectPath = typeof route.query.redirect === 'string' ? route.query.redirect : '/dashboard'
const error = ref<string | null>(null) const error = ref<string | null>(null)
const loggingIn = ref(false) const loggingIn = ref(false)
const user = reactive({ const user = reactive({
@@ -113,7 +117,7 @@ const handleLogin = async () => {
localStorage.removeItem('rember'); localStorage.removeItem('rember');
} }
setToast('Logged in successfully.', 'success'); setToast('Logged in successfully.', 'success');
router.push('/dashboard'); router.push(redirectPath);
} }
} catch (err: any) { } catch (err: any) {
console.error('Login error:', err); console.error('Login error:', err);
@@ -130,7 +134,7 @@ const handlePasskeyLogin = async () => {
const res = await webauthStore.loginPasskey(); const res = await webauthStore.loginPasskey();
if (!!res?.code && res.code === 200) { if (!!res?.code && res.code === 200) {
setToast('Logged in successfully.', 'success'); setToast('Logged in successfully.', 'success');
router.push('/dashboard'); router.push(redirectPath);
} }
} catch (err: any) { } catch (err: any) {
console.error('Passkey login error:', err); console.error('Passkey login error:', err);
@@ -0,0 +1,183 @@
<script setup lang="ts">
import { onMounted, reactive, ref } from 'vue'
import { RefreshCw, Plus, X } from '@lucide/vue'
import request from '@/api/client'
import { useToast } from '@/composables/toast'
import Button from '@/components/ui/Button.vue'
import type { Channel, ChannelModelMapping } from '@/types'
function errMsg(e: unknown) {
return (e as any)?.response?.data?.error || (e as any)?.message || '请求失败'
}
const props = defineProps<{ channel: Channel }>()
const { setToast } = useToast()
const mappings = ref<ChannelModelMapping[]>([])
const remote = ref<string[]>([])
const selected = ref<string[]>([])
const loading = ref(false)
const fetched = ref(false)
const addForm = reactive({ custom_name: '', upstream_model: '' })
async function load() {
try {
const { data } = await request.get(`/admin/channels/${props.channel.id}/models`)
mappings.value = data.data?.items || data.data || []
} catch (e) {
setToast(errMsg(e), 'error')
}
}
async function fetchRemote() {
loading.value = true
try {
const { data } = await request.get(`/admin/channels/${props.channel.id}/models/remote`)
remote.value = data.data || []
selected.value = []
fetched.value = true
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
loading.value = false
}
}
async function addSelected() {
let added = 0
for (const name of selected.value) {
try {
await request.post(`/admin/channels/${props.channel.id}/models`, {
upstream_model: name,
})
added++
} catch {
/* 单个失败不中断 */
}
}
selected.value = []
setToast(added ? `已添加 ${added} 个模型` : '所选均已添加', 'success')
await load()
await fetchRemote()
}
async function addManual() {
if (!addForm.upstream_model.trim()) return
try {
await request.post(`/admin/channels/${props.channel.id}/models`, {
upstream_model: addForm.upstream_model.trim(),
custom_name: addForm.custom_name.trim() || undefined,
})
setToast('已添加', 'success')
addForm.custom_name = ''
addForm.upstream_model = ''
await load()
} catch (e) {
setToast(errMsg(e), 'error')
}
}
async function saveUpstream(b: ChannelModelMapping) {
try {
await request.patch(`/admin/channels/${props.channel.id}/models/${b.id}`, {
upstream_model: b.upstream_model,
})
setToast('已更新', 'success')
await load()
} catch (e) {
setToast(errMsg(e), 'error')
}
}
async function remove(b: ChannelModelMapping) {
if (!confirm(`解除模型 ${b.model_name} 的绑定?`)) return
try {
await request.delete(`/admin/channels/${props.channel.id}/models/${b.id}`)
setToast('已解除', 'success')
await load()
} catch (e) {
setToast(errMsg(e), 'error')
}
}
onMounted(load)
</script>
<template>
<div class="space-y-3">
<!-- 已允许的模型 -->
<div>
<p class="mb-1.5 text-xs font-medium text-base-content/50">已允许的模型({{ mappings.length }})</p>
<div v-if="mappings.length" class="flex flex-wrap gap-2">
<div
v-for="b in mappings"
:key="b.id"
class="inline-flex items-center gap-1.5 rounded-md border border-base-300/60 bg-base-100 px-2 py-1 font-mono text-[11px] text-base-content/60"
>
<span class="text-base-content">{{ b.model_name }}</span>
<span class="opacity-60">→</span>
<input
v-model="b.upstream_model"
class="w-28 rounded border border-transparent bg-transparent px-1 text-[11px] text-primary outline-none transition focus:border-primary/50 focus:bg-base-200/50"
@change="saveUpstream(b)"
/>
<button class="text-base-content/40 hover:text-error" aria-label="解除" @click="remove(b)">
<X :size="12" />
</button>
</div>
</div>
<p v-else class="text-xs text-base-content/50">尚未允许任何模型</p>
</div>
<!-- 从接口拉取 + 勾选 -->
<div class="border-t border-base-300/60 pt-3">
<div class="mb-1.5 flex items-center justify-between">
<p class="text-xs font-medium text-base-content/50">从接口拉取模型</p>
<Button size="sm" variant="ghost" :loading="loading" @click="fetchRemote">
<RefreshCw :size="13" />
拉取
</Button>
</div>
<div v-if="remote.length" class="flex max-h-36 flex-wrap gap-2 overflow-y-auto">
<label
v-for="m in remote"
:key="m"
class="flex cursor-pointer items-center gap-1.5 rounded-md border px-2 py-1 font-mono text-[11px] text-base-content/60 transition select-none"
:class="selected.includes(m) ? 'border-primary bg-primary/10 text-base-content' : 'border-base-300/60 hover:border-base-content/30'"
>
<input v-model="selected" type="checkbox" :value="m" class="size-3.5 accent-primary" />
{{ m }}
</label>
</div>
<div v-if="remote.length" class="mt-2">
<Button size="sm" @click="addSelected">
<Plus :size="13" />
添加所选({{ selected.length }})
</Button>
</div>
<p v-else-if="!loading" class="text-xs text-base-content/50">
{{ remote.length === 0 && fetched ? '接口返回的模型均已允许,无新增候选' : '点「拉取」获取渠道接口返回的新模型,勾选需要的加入' }}
</p>
</div>
<!-- 手动添加 -->
<div class="flex items-center gap-2 border-t border-base-300/60 pt-3">
<input
v-model="addForm.custom_name"
placeholder="自定义名称(可选)"
class="h-8 min-w-0 flex-1 rounded-md border border-base-300/60 bg-base-100 px-2 font-mono text-xs outline-none focus:border-primary"
@keyup.enter="addManual"
/>
<input
v-model="addForm.upstream_model"
placeholder="上游模型名"
class="h-8 min-w-0 flex-1 rounded-md border border-base-300/60 bg-base-100 px-2 font-mono text-xs outline-none focus:border-primary"
@keyup.enter="addManual"
/>
<Button size="sm" class="shrink-0" @click="addManual">
<Plus :size="13" />
添加
</Button>
</div>
</div>
</template>
@@ -0,0 +1,342 @@
<script setup lang="ts">
import { onMounted, reactive, ref } from 'vue'
import { ChevronDown, Zap, Pencil, Trash2, Layers } from '@lucide/vue'
import request from '@/api/client'
import { useToast } from '@/composables/toast'
import { PROTOCOL_OPTIONS, protocolShort } from '@/lib/protocol'
function errMsg(e: unknown) {
return (e as any)?.response?.data?.error || (e as any)?.message || '请求失败'
}
import ChannelModelsDrawer from '@/views/dashboard/ChannelModelsDrawer.vue'
import Button from '@/components/ui/Button.vue'
import Input from '@/components/ui/Input.vue'
import Modal from '@/components/ui/Modal.vue'
import Badge from '@/components/ui/Badge.vue'
import type { Channel } from '@/types'
const { setToast } = useToast()
const channels = ref<Channel[]>([])
const editOpen = ref(false)
const editing = ref<Channel | null>(null)
const saving = ref(false)
const busyId = ref<number | null>(null)
const expandedId = ref<number | null>(null)
function toggleDrawer(ch: Channel) {
expandedId.value = expandedId.value === ch.id ? null : ch.id
}
const form = reactive({
name: '',
formats: ['chat'] as string[],
base_url: '',
base_urls: { chat: '', responses: '', messages: '' } as Record<string, string>,
api_key: '',
weight: 1,
priority: 0,
timeout_ms: 120000,
max_concurrency: 16,
enabled: true,
})
async function load() {
try {
const { data } = await request.get('/admin/channels')
channels.value = data.data.items || data.data
} catch (e) {
setToast(errMsg(e), 'error')
}
}
function openCreate() {
editing.value = null
Object.assign(form, {
name: '', formats: ['chat'], base_url: '',
base_urls: { chat: '', responses: '', messages: '' },
api_key: '',
weight: 1, priority: 0, timeout_ms: 120000, max_concurrency: 16, enabled: true,
})
editOpen.value = true
}
function openEdit(ch: Channel) {
editing.value = ch
Object.assign(form, {
name: ch.name, formats: [...(ch.formats?.length ? ch.formats : ['chat'])],
base_url: ch.base_url,
base_urls: {
chat: ch.base_urls?.chat ?? '',
responses: ch.base_urls?.responses ?? '',
messages: ch.base_urls?.messages ?? '',
},
api_key: '',
weight: ch.weight, priority: ch.priority, timeout_ms: ch.timeout_ms,
max_concurrency: ch.max_concurrency, enabled: ch.enabled,
})
editOpen.value = true
}
async function save() {
if (form.formats.length === 0) {
setToast('请至少选择一种 API 格式', 'error')
return
}
saving.value = true
const payload = {
...form,
weight: Number(form.weight),
priority: Number(form.priority),
timeout_ms: Number(form.timeout_ms),
max_concurrency: Number(form.max_concurrency),
}
try {
if (editing.value) {
await request.put(`/admin/channels/${editing.value.id}`, payload)
setToast('渠道已更新', 'success')
} else {
await request.post('/admin/channels', payload)
setToast('渠道已创建', 'success')
}
editOpen.value = false
await load()
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
saving.value = false
}
}
async function remove(ch: Channel) {
if (!confirm(`删除渠道 ${ch.name}?关联的模型绑定也会清除。`)) return
try {
await request.delete(`/admin/channels/${ch.id}`)
setToast('渠道已删除', 'success')
await load()
} catch (e) {
setToast(errMsg(e), 'error')
}
}
async function testChannel(ch: Channel) {
busyId.value = ch.id
try {
await request.post(`/admin/channels/${ch.id}/test`)
setToast(`渠道 ${ch.name} 连接正常`, 'success')
} catch (e) {
setToast(`连接失败: ${errMsg(e)}`, 'error')
} finally {
busyId.value = null
await load()
}
}
onMounted(load)
</script>
<template>
<div class="mx-auto max-w-6xl">
<div class="mb-6 flex flex-wrap items-center justify-between gap-3">
<div>
<h1 class="text-lg font-semibold">渠道</h1>
<p class="text-sm text-base-content/60">接入上游服务,API Key 加密存储</p>
</div>
<Button class="shrink-0" @click="openCreate">添加渠道</Button>
</div>
<!-- 移动端:卡片列表 -->
<div class="space-y-3 md:hidden">
<div v-for="ch in channels" :key="ch.id" class="card border border-base-300/60 bg-base-100 p-4 shadow-sm" :class="ch.enabled ? 'border-l-2 border-l-success' : ''">
<div class="flex flex-wrap items-start justify-between gap-2">
<div class="min-w-0">
<p class="text-sm font-medium">{{ ch.name }}</p>
<div class="mt-1.5 flex flex-wrap gap-1">
<code
v-for="f in ch.formats || []"
:key="f"
class="rounded bg-base-200 px-1.5 py-0.5 font-mono text-[10px] text-base-content/60"
>{{ protocolShort(f) }}</code>
</div>
</div>
<div class="flex shrink-0 gap-1.5">
<Badge :variant="ch.health_status === 'healthy' ? 'ok' : ch.health_status === 'cooldown' ? 'err' : 'warn'">
{{ ch.health_status }}
</Badge>
<Badge :variant="ch.enabled ? 'ok' : 'neutral'">{{ ch.enabled ? '启用' : '停用' }}</Badge>
</div>
</div>
<p class="mt-2 truncate font-mono text-[11px] text-base-content/60">{{ ch.base_url }}</p>
<div class="mt-3 flex flex-wrap gap-x-3 gap-y-1.5 border-t border-base-300/60 pt-3">
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-primary" :disabled="busyId === ch.id" @click="testChannel(ch)">
<Zap :size="13" />
{{ busyId === ch.id ? '测试中…' : '测试' }}
</button>
<button class="inline-flex items-center gap-1 text-xs text-primary hover:text-primary/80" @click="toggleDrawer(ch)">
<Layers :size="13" />
支持的模型 {{ expandedId === ch.id ? '▴' : '▾' }}
</button>
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-base-content" @click="openEdit(ch)">
<Pencil :size="13" />
编辑
</button>
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-error" @click="remove(ch)">
<Trash2 :size="13" />
删除
</button>
</div>
<div v-if="expandedId === ch.id" class="mt-3 border-t border-base-300/60 pt-3">
<ChannelModelsDrawer :channel="ch" />
</div>
</div>
<p v-if="channels.length === 0" class="card border border-base-300/60 bg-base-100 px-4 py-10 text-center text-sm text-base-content/60">
还没有渠道,点击「添加渠道」
</p>
</div>
<!-- 桌面端:表格 -->
<div class="card hidden border border-base-300/60 bg-base-100 shadow-sm md:block">
<div class="overflow-x-auto">
<table class="w-full text-sm min-w-[820px]">
<thead>
<tr class="border-b border-base-300/60 text-left text-xs text-base-content/50">
<th scope="col" class="px-4 py-2.5 font-medium">名称</th>
<th scope="col" class="px-4 py-2.5 font-medium">API 格式</th>
<th scope="col" class="px-4 py-2.5 font-medium">Base URL</th>
<th scope="col" class="px-4 py-2.5 font-medium">Key</th>
<th scope="col" class="px-4 py-2.5 font-medium">健康</th>
<th scope="col" class="px-4 py-2.5 font-medium">启用</th>
<th scope="col" class="px-4 py-2.5" />
</tr>
</thead>
<tbody>
<template v-for="ch in channels" :key="ch.id">
<tr class="border-b border-base-300/40 last:border-0 hover:bg-base-200/50" :style="ch.enabled ? { borderLeft: '2px solid oklch(var(--p))' } : {}">
<td class="px-4 py-2.5">
<button class="inline-flex items-center gap-1.5 transition hover:text-primary" @click="toggleDrawer(ch)">
<span class="truncate">{{ ch.name }}</span>
<ChevronDown :size="12" class="shrink-0 text-base-content/50 transition-transform" :class="expandedId === ch.id ? 'rotate-180' : ''" />
</button>
</td>
<td class="px-4 py-2.5">
<div class="flex flex-col gap-0.5">
<code
v-for="f in ch.formats || []"
:key="f"
class="font-mono text-[11px] leading-4 text-base-content/60"
>{{ protocolShort(f) }}</code>
</div>
</td>
<td class="max-w-[220px] truncate px-4 py-2.5 font-mono text-xs text-base-content/60">{{ ch.base_url }}</td>
<td class="px-4 py-2.5 font-mono text-xs text-base-content/60">{{ ch.api_key_masked || '****' }}</td>
<td class="px-4 py-2.5">
<Badge :variant="ch.health_status === 'healthy' ? 'ok' : ch.health_status === 'cooldown' ? 'err' : 'warn'">
{{ ch.health_status }}
</Badge>
</td>
<td class="px-4 py-2.5 text-xs text-base-content/60">{{ ch.enabled ? '是' : '否' }}</td>
<td class="px-4 py-2.5 text-right">
<div class="flex justify-end gap-2">
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-primary" :disabled="busyId === ch.id" @click="testChannel(ch)">
<Zap :size="13" />
{{ busyId === ch.id ? '测试中…' : '测试' }}
</button>
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-base-content" @click="openEdit(ch)">
<Pencil :size="13" />
编辑
</button>
<button class="inline-flex items-center gap-1 text-xs text-base-content/60 hover:text-error" @click="remove(ch)">
<Trash2 :size="13" />
删除
</button>
</div>
</td>
</tr>
<tr v-if="expandedId === ch.id" class="bg-base-200/30">
<td colspan="7" class="px-4 py-3">
<ChannelModelsDrawer :channel="ch" />
</td>
</tr>
</template>
<tr v-if="channels.length === 0">
<td colspan="7" class="px-4 py-10 text-center text-sm text-base-content/60">还没有渠道,点击「添加渠道」</td>
</tr>
</tbody>
</table>
</div>
</div>
<Modal :open="editOpen" :title="editing ? '编辑渠道' : '添加渠道'" @close="editOpen = false">
<div class="space-y-4">
<Input v-model="form.name" label="名称" placeholder="openai" />
<div>
<span class="mb-1.5 block text-xs font-medium text-base-content/50">支持的 API 格式</span>
<div class="flex flex-wrap gap-2">
<label
v-for="opt in PROTOCOL_OPTIONS"
:key="opt.value"
class="flex cursor-pointer items-center gap-1.5 rounded-md border px-2.5 py-1.5 text-xs transition select-none"
:class="form.formats.includes(opt.value) ? 'border-primary bg-primary/10' : 'border-base-300 text-base-content/60 hover:border-base-content/30'"
>
<input
v-model="form.formats"
type="checkbox"
:value="opt.value"
class="size-3.5 accent-primary"
/>
{{ opt.label }}
</label>
</div>
<p class="mt-1.5 text-xs text-base-content/50">客户端协议不在其中时,网关自动转换为其支持的格式</p>
</div>
<Input
v-model="form.base_url"
label="Base URL(可选)"
placeholder="https://api.openai.com/v1"
:maxlength="255"
hint="支持前缀或完整端点,如 https://api.openai.com/v1 或 https://api.openai.com/v1/chat/completions;留空按供应商默认"
/>
<div class="space-y-3 rounded-md border border-base-300/60 p-3">
<p class="text-xs font-medium text-base-content/50">分协议 Base URL(可选,如智谱三种格式不同)</p>
<Input v-model="form.base_urls.chat" label="OpenAI Chat Completions" placeholder="留空用主 Base URL" :maxlength="255" />
<Input v-model="form.base_urls.responses" label="OpenAI Responses" placeholder="留空用主 Base URL" :maxlength="255" />
<Input v-model="form.base_urls.messages" label="Anthropic Messages" placeholder="留空用主 Base URL" :maxlength="255" />
<p class="text-xs text-base-content/50">网关按协议选对应 base_url 直通,无需为每种格式建多个渠道</p>
</div>
<Input
v-model="form.api_key"
label="上游 API Key"
:placeholder="editing ? '留空则不修改' : 'sk-...'"
/>
<div class="grid grid-cols-1 gap-4 sm:grid-cols-2">
<Input v-model="form.weight" label="权重" type="number" />
<Input v-model="form.priority" label="优先级" type="number" />
<Input v-model="form.timeout_ms" label="超时 (ms)" type="number" />
<Input v-model="form.max_concurrency" label="最大并发" type="number" />
</div>
<div class="flex items-center justify-between rounded-md border border-base-300/60 p-3">
<div>
<p class="text-sm font-medium">启用渠道</p>
<p class="text-xs text-base-content/50">禁用后该渠道不会被用于请求转发</p>
</div>
<button
type="button"
role="switch"
:aria-checked="form.enabled"
class="relative inline-flex h-6 w-11 shrink-0 cursor-pointer items-center rounded-full transition-colors focus-visible:outline-2 focus-visible:outline-offset-2 focus-visible:outline-primary"
:class="form.enabled ? 'bg-primary' : 'bg-base-200'"
@click="form.enabled = !form.enabled"
>
<span
class="pointer-events-none inline-block h-4 w-4 rounded-full bg-white shadow-sm ring-0 transition-transform"
:class="form.enabled ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
</div>
<template #footer>
<Button variant="ghost" @click="editOpen = false">取消</Button>
<Button :loading="saving" @click="save">{{ editing ? '保存' : '创建' }}</Button>
</template>
</Modal>
</div>
</template>
+148 -1
View File
@@ -87,6 +87,10 @@
<div class="flex items-center justify-end gap-3 border-t border-base-300/40 pt-4"> <div class="flex items-center justify-end gap-3 border-t border-base-300/40 pt-4">
<button type="button" @click="goBack" class="btn btn-ghost btn-sm">Back</button> <button type="button" @click="goBack" class="btn btn-ghost btn-sm">Back</button>
<button type="button" @click="testChannel" class="btn btn-warning btn-sm" :disabled="testing">
<span v-if="testing" class="loading loading-spinner loading-xs" aria-hidden="true"></span>
Test Connection
</button>
<button type="submit" class="btn btn-primary btn-sm px-5" :disabled="updating"> <button type="submit" class="btn btn-primary btn-sm px-5" :disabled="updating">
<span v-if="updating" class="loading loading-spinner loading-xs" aria-hidden="true"></span> <span v-if="updating" class="loading loading-spinner loading-xs" aria-hidden="true"></span>
Save Changes Save Changes
@@ -94,6 +98,47 @@
</div> </div>
</form> </form>
</div> </div>
<!-- Model Bindings -->
<div class="card border border-base-300/60 bg-base-100 shadow-sm">
<div class="card-body gap-4 p-4 sm:p-6">
<div class="flex items-center justify-between">
<h2 class="text-xs font-semibold uppercase tracking-wider text-base-content/50">Model Bindings</h2>
<button class="btn btn-primary btn-sm" @click="openAddModelModal">
<PlusIcon class="h-4 w-4" aria-hidden="true" />Add Model
</button>
</div>
<div v-if="bindings.length > 0" class="overflow-x-auto">
<table class="table table-sm">
<thead>
<tr class="text-xs uppercase tracking-wider text-base-content/50">
<th>Model Name</th>
<th>Upstream Model</th>
<th class="text-right">Weight</th>
<th class="text-right"><span class="sr-only">Actions</span></th>
</tr>
</thead>
<tbody>
<tr v-for="b in bindings" :key="b.id" class="border-base-300/40">
<td class="font-medium">{{ b.model_name }}</td>
<td class="font-mono text-xs">{{ b.upstream_model }}</td>
<td class="text-right">{{ b.weight }}</td>
<td class="text-right">
<button class="btn btn-ghost btn-xs btn-square text-error" @click="confirmDeleteBinding(b)"
aria-label="Delete binding">
<TrashIcon class="h-4 w-4" aria-hidden="true" />
</button>
</td>
</tr>
</tbody>
</table>
</div>
<div v-else class="py-6 text-center text-sm text-base-content/50">
No model bindings configured.
</div>
</div>
</div>
</div> </div>
<!-- Loading state --> <!-- Loading state -->
@@ -104,22 +149,59 @@
</div> </div>
</div> </div>
</div> </div>
<!-- Add Model Modal -->
<dialog ref="addModelModalRef" class="modal">
<div class="modal-box max-w-lg px-0 sm:px-6">
<form method="dialog">
<button class="btn btn-circle btn-ghost btn-sm absolute right-2 top-2" aria-label="Close dialog">✕</button>
</form>
<h3 class="mb-4 text-lg font-bold">Add Model Binding</h3>
<form @submit.prevent="addModelBinding" class="space-y-4">
<label class="floating-label">
<span>Model ID *</span>
<input v-model.number="newBinding.model_id" type="number" placeholder="Model ID" class="input w-full" required />
</label>
<label class="floating-label">
<span>Upstream Model Name *</span>
<input v-model="newBinding.upstream_model" type="text" placeholder="e.g. gpt-4o" class="input w-full" required />
</label>
<label class="floating-label">
<span>Weight</span>
<input v-model.number="newBinding.weight" type="number" min="1" placeholder="1" class="input w-full" />
</label>
<div class="modal-action">
<button type="button" class="btn btn-ghost" @click="closeAddModelModal">Cancel</button>
<button type="submit" class="btn btn-primary" :disabled="addingModel">
{{ addingModel ? 'Adding...' : 'Add' }}
</button>
</div>
</form>
</div>
<form method="dialog" class="modal-backdrop">
<button aria-label="Close dialog">close</button>
</form>
</dialog>
</div> </div>
</template> </template>
<script setup lang="ts"> <script setup lang="ts">
import { computed, onMounted, ref } from 'vue'; import { computed, onMounted, ref } from 'vue';
import { useRoute, useRouter } from 'vue-router'; import { useRoute, useRouter } from 'vue-router';
import { useChannelStore, type Channel } from '../../stores/channel'; import { useChannelStore, type Channel, type ChannelModelBinding } from '../../stores/channel';
import BreadcrumbHeader from '@/components/dashboard/BreadcrumbHeader.vue'; import BreadcrumbHeader from '@/components/dashboard/BreadcrumbHeader.vue';
import { useToast } from '@/composables/toast'; import { useToast } from '@/composables/toast';
import { PlusIcon, TrashIcon } from '@lucide/vue';
const route = useRoute(); const route = useRoute();
const router = useRouter(); const router = useRouter();
const channelStore = useChannelStore(); const channelStore = useChannelStore();
const { setToast } = useToast(); const { setToast } = useToast();
const updating = ref(false); const updating = ref(false);
const testing = ref(false);
const api_key = ref(''); const api_key = ref('');
const bindings = ref<ChannelModelBinding[]>([]);
const channelId = computed(() => route.query.id); const channelId = computed(() => route.query.id);
const ch = computed(() => channelStore.channel); const ch = computed(() => channelStore.channel);
@@ -127,9 +209,16 @@ const ch = computed(() => channelStore.channel);
onMounted(async () => { onMounted(async () => {
if (channelId.value) { if (channelId.value) {
await channelStore.fetchChannel(channelId.value as string); await channelStore.fetchChannel(channelId.value as string);
await fetchBindings();
} }
}); });
const fetchBindings = async () => {
if (channelId.value) {
bindings.value = await channelStore.fetchChannelModels(channelId.value as string);
}
};
const toggleEnabled = () => { const toggleEnabled = () => {
if (!ch.value) return; if (!ch.value) return;
ch.value.enabled = !ch.value.enabled; ch.value.enabled = !ch.value.enabled;
@@ -163,7 +252,65 @@ const updateCh = async () => {
} }
}; };
const testChannel = async () => {
if (!ch.value) return;
testing.value = true;
try {
const result = await channelStore.testChannel(ch.value.id);
setToast(`Connection OK (${result.data?.latency_ms}ms)`, 'success');
} catch (err: any) {
setToast(err.response?.data?.error || 'Connection test failed', 'error');
} finally {
testing.value = false;
}
};
const goBack = () => { const goBack = () => {
router.push({ name: 'Channels' }); router.push({ name: 'Channels' });
}; };
// Model binding
const addModelModalRef = ref<HTMLDialogElement | null>(null);
const addingModel = ref(false);
const newBinding = ref({
model_id: 0,
upstream_model: '',
weight: 1,
});
const openAddModelModal = () => {
newBinding.value = { model_id: 0, upstream_model: '', weight: 1 };
addModelModalRef.value?.showModal();
};
const closeAddModelModal = () => {
addModelModalRef.value?.close();
};
const addModelBinding = async () => {
if (!channelId.value || !newBinding.value.model_id || !newBinding.value.upstream_model) return;
addingModel.value = true;
try {
await channelStore.addChannelModel(channelId.value as string, newBinding.value);
setToast('Model binding added', 'success');
closeAddModelModal();
await fetchBindings();
} catch (err: any) {
setToast(err.response?.data?.error || 'Failed to add binding', 'error');
} finally {
addingModel.value = false;
}
};
const confirmDeleteBinding = async (b: ChannelModelBinding) => {
if (confirm(`Remove binding for model "${b.model_name}"?`)) {
try {
await channelStore.deleteChannelModel(channelId.value as string, b.id);
setToast('Binding removed', 'success');
await fetchBindings();
} catch (err: any) {
setToast('Delete failed', 'error');
}
}
};
</script> </script>
+277
View File
@@ -0,0 +1,277 @@
<template>
<div class="space-y-5">
<BreadcrumbHeader />
<div class="flex flex-wrap items-center justify-between gap-3">
<div>
<p class="text-sm text-base-content/60">Manage model pricing and channel bindings.</p>
<div v-if="summary" class="mt-1 flex gap-3 text-xs text-base-content/50">
<span>Total: {{ summary.total }}</span>
<span v-if="summary.unpriced > 0" class="text-warning">{{ summary.unpriced }} unpriced</span>
<span v-if="summary.missing.length > 0" class="text-error">{{ summary.missing.length }} orphan bindings</span>
</div>
</div>
<div class="flex items-center gap-2">
<button v-if="models.length > 0" class="btn btn-ghost btn-sm" @click="confirmDeleteUnused">
<TrashIcon class="h-4 w-4" aria-hidden="true" />Clean Unused
</button>
<button class="btn btn-primary btn-sm" @click="openCreateModal" aria-label="Create new model">
<PlusIcon class="h-4 w-4" aria-hidden="true" />New Model
</button>
</div>
</div>
<!-- Model cards -->
<div class="space-y-3">
<div v-for="m in models" :key="m.id"
class="card border bg-base-100 shadow-sm"
:class="m.channels && m.channels.length > 0 ? 'border-base-300/60' : 'border-warning/60 bg-warning/5'">
<div class="px-4 py-3">
<div class="flex flex-wrap items-center justify-between gap-x-4 gap-y-2">
<div class="flex flex-wrap items-center gap-2">
<span class="font-mono text-sm font-medium text-base-content">{{ m.name }}</span>
<span v-if="m.display_name" class="text-xs text-base-content/50">{{ m.display_name }}</span>
<span v-if="m.channels && m.channels.length > 0" class="badge badge-xs badge-ghost">渠道允许</span>
<span v-else class="badge badge-xs bg-yellow-200 text-yellow-800 dark:bg-yellow-900/50 dark:text-yellow-300">悬空</span>
<span :class="m.enabled ? 'badge badge-xs bg-green-200 text-green-800 dark:bg-green-900/50 dark:text-green-300' : 'badge badge-xs badge-ghost'">{{ m.enabled ? '启用' : '停用' }}</span>
<span v-if="m.denied" class="badge badge-xs badge-error">已禁止</span>
<span v-if="m.needs_pricing" class="badge badge-xs bg-orange-200 text-orange-800 dark:bg-orange-900/50 dark:text-orange-300">未定价</span>
</div>
<div class="flex gap-2">
<button class="btn btn-ghost btn-xs" @click="openEditModal(m)">编辑</button>
<button class="btn btn-ghost btn-xs text-error" @click="confirmDeleteModel(m)">删除</button>
</div>
</div>
<div class="mt-2 flex flex-wrap items-center gap-3">
<span class="font-mono text-xs text-base-content/60">入 {{ formatPrice(m.input_price) }}</span>
<span class="font-mono text-xs text-base-content/60">出 {{ formatPrice(m.output_price) }}</span>
<span class="font-mono text-xs text-base-content/60">缓存读 {{ formatPrice(m.cache_read_price) }}</span>
</div>
</div>
<div v-if="m.channels && m.channels.length > 0" class="border-t border-base-300/60 px-4 py-2">
<p class="mb-1.5 text-[11px] font-medium text-base-content/50">允许渠道</p>
<div class="flex flex-wrap gap-2">
<span v-for="ch in m.channels" :key="ch.id"
class="inline-flex items-center rounded-md border border-base-300/60 bg-base-100 px-2 py-0.5 font-mono text-[11px] text-base-content/60">
{{ ch.channel_name }} → {{ ch.upstream_model }}
</span>
</div>
</div>
<p v-else class="border-t border-base-300/60 bg-amber-100 px-4 py-2 text-xs font-medium text-amber-900 dark:bg-amber-900/40 dark:text-amber-100">
悬空模型:无任何渠道提供,客户端无法调用
</p>
</div>
<!-- Empty state -->
<div v-if="models.length === 0" class="card border border-base-300/60 bg-base-100 px-4 py-14 text-center">
<BoxesIcon class="mx-auto h-10 w-10 text-base-content/20" aria-hidden="true" />
<h2 class="mt-2 text-sm font-semibold">No models yet</h2>
<p class="mt-1 max-w-xs text-sm text-base-content/60">
Add models to manage pricing and channel bindings.
</p>
<button class="btn btn-primary btn-sm mt-3" @click="openCreateModal">
<PlusIcon class="h-4 w-4" aria-hidden="true" />Create Model
</button>
</div>
</div>
<!-- Create/Edit modal -->
<dialog ref="modalRef" class="modal">
<div class="modal-box max-w-lg px-0 sm:px-6">
<form method="dialog">
<button class="btn btn-circle btn-ghost btn-sm absolute right-2 top-2" aria-label="Close dialog">✕</button>
</form>
<h3 class="mb-4 text-lg font-bold">{{ editingModel ? 'Edit Model' : 'New Model' }}</h3>
<form @submit.prevent="saveModel" class="space-y-4">
<label class="floating-label">
<span>Model Name *</span>
<input v-model="form.name" type="text" placeholder="e.g. gpt-4o" class="input w-full" required
:disabled="!!editingModel" />
</label>
<label class="floating-label">
<span>Display Name</span>
<input v-model="form.display_name" type="text" placeholder="e.g. GPT-4o" class="input w-full" />
</label>
<div class="grid grid-cols-3 gap-3">
<label class="floating-label">
<span>Input $/M tokens</span>
<input v-model.number="form.input_price" type="number" step="0.01" min="0" placeholder="0"
class="input w-full" />
</label>
<label class="floating-label">
<span>Output $/M tokens</span>
<input v-model.number="form.output_price" type="number" step="0.01" min="0" placeholder="0"
class="input w-full" />
</label>
<label class="floating-label">
<span>Cache Read $/M</span>
<input v-model.number="form.cache_read_price" type="number" step="0.01" min="0" placeholder="0"
class="input w-full" />
</label>
</div>
<div class="grid grid-cols-2 gap-3">
<label class="floating-label">
<span>Sort Order</span>
<input v-model.number="form.sort" type="number" min="0" placeholder="0" class="input w-full" />
</label>
<div class="flex items-center gap-2 pt-6">
<input type="checkbox" class="toggle toggle-success toggle-sm" v-model="form.enabled" />
<span class="text-sm">{{ form.enabled ? 'Enabled' : 'Disabled' }}</span>
</div>
</div>
<div class="modal-action">
<button type="button" class="btn btn-ghost" @click="closeModal">Cancel</button>
<button type="submit" class="btn btn-primary" :disabled="saving">
{{ saving ? 'Saving...' : (editingModel ? 'Update' : 'Create') }}
</button>
</div>
</form>
</div>
<form method="dialog" class="modal-backdrop">
<button aria-label="Close dialog">close</button>
</form>
</dialog>
</div>
</template>
<script setup lang="ts">
import { ref, reactive, onMounted } from 'vue';
import BreadcrumbHeader from '@/components/dashboard/BreadcrumbHeader.vue';
import { useModelStore, type Model, type NewModelPayload } from '@/stores/model';
import { useToast } from '@/composables/toast';
import {
BoxesIcon, PencilIcon, PlusIcon, TrashIcon
} from '@lucide/vue';
const modelStore = useModelStore();
const { setToast } = useToast();
const models = ref<Model[]>([]);
const summary = ref(modelStore.summary);
const editingModel = ref<Model | null>(null);
const saving = ref(false);
const form = reactive<NewModelPayload & { enabled: boolean }>({
name: '',
display_name: '',
input_price: 0,
output_price: 0,
cache_read_price: 0,
sort: 0,
enabled: true,
});
onMounted(async () => {
await fetchModels();
});
const fetchModels = async () => {
await modelStore.fetchModels();
models.value = modelStore.models;
summary.value = modelStore.summary;
};
const formatPrice = (price: number) => {
return price === 0 ? '-' : `$${price.toFixed(2)}`;
};
const toggleEnabled = async (m: Model) => {
try {
await modelStore.updateModel(m.id, { enabled: !m.enabled });
setToast(`Model ${m.name} ${m.enabled ? 'disabled' : 'enabled'}`, 'success');
await fetchModels();
} catch (error: any) {
setToast('Status update failed', 'error');
}
};
const openCreateModal = () => {
editingModel.value = null;
form.name = '';
form.display_name = '';
form.input_price = 0;
form.output_price = 0;
form.cache_read_price = 0;
form.sort = 0;
form.enabled = true;
modalRef.value?.showModal();
};
const openEditModal = (m: Model) => {
editingModel.value = m;
form.name = m.name;
form.display_name = m.display_name || '';
form.input_price = m.input_price;
form.output_price = m.output_price;
form.cache_read_price = m.cache_read_price;
form.sort = m.sort;
form.enabled = m.enabled;
modalRef.value?.showModal();
};
const saveModel = async () => {
saving.value = true;
try {
if (editingModel.value) {
await modelStore.updateModel(editingModel.value.id, {
display_name: form.display_name,
input_price: form.input_price,
output_price: form.output_price,
cache_read_price: form.cache_read_price,
sort: form.sort,
enabled: form.enabled,
});
setToast('Model updated', 'success');
} else {
await modelStore.createModel({
name: form.name,
display_name: form.display_name,
input_price: form.input_price,
output_price: form.output_price,
cache_read_price: form.cache_read_price,
sort: form.sort,
enabled: form.enabled,
});
setToast('Model created', 'success');
}
closeModal();
await fetchModels();
} catch (error: any) {
setToast(error.message || 'Save failed', 'error');
} finally {
saving.value = false;
}
};
const confirmDeleteModel = async (m: Model) => {
if (confirm(`Delete model "${m.name}"? This will also remove all channel bindings.`)) {
try {
await modelStore.deleteModel(m.id);
setToast(`Model ${m.name} deleted`, 'success');
await fetchModels();
} catch (error: any) {
setToast('Delete failed', 'error');
}
}
};
const confirmDeleteUnused = async () => {
if (confirm('Delete all models that are not bound to any channel?')) {
try {
const result = await modelStore.deleteUnusedModels();
const count = result.data?.count || 0;
setToast(`Deleted ${count} unused models`, 'success');
await fetchModels();
} catch (error: any) {
setToast('Cleanup failed', 'error');
}
}
};
const modalRef = ref<HTMLDialogElement | null>(null);
const closeModal = () => {
modalRef.value?.close();
};
</script>
+225
View File
@@ -0,0 +1,225 @@
<script setup lang="ts">
import { computed, onMounted, reactive, ref } from 'vue'
import request from '@/api/client'
import { useToast } from '@/composables/toast'
import Button from '@/components/ui/Button.vue'
import Input from '@/components/ui/Input.vue'
import Modal from '@/components/ui/Modal.vue'
import Badge from '@/components/ui/Badge.vue'
import type { Model, ModelSummary } from '@/types'
function errMsg(e: unknown) {
return (e as any)?.response?.data?.error || (e as any)?.message || '请求失败'
}
const { setToast } = useToast()
const models = ref<Model[]>([])
const summary = ref<ModelSummary>({ total: 0, unpriced: 0, missing: [], denied_count: 0 })
const editOpen = ref(false)
const editing = ref<Model | null>(null)
const saving = ref(false)
const quickName = ref('')
const clearing = ref(false)
const unused = computed(() => models.value.filter((m) => m.channels.length === 0))
async function clearUnused() {
if (!unused.value.length) {
setToast('没有未绑定渠道的模型', 'info')
return
}
const names = unused.value.map((m) => m.name)
if (!confirm(`确定删除 ${names.length} 个未绑定渠道的模型?\n\n${names.join('\n')}`)) return
clearing.value = true
try {
const { data } = await request.delete('/admin/models/unused')
setToast(`已清除 ${data.data.count} 个模型`, 'success')
await load()
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
clearing.value = false
}
}
function quickAdd() {
openCreate()
if (quickName.value) form.name = quickName.value.trim()
}
const form = reactive({
name: '',
input_price: 0,
output_price: 0,
cache_read_price: 0,
enabled: true,
})
async function load() {
try {
const { data } = await request.get('/admin/models')
models.value = data.data
summary.value = data.summary
} catch (e) {
setToast(errMsg(e), 'error')
}
}
function openCreate() {
editing.value = null
Object.assign(form, { name: '', input_price: 0, output_price: 0, cache_read_price: 0, enabled: true })
editOpen.value = true
}
function openEdit(m: Model) {
editing.value = m
Object.assign(form, {
name: m.name,
input_price: m.input_price, output_price: m.output_price, cache_read_price: m.cache_read_price,
enabled: m.enabled,
})
editOpen.value = true
}
async function save() {
saving.value = true
const payload = {
input_price: Number(form.input_price),
output_price: Number(form.output_price),
cache_read_price: Number(form.cache_read_price),
enabled: form.enabled,
}
try {
if (editing.value) {
await request.put(`/admin/models/${editing.value.id}`, payload)
setToast('模型已更新', 'success')
} else {
await request.post('/admin/models', { name: form.name, ...payload })
setToast('模型已创建', 'success')
}
editOpen.value = false
await load()
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
saving.value = false
}
}
async function removeModel(m: Model) {
if (!confirm(`删除模型 ${m.name}?`)) return
try {
await request.delete(`/admin/models/${m.id}`)
setToast('模型已删除', 'success')
await load()
} catch (e) {
setToast(errMsg(e), 'error')
}
}
onMounted(load)
</script>
<template>
<div class="mx-auto max-w-6xl">
<div class="mb-6 flex flex-wrap items-center justify-between gap-3">
<div>
<h1 class="text-lg font-semibold">模型与定价</h1>
<p class="text-sm text-base-content/60">接口导入不全时可直接输入模型名添加,如 glm-4.7-flash</p>
</div>
<div class="flex w-full flex-wrap gap-2 sm:w-auto sm:flex-nowrap">
<input
v-model="quickName"
placeholder="模型名,如 glm-4.7-flash"
class="h-10 min-w-0 flex-1 rounded-md border border-base-300/60 bg-base-100 px-3 font-mono text-xs outline-none focus:border-primary sm:w-52 sm:flex-none"
@keyup.enter="quickAdd"
/>
<Button class="shrink-0" @click="quickAdd">添加模型</Button>
<Button
size="md"
variant="danger"
class="shrink-0 px-1!"
:loading="clearing"
:disabled="!unused.length"
@click="clearUnused"
>
清除悬空{{ unused.length ? `(${unused.length})` : '' }}
</Button>
</div>
</div>
<!-- 提示:定价目录 = 渠道选中的模型 + 手动添加的模型 -->
<div v-if="summary.missing.length" class="card border border-error/50 bg-error/5 p-4">
<p class="text-sm font-medium text-error">以下渠道选中的模型不在定价目录</p>
<p v-for="(x, i) in summary.missing" :key="i" class="mt-1 font-mono text-xs text-base-content/60">
{{ x.channel }} → {{ x.upstream_model || '模型 #' + x.model_id }}(请到渠道抽屉重新选中,或手动添加)
</p>
</div>
<p v-else-if="summary.unpriced > 0" class="text-xs text-base-content/60">
有 <span class="font-mono text-warning">{{ summary.unpriced }}</span> 个渠道允许的模型未定价,网关将按示例价计费
</p>
<p v-else class="text-xs text-base-content/60">定价目录中渠道允许的模型均已定价</p>
<div class="space-y-3">
<div v-for="m in models" :key="m.id" :class="m.channels.length ? 'card border border-base-300/60 bg-base-100' : 'card border border-warning/60 bg-warning/5'">
<div class="px-4 py-3">
<div class="flex flex-wrap items-center justify-between gap-x-4 gap-y-2">
<div class="flex flex-wrap items-center gap-2">
<span class="font-mono text-sm text-base-content">{{ m.name }}</span>
<Badge v-if="m.channels.length" variant="neutral">渠道允许</Badge>
<Badge v-else variant="warn">悬空</Badge>
<Badge :variant="m.enabled ? 'ok' : 'neutral'">{{ m.enabled ? '启用' : '停用' }}</Badge>
<Badge v-if="m.denied" variant="err">已禁止</Badge>
<Badge v-if="m.needs_pricing" variant="warn">未定价</Badge>
</div>
<div class="flex gap-2">
<button class="text-xs text-base-content/60 hover:text-base-content" @click="openEdit(m)">编辑</button>
<button class="text-xs text-base-content/60 hover:text-error" @click="removeModel(m)">删除</button>
</div>
</div>
<div class="mt-2 flex flex-wrap items-center gap-3">
<span class="font-mono text-xs text-base-content/60">入 {{ m.input_price }}</span>
<span class="font-mono text-xs text-base-content/60">出 {{ m.output_price }}</span>
<span class="font-mono text-xs text-base-content/60">缓存读 {{ m.cache_read_price }}</span>
</div>
</div>
<div v-if="m.channels.length" class="border-t border-base-300/60 px-4 py-2">
<p class="mb-1.5 text-[11px] font-medium text-base-content/50">允许渠道(渠道抽屉中管理)</p>
<div class="flex flex-wrap gap-2">
<span
v-for="b in m.channels"
:key="b.id"
class="inline-flex items-center rounded-md border border-base-300/60 bg-base-100 px-2 py-0.5 font-mono text-[11px] text-base-content/60"
>
{{ b.channel_name }} → {{ b.upstream_model }}
</span>
</div>
</div>
<p v-else class="border-t border-base-300/60 px-4 py-2 text-xs text-warning">
悬空模型:无任何渠道提供,客户端无法调用
</p>
</div>
<p v-if="models.length === 0" class="card border border-base-300/60 bg-base-100 px-4 py-10 text-center text-sm text-base-content/60">
还没有模型,点击「添加模型」或到渠道页「导入模型」
</p>
</div>
<!-- 模型编辑 -->
<Modal :open="editOpen" :title="editing ? '编辑模型' : '添加模型'" @close="editOpen = false">
<div class="space-y-4">
<Input v-model="form.name" label="模型名" placeholder="claude-sonnet-5" :disabled="!!editing" />
<div class="grid grid-cols-1 gap-4 sm:grid-cols-2">
<Input v-model="form.input_price" label="输入价格 /1M" type="number" />
<Input v-model="form.output_price" label="输出价格 /1M" type="number" />
<Input v-model="form.cache_read_price" label="缓存读价格 /1M" type="number" />
</div>
</div>
<template #footer>
<Button variant="ghost" @click="editOpen = false">取消</Button>
<Button :loading="saving" @click="save">{{ editing ? '保存' : '创建' }}</Button>
</template>
</Modal>
</div>
</template>
-4
View File
@@ -170,8 +170,6 @@
<tr class="text-xs uppercase tracking-wider text-base-content/50"> <tr class="text-xs uppercase tracking-wider text-base-content/50">
<th class="pl-4">Name</th> <th class="pl-4">Name</th>
<th>Create Time</th> <th>Create Time</th>
<th>Sign Count</th>
<th>Device</th>
<th class="pr-4 text-right"><span class="sr-only">Actions</span></th> <th class="pr-4 text-right"><span class="sr-only">Actions</span></th>
</tr> </tr>
</thead> </thead>
@@ -179,8 +177,6 @@
<tr v-for="passkey in passkeys" :key="passkey.id" class="border-base-300/40 hover:bg-base-200/50"> <tr v-for="passkey in passkeys" :key="passkey.id" class="border-base-300/40 hover:bg-base-200/50">
<td class="pl-4 font-medium">{{ passkey.name }}</td> <td class="pl-4 font-medium">{{ passkey.name }}</td>
<td class="tabular-nums text-base-content/70">{{ formatDateTime(passkey.created_at) }}</td> <td class="tabular-nums text-base-content/70">{{ formatDateTime(passkey.created_at) }}</td>
<td class="tabular-nums">{{ passkey.sign_count }}</td>
<td class="text-base-content/70">{{ passkey.device_type }}</td>
<td class="pr-4 text-right"> <td class="pr-4 text-right">
<button class="btn btn-ghost btn-xs btn-square text-error" <button class="btn btn-ghost btn-xs btn-square text-error"
@click="confirmRmovePasskey(passkey)" aria-label="Delete passkey"> @click="confirmRmovePasskey(passkey)" aria-label="Delete passkey">
@@ -0,0 +1,162 @@
<script setup lang="ts">
import { ref, onMounted } from 'vue'
import request from '@/api/client'
import { useToast } from '@/composables/toast'
function errMsg(e: unknown) {
return (e as any)?.response?.data?.error || (e as any)?.message || '请求失败'
}
const { setToast } = useToast()
const loading = ref(false)
const saving = ref(false)
const registrationEnabled = ref(true)
const passwordLoginEnabled = ref(true)
const logRawRequests = ref(false)
async function load() {
loading.value = true
try {
const [regRes, pwdRes] = await Promise.all([
request.get('/admin/config/registration'),
request.get('/admin/config/password-login'),
])
registrationEnabled.value = regRes.data.data.enabled
passwordLoginEnabled.value = pwdRes.data.data.enabled
// 原始请求/响应记录开关(通用配置键 log_raw_requests)
const cfgRes = await request.get('/admin/config')
logRawRequests.value = cfgRes.data?.data?.log_raw_requests === 'true'
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
loading.value = false
}
}
async function saveRegistration(enabled: boolean) {
saving.value = true
try {
await request.put('/admin/config/registration', { enabled })
registrationEnabled.value = enabled
setToast('注册设置已更新', 'success')
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
saving.value = false
}
}
async function savePasswordLogin(enabled: boolean) {
saving.value = true
try {
await request.put('/admin/config/password-login', { enabled })
passwordLoginEnabled.value = enabled
setToast('密码登录设置已更新', 'success')
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
saving.value = false
}
}
async function saveLogRawRequests(enabled: boolean) {
saving.value = true
try {
await request.put('/admin/config', { log_raw_requests: enabled ? 'true' : 'false' })
logRawRequests.value = enabled
setToast(enabled ? '已开启原始请求/响应记录' : '已关闭原始请求/响应记录', 'success')
} catch (e) {
setToast(errMsg(e), 'error')
} finally {
saving.value = false
}
}
onMounted(load)
</script>
<template>
<div class="mx-auto max-w-2xl space-y-6">
<div>
<h1 class="text-lg font-semibold">系统配置</h1>
<p class="text-sm text-base-content/60">管理平台全局设置</p>
</div>
<div v-if="loading" class="py-10 text-center text-sm text-base-content/50">加载中…</div>
<template v-else>
<!-- 开放注册 -->
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<div class="flex items-center justify-between">
<div>
<h3 class="text-sm font-medium">开放注册</h3>
<p class="mt-1 text-xs text-base-content/50">允许新用户通过注册页面创建账号</p>
</div>
<button
type="button"
role="switch"
:aria-checked="registrationEnabled"
class="relative inline-flex h-6 w-11 shrink-0 cursor-pointer items-center rounded-full transition-colors focus-visible:outline-2 focus-visible:outline-offset-2 focus-visible:outline-primary"
:class="registrationEnabled ? 'bg-primary' : 'bg-base-200'"
:disabled="saving"
@click="saveRegistration(!registrationEnabled)"
>
<span
class="pointer-events-none inline-block h-4 w-4 rounded-full bg-white shadow-sm ring-0 transition-transform"
:class="registrationEnabled ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
</div>
<!-- 密码登录 -->
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<div class="flex items-center justify-between">
<div>
<h3 class="text-sm font-medium">密码登录</h3>
<p class="mt-1 text-xs text-base-content/50">允许用户通过用户名和密码登录</p>
</div>
<button
type="button"
role="switch"
:aria-checked="passwordLoginEnabled"
class="relative inline-flex h-6 w-11 shrink-0 cursor-pointer items-center rounded-full transition-colors focus-visible:outline-2 focus-visible:outline-offset-2 focus-visible:outline-primary"
:class="passwordLoginEnabled ? 'bg-primary' : 'bg-base-200'"
:disabled="saving"
@click="savePasswordLogin(!passwordLoginEnabled)"
>
<span
class="pointer-events-none inline-block h-4 w-4 rounded-full bg-white shadow-sm ring-0 transition-transform"
:class="passwordLoginEnabled ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
</div>
<!-- 原始请求/响应记录(仅管理员) -->
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<div class="flex items-center justify-between">
<div>
<h3 class="text-sm font-medium">记录原始请求/响应</h3>
<p class="mt-1 text-xs text-base-content/50">仅对管理员账号生效:在用量明细中保存每次请求的客户端原始请求体与上游原始响应体(流式含全部 SSE 事件),用于排障。会显著增加存储。</p>
</div>
<button
type="button"
role="switch"
:aria-checked="logRawRequests"
class="relative inline-flex h-6 w-11 shrink-0 cursor-pointer items-center rounded-full transition-colors focus-visible:outline-2 focus-visible:outline-offset-2 focus-visible:outline-primary"
:class="logRawRequests ? 'bg-primary' : 'bg-base-200'"
:disabled="saving"
@click="saveLogRawRequests(!logRawRequests)"
>
<span
class="pointer-events-none inline-block h-4 w-4 rounded-full bg-white shadow-sm ring-0 transition-transform"
:class="logRawRequests ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
</div>
</template>
</div>
</template>
+265
View File
@@ -0,0 +1,265 @@
<template>
<div class="space-y-5">
<BreadcrumbHeader />
<!-- 汇总卡片 -->
<div class="grid grid-cols-2 gap-4 lg:grid-cols-4">
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">请求总数</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtInt(summary?.requests) }}</p>
</div>
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">输入 Tokens</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtInt(summary?.input_tokens) }}</p>
</div>
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">输出 Tokens</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtInt(summary?.output_tokens) }}</p>
</div>
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">总费用 (USD)</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtCost(summary?.cost) }}</p>
</div>
</div>
<!-- 筛选栏 -->
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<div class="flex flex-wrap items-center gap-2">
<select v-model="filters.protocol" class="select select-sm border-base-300 bg-base-100" @change="applyFilters">
<option value="">全部协议</option>
<option value="chat">chat</option>
<option value="messages">messages</option>
<option value="responses">responses</option>
</select>
<select v-model="filters.status" class="select select-sm border-base-300 bg-base-100" @change="applyFilters">
<option value="">全部状态</option>
<option value="success">成功</option>
<option value="error">失败</option>
<option value="canceled">已取消</option>
</select>
<input
v-model="filters.model"
class="input input-sm w-44 border-base-300 bg-base-100"
placeholder="模型名称"
@keyup.enter="applyFilters"
/>
<input
v-model="filters.userId"
class="input input-sm w-32 border-base-300 bg-base-100"
placeholder="用户 ID"
@keyup.enter="applyFilters"
/>
<button class="btn btn-primary btn-sm" @click="applyFilters">
<SearchIcon class="h-4 w-4" aria-hidden="true" />筛选
</button>
<button class="btn btn-ghost btn-sm" @click="resetFilters">
<RotateCcwIcon class="h-4 w-4" aria-hidden="true" />重置
</button>
</div>
</div>
<!-- 明细表格 -->
<div class="card border border-base-300/60 bg-base-100 shadow-sm">
<div v-if="store.loading && !store.adminLogs.length" class="px-4 py-12 text-center text-sm text-base-content/50">加载中…</div>
<div v-else-if="store.adminLogs.length" class="overflow-x-auto">
<table class="table table-sm">
<thead>
<tr class="text-xs uppercase tracking-wider text-base-content/50">
<th>ID</th>
<th>用户</th>
<th>时间</th>
<th>模型</th>
<th>协议</th>
<th class="text-right">输入</th>
<th class="text-right">输出</th>
<th class="text-right">费用</th>
<th>状态</th>
<th class="pr-4 text-right"><span class="sr-only">Actions</span></th>
</tr>
</thead>
<tbody>
<tr v-for="l in store.adminLogs" :key="l.id" class="border-base-300/40 hover:bg-base-200/50">
<td class="tabular-nums text-base-content/60">{{ l.id }}</td>
<td class="whitespace-nowrap font-medium">
<span v-if="l.username">{{ l.username }}</span>
<span v-else class="text-base-content/50">#{{ l.user_id }}</span>
</td>
<td class="whitespace-nowrap tabular-nums text-base-content/70">{{ fmtTime(l.created_at) }}</td>
<td class="max-w-40 truncate" :title="l.model_name">{{ l.model_name }}</td>
<td><span class="badge badge-ghost badge-sm">{{ l.protocol }}</span></td>
<td class="text-right tabular-nums">{{ fmtInt(l.input_tokens) }}</td>
<td class="text-right tabular-nums">{{ fmtInt(l.output_tokens) }}</td>
<td class="text-right tabular-nums">{{ fmtCost(l.cost) }}</td>
<td><span class="badge badge-sm" :class="statusClass(l.status)">{{ statusLabel(l.status) }}</span></td>
<td class="pr-3">
<div class="flex items-center justify-end gap-1">
<button
v-if="l.raw_request || l.raw_response"
class="btn btn-ghost btn-xs btn-square"
:aria-label="`View raw data for request ${l.id}`"
@click="viewRaw(l)"
>
<FileTextIcon class="h-4 w-4" aria-hidden="true" />
</button>
</div>
</td>
</tr>
</tbody>
</table>
</div>
<div v-else class="px-4 py-12 text-center text-sm text-base-content/50">暂无用量记录</div>
<Pagination
v-if="store.adminLogsTotal > 0"
:current-page="page"
:total-items="store.adminLogsTotal"
:page-size="pageSize"
:page-size-options="[10, 20, 50, 100]"
@change-page="changePage"
/>
</div>
<!-- 原始请求/响应 查看弹窗 -->
<dialog ref="rawModal" class="modal">
<div class="modal-box max-w-3xl">
<form method="dialog">
<button class="btn btn-circle btn-ghost btn-sm absolute right-2 top-2" aria-label="Close">✕</button>
</form>
<h3 class="text-lg font-semibold">原始数据 #{{ currentRaw?.id }}</h3>
<p class="mt-1 text-xs text-base-content/50">
{{ currentRaw?.model_name }} · {{ currentRaw?.protocol }}
</p>
<!-- 标签页切换:请求 / 响应,避免上下堆叠,手机友好 -->
<div v-if="hasAnyRaw" class="mt-4">
<div class="tabs tabs-boxed w-fit max-w-full overflow-x-auto">
<button
v-if="currentRaw?.raw_request"
type="button"
class="tab tab-sm"
:class="rawTab === 'request' && 'tab-active'"
@click="rawTab = 'request'"
>请求</button>
<button
v-if="currentRaw?.raw_response"
type="button"
class="tab tab-sm"
:class="rawTab === 'response' && 'tab-active'"
@click="rawTab = 'response'"
>响应</button>
</div>
<pre class="mt-3 max-h-[55vh] overflow-auto rounded-lg bg-base-200/50 p-3 text-xs leading-relaxed whitespace-pre-wrap break-words">{{ activeRawContent }}</pre>
</div>
<p v-else class="mt-4 text-sm text-base-content/50">该请求未记录原始数据(仅管理员且开关开启时记录)。</p>
</div>
<form method="dialog" class="modal-backdrop">
<button aria-label="Close">close</button>
</form>
</dialog>
</div>
</template>
<script setup lang="ts">
import { ref, reactive, onMounted, computed } from 'vue'
import BreadcrumbHeader from '@/components/dashboard/BreadcrumbHeader.vue'
import Pagination from '@/components/common/Pagination.vue'
import { useUsageStore } from '@/stores/usage'
import type { UsageLogItem } from '@/types'
import { SearchIcon, RotateCcwIcon, FileTextIcon } from '@lucide/vue'
const store = useUsageStore()
const page = ref(1)
const pageSize = ref(20)
const filters = reactive({ protocol: '', status: '', model: '', userId: '' })
const summary = computed(() => store.adminSummary?.totals)
// 原始数据弹窗
const rawModal = ref<HTMLDialogElement | null>(null)
const currentRaw = ref<UsageLogItem | null>(null)
const rawTab = ref<'request' | 'response'>('request')
const hasAnyRaw = computed(() => !!currentRaw.value?.raw_request || !!currentRaw.value?.raw_response)
const activeRawContent = computed(() => {
const item = currentRaw.value
if (!item) return ''
return rawTab.value === 'request' ? item.raw_request ?? '' : item.raw_response ?? ''
})
function viewRaw(item: UsageLogItem) {
currentRaw.value = item
// 默认停在第一个有内容的标签(请求优先)
rawTab.value = item.raw_request ? 'request' : 'response'
rawModal.value?.showModal()
}
async function loadLogs() {
const params: Record<string, any> = { page: page.value, pageSize: pageSize.value }
if (filters.protocol) params.protocol = filters.protocol
if (filters.status) params.status = filters.status
if (filters.model) params.model = filters.model
if (filters.userId) params.user_id = filters.userId
try {
await store.fetchAdminLogs(params)
} catch { /* store 已抛错 */ }
}
async function loadSummary() {
try {
await store.fetchAdminSummary()
} catch { /* 同上 */ }
}
function applyFilters() {
page.value = 1
loadLogs()
}
function resetFilters() {
filters.protocol = ''
filters.status = ''
filters.model = ''
filters.userId = ''
applyFilters()
}
function changePage(p: number, s: number) {
page.value = p
pageSize.value = s
loadLogs()
}
function fmtInt(n?: number): string {
return (n ?? 0).toLocaleString()
}
function fmtCost(n?: number): string {
return `$${(n ?? 0).toFixed(6)}`
}
function fmtTime(t?: string): string {
if (!t) return '—'
const d = new Date(t)
const pad = (x: number) => String(x).padStart(2, '0')
return `${d.getFullYear()}-${pad(d.getMonth() + 1)}-${pad(d.getDate())} ${pad(d.getHours())}:${pad(d.getMinutes())}`
}
function statusLabel(s: string): string {
switch (s) {
case 'success': return '成功'
case 'error': return '失败'
case 'canceled': return '已取消'
default: return s
}
}
function statusClass(s: string): string {
switch (s) {
case 'success': return 'badge-success badge-soft'
case 'error': return 'badge-error badge-soft'
case 'canceled': return 'badge-warning badge-soft'
default: return 'badge-ghost'
}
}
onMounted(() => {
loadLogs()
loadSummary()
})
</script>
+361
View File
@@ -0,0 +1,361 @@
<template>
<div class="space-y-5">
<BreadcrumbHeader />
<div v-if="store.loading && !store.monthly" class="py-16 text-center text-sm text-base-content/50">
加载中…
</div>
<template v-else>
<!-- 年份切换 + 选中月份概览卡片 -->
<div class="flex items-center justify-between">
<div class="flex items-center gap-1">
<button class="btn btn-ghost btn-square btn-sm" aria-label="上一年" :disabled="year <= 2000" @click="switchYear(-1)">
<ChevronLeft class="size-4" aria-hidden="true" />
</button>
<span class="min-w-16 text-center text-lg font-semibold tabular-nums">{{ year }}</span>
<button class="btn btn-ghost btn-square btn-sm" aria-label="下一年" :disabled="year >= currentYear" @click="switchYear(1)">
<ChevronRight class="size-4" aria-hidden="true" />
</button>
</div>
<span class="text-xs text-base-content/50">{{ selectedMonthLabel }}用量概览</span>
</div>
<div class="grid grid-cols-1 gap-4 sm:grid-cols-3">
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">消费金额 (USD)</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtCost(selectedMonth?.cost) }}</p>
</div>
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">调用次数</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtInt(selectedMonth?.requests) }}</p>
</div>
<div class="card border border-base-300/60 bg-base-100 p-4 shadow-sm">
<p class="text-xs text-base-content/50">Token 消耗</p>
<p class="mt-1 text-2xl font-semibold tabular-nums">{{ fmtInt(monthTokens(selectedMonth)) }}</p>
<p class="mt-0.5 text-xs tabular-nums text-base-content/50">
输入 {{ fmtCompact(selectedMonth?.input_tokens) }} · 输出 {{ fmtCompact(selectedMonth?.output_tokens) }} · 缓存 {{
fmtCompact(selectedMonth?.cache_read_tokens) }}
</p>
</div>
</div>
<!-- 月度汇总图表:三种指标均按模型分色堆叠 -->
<div class="card border border-base-300/60 bg-base-100 p-5 shadow-sm">
<div class="mb-4 flex flex-wrap items-center justify-between gap-2">
<h3 class="text-sm font-semibold">月度汇总</h3>
<div class="flex items-center gap-3">
<div class="join">
<button v-for="opt in METRICS" :key="opt.key" class="btn btn-xs join-item"
:class="metric === opt.key ? 'btn-primary' : 'btn-ghost border-base-300/60'" @click="metric = opt.key">
{{ opt.label }}
</button>
</div>
<span v-if="maxMetricValue > 0" class="text-xs text-base-content/40">峰值 {{ fmtMetricValue(maxMetricValue) }}</span>
</div>
</div>
<div v-if="maxMetricValue > 0" class="flex h-44 items-end gap-1.5 sm:gap-3">
<div v-for="(m, i) in months" :key="m.month"
class="group flex h-full min-w-0 flex-1 cursor-pointer flex-col items-center justify-end gap-1"
:title="barTitle(m)" @click="selectedMonthIndex = i">
<!-- 柱顶总量 -->
<span class="text-[9px] leading-none tabular-nums text-base-content/40"
:class="{ 'font-semibold text-base-content/70': i === selectedMonthIndex }">
{{ metricValue(m) > 0 ? fmtMetricValue(metricValue(m)) : '' }}
</span>
<!-- 堆叠柱体:图例顺序堆叠,用量最大的模型在底部 -->
<div class="flex w-full max-w-10 flex-col-reverse overflow-hidden rounded-t transition-opacity"
:class="i === selectedMonthIndex ? 'opacity-100 ring-2 ring-primary/60' : 'opacity-80 group-hover:opacity-100'"
:style="{ height: barHeightPct(m) }">
<div v-for="seg in barSegments(m)" :key="seg.name" class="w-full"
:style="{ height: seg.pct + '%', backgroundColor: seg.color }"
:title="`${seg.name}: ${fmtMetricValue(seg.value)}(${seg.share}%)`">
</div>
</div>
<span class="text-[10px] leading-none tabular-nums"
:class="i === selectedMonthIndex ? 'font-semibold text-primary' : 'text-base-content/50'">
{{ i + 1 }}月
</span>
</div>
</div>
<div v-else class="py-10 text-center text-sm text-base-content/50">{{ year }} 年暂无用量数据</div>
<!-- 图例 -->
<div v-if="legend.length" class="mt-4 flex flex-wrap items-center gap-x-4 gap-y-1.5">
<span v-for="item in legend" :key="item.name" class="flex items-center gap-1.5 text-xs text-base-content/70"
:title="`${item.name}:全年 ${fmtMetricValue(item.value)}`">
<span class="size-2.5 rounded-sm" :style="{ backgroundColor: item.color }" aria-hidden="true"></span>
<span class="max-w-40 truncate">{{ item.name }}</span>
<span class="tabular-nums text-base-content/40">{{ fmtMetricValue(item.value) }}</span>
</span>
</div>
</div>
<!-- 请求明细 -->
<div class="card border border-base-300/60 bg-base-100 shadow-sm">
<div class="flex items-center justify-between px-5 pt-4">
<h3 class="text-sm font-semibold">请求明细</h3>
</div>
<div v-if="store.myLogs.length" class="overflow-x-auto">
<table class="table table-sm">
<thead>
<tr class="text-xs uppercase tracking-wider text-base-content/50">
<th>时间</th>
<th>模型</th>
<th>协议</th>
<th class="text-right">输入</th>
<th class="text-right">输出</th>
<th class="text-right">缓存</th>
<th class="text-right">费用</th>
<th>状态</th>
<th class="text-right">延迟</th>
</tr>
</thead>
<tbody>
<tr v-for="l in store.myLogs" :key="l.id" class="border-base-300/40 hover:bg-base-200/50">
<td class="whitespace-nowrap tabular-nums text-base-content/70">{{ fmtTime(l.created_at) }}</td>
<td class="max-w-40 truncate font-medium">{{ l.model_name }}</td>
<td><span class="badge badge-ghost badge-sm">{{ l.protocol }}</span></td>
<td class="text-right tabular-nums">{{ fmtInt(l.input_tokens) }}</td>
<td class="text-right tabular-nums">{{ fmtInt(l.output_tokens) }}</td>
<td class="text-right tabular-nums">{{ fmtInt(l.cache_read_tokens) }}</td>
<td class="text-right tabular-nums">{{ fmtCost(l.cost) }}</td>
<td>
<span class="badge badge-sm" :class="statusClass(l.status)">
{{ statusLabel(l.status) }}
</span>
</td>
<td class="text-right tabular-nums text-base-content/70">{{ l.latency_ms }}ms</td>
</tr>
</tbody>
</table>
</div>
<div v-else class="px-5 py-12 text-center text-sm text-base-content/50">暂无请求记录</div>
<Pagination v-if="myLogsTotal > 0" :current-page="page" :total-items="myLogsTotal" :page-size="pageSize"
:page-size-options="[10, 20, 50]" @change-page="changePage" />
</div>
</template>
</div>
</template>
<script setup lang="ts">
import { ref, computed, onMounted } from 'vue'
import { ChevronLeft, ChevronRight } from '@lucide/vue'
import BreadcrumbHeader from '@/components/dashboard/BreadcrumbHeader.vue'
import Pagination from '@/components/common/Pagination.vue'
import { useUsageStore } from '@/stores/usage'
import type { MonthlyUsage, MonthlyModelUsage } from '@/types'
const store = useUsageStore()
const currentYear = new Date().getFullYear()
const year = ref(currentYear)
const selectedMonthIndex = ref(new Date().getMonth())
const months = computed<MonthlyUsage[]>(() => {
const data = store.monthly
if (data && data.year === year.value) return data.months
// 数据未就绪/年份不匹配时给出 12 个月空骨架,保持布局稳定
return Array.from({ length: 12 }, (_, i) => ({
month: `${year.value}-${String(i + 1).padStart(2, '0')}`,
requests: 0, input_tokens: 0, output_tokens: 0, cache_read_tokens: 0, cost: 0, models: [],
}))
})
const selectedMonth = computed(() => months.value[selectedMonthIndex.value])
const selectedMonthLabel = computed(() => `${year.value} 年 ${selectedMonthIndex.value + 1} 月`)
function monthTokens(m?: MonthlyUsage): number {
if (!m) return 0
return m.input_tokens + m.output_tokens + m.cache_read_tokens
}
// --- 月度汇总图表:指标切换 + 按模型分色堆叠 ---
type MetricKey = 'tokens' | 'cost' | 'requests'
const METRICS: { key: MetricKey; label: string }[] = [
{ key: 'tokens', label: 'Token' },
{ key: 'cost', label: '消费金额' },
{ key: 'requests', label: '调用次数' },
]
const metric = ref<MetricKey>('tokens')
const PALETTE = [
'#6366f1', '#0ea5e9', '#10b981', '#f59e0b', '#ef4444', '#8b5cf6',
'#14b8a6', '#f97316', '#3b82f6', '#ec4899', '#84cc16', '#eab308',
]
const OTHER_COLOR = '#94a3b8'
const MAX_LEGEND = 8 // 图例最多展示 8 个模型,其余归入「其他」
const OTHER_NAME = '其他'
function monthTokensOf(mm: MonthlyModelUsage): number {
return mm.input_tokens + mm.output_tokens + mm.cache_read_tokens
}
// 当前指标下的数值(柱高、图例、峰值共用)
function metricValue(m: MonthlyUsage): number {
switch (metric.value) {
case 'cost': return m.cost
case 'requests': return m.requests
default: return monthTokens(m)
}
}
function metricValueOf(mm: MonthlyModelUsage): number {
switch (metric.value) {
case 'cost': return mm.cost
case 'requests': return mm.requests
default: return monthTokensOf(mm)
}
}
const maxMetricValue = computed(() => Math.max(0, ...months.value.map(metricValue)))
// 全年维度统计每个模型在当前指标下的总量,取前 MAX_LEGEND 个进入图例
const legend = computed(() => {
const totals = new Map<string, number>()
for (const m of months.value) {
for (const mm of m.models) {
totals.set(mm.model_name, (totals.get(mm.model_name) ?? 0) + metricValueOf(mm))
}
}
const sorted = [...totals.entries()].sort((a, b) => b[1] - a[1])
const top = sorted.slice(0, MAX_LEGEND).map(([name, value], i) => ({
name, value, color: PALETTE[i % PALETTE.length],
}))
const restValue = sorted.slice(MAX_LEGEND).reduce((s, [, v]) => s + v, 0)
if (restValue > 0) top.push({ name: OTHER_NAME, value: restValue, color: OTHER_COLOR })
return top
})
const legendIndex = computed(() => {
const idx = new Map<string, number>()
legend.value.forEach((item, i) => idx.set(item.name, i))
return idx
})
// 单月柱体:按图例顺序堆叠(保持各月颜色顺序一致),未进图例的模型归入「其他」
function barSegments(m: MonthlyUsage) {
const total = metricValue(m)
if (total === 0) return []
const byName = new Map<string, number>()
for (const mm of m.models) byName.set(mm.model_name, metricValueOf(mm))
const segs: { name: string; value: number; pct: number; share: number; color: string }[] = []
let other = 0
for (const [name, value] of byName) {
if (legendIndex.value.has(name)) continue
other += value
}
for (const item of legend.value) {
const value = item.name === OTHER_NAME ? other : (byName.get(item.name) ?? 0)
if (value <= 0) continue
const pct = (value / total) * 100
segs.push({ name: item.name, value, pct, share: Math.round(pct), color: item.color })
}
return segs
}
function barHeightPct(m: MonthlyUsage): string {
if (maxMetricValue.value === 0) return '0%'
return `${(metricValue(m) / maxMetricValue.value) * 100}%`
}
function barTitle(m: MonthlyUsage): string {
if (monthTokens(m) === 0 && m.requests === 0) return `${m.month}:无用量`
const parts = m.models
.slice()
.sort((a, b) => metricValueOf(b) - metricValueOf(a))
.map(mm => `${mm.model_name} ${fmtMetricValue(metricValueOf(mm))}`)
return `${m.month}:${fmtInt(m.requests)} 次调用,${fmtInt(monthTokens(m))} tokens,${fmtCost(m.cost)}\n${parts.join('\n')}`
}
const switchYear = (delta: number) => {
const next = year.value + delta
if (next < 2000 || next > currentYear) return
year.value = next
selectedMonthIndex.value = next === currentYear ? new Date().getMonth() : 11
loadMonthly()
}
// --- 请求明细(保留原有功能) ---
const page = ref(1)
const pageSize = ref(20)
const myLogsTotal = computed(() => store.myLogsTotal)
async function loadMonthly() {
try {
await store.fetchMonthly(year.value)
} catch { /* toast 由 store 抛错,页面保持静默 */ }
}
async function loadLogs() {
try {
await store.fetchMyLogs(pageSize.value, page.value)
} catch { /* 同上 */ }
}
const changePage = (p: number, s: number) => {
page.value = p
pageSize.value = s
loadLogs()
}
function fmtInt(n?: number): string {
return (n ?? 0).toLocaleString()
}
function fmtCost(n?: number): string {
return `$${(n ?? 0).toFixed(4)}`
}
// 金额紧凑格式(图表柱顶/图例使用)
function fmtMoney(v: number): string {
if (v >= 1e6) return '$' + (v / 1e6).toFixed(2) + 'M'
if (v >= 1e3) return '$' + (v / 1e3).toFixed(2) + 'k'
if (v >= 1) return '$' + v.toFixed(2)
return '$' + v.toFixed(4)
}
// 当前指标数值格式化
function fmtMetricValue(v: number): string {
switch (metric.value) {
case 'cost': return fmtMoney(v)
case 'requests': return fmtCompact(v)
default: return fmtCompact(v)
}
}
// 紧凑数字:柱顶/图例等小空间使用
function fmtCompact(n?: number): string {
const v = n ?? 0
if (v >= 1e9) return (v / 1e9).toFixed(1) + 'B'
if (v >= 1e6) return (v / 1e6).toFixed(1) + 'M'
if (v >= 1e3) return (v / 1e3).toFixed(1) + 'k'
return String(v)
}
function fmtTime(t?: string): string {
if (!t) return '—'
const d = new Date(t)
const pad = (x: number) => String(x).padStart(2, '0')
return `${d.getFullYear()}-${pad(d.getMonth() + 1)}-${pad(d.getDate())} ${pad(d.getHours())}:${pad(d.getMinutes())}`
}
function statusLabel(s: string): string {
switch (s) {
case 'success': return '成功'
case 'error': return '失败'
case 'canceled': return '已取消'
default: return s
}
}
function statusClass(s: string): string {
switch (s) {
case 'success': return 'badge-success badge-soft'
case 'error': return 'badge-error badge-soft'
case 'canceled': return 'badge-warning badge-soft'
default: return 'badge-ghost'
}
}
onMounted(() => {
loadMonthly()
loadLogs()
})
</script>
+1 -1
View File
@@ -9,7 +9,7 @@ import path from 'path'
// 需要自签名 HTTPS 时设置 VITE_DEV_HTTPS=true // 需要自签名 HTTPS 时设置 VITE_DEV_HTTPS=true
const useHttps = process.env.VITE_DEV_HTTPS === 'true' const useHttps = process.env.VITE_DEV_HTTPS === 'true'
// 后端地址:默认 make dev-backend 启动的 8080,可用 VITE_DEV_API_TARGET 覆盖 // 后端地址:默认 make dev-backend 启动的 8080,可用 VITE_DEV_API_TARGET 覆盖
const apiTarget = process.env.VITE_DEV_API_TARGET || 'http://localhost:8080' const apiTarget = process.env.VITE_DEV_API_TARGET || 'http://localhost:3000'
// https://vite.dev/config/ // https://vite.dev/config/
export default defineConfig({ export default defineConfig({