Merge remote-tracking branch 'upstream/main' into fix/issue-5843-antigravity-mixed-tools

# Conflicts:
#	backend/internal/repository/group_usage_rollup_trigger_integration_test.go
This commit is contained in:
wucm667
2026-08-28 17:36:43 +08:00
655 changed files with 59827 additions and 4350 deletions
+3
View File
@@ -49,6 +49,9 @@ coverage/
.env.*
!.env.example
# 本地闭源插件目录可能包含发布签名私钥,绝不能进入 Docker 构建上下文。
/plugins/
# Local config
config.yaml
config.local.yaml
+7
View File
@@ -35,3 +35,10 @@ exceptions:
mitigation: "Proxy configuration not user-controlled; upgrade when axios releases fix"
expires_on: "2026-07-10"
owner: "security@your-domain"
- package: nanoid
advisory: "GHSA-2v37-7h3g-55p8"
severity: high
reason: "Custom generator with size=0 not used; all nanoid calls use default size or positive values"
mitigation: "No zero-size custom generators in codebase; upgrade when nanoid releases fix"
expires_on: "2026-10-13"
owner: "security@your-domain"
+4 -3
View File
@@ -17,6 +17,7 @@ jobs:
/bin/bash -n deploy/apple-container.sh
/bin/bash deploy/tests/apple-container-test.sh
/bin/sh deploy/tests/docker-compose-security-test.sh
/bin/sh deploy/tests/docker-compose-gateway-env-test.sh
/bin/sh deploy/tests/docker-runtime-resources-test.sh
/bin/sh deploy/test-caddyfile-cache.sh
@@ -32,7 +33,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
go version | grep -q 'go1.26.6'
go version | grep -q 'go1.27.0'
- name: Unit tests
working-directory: backend
run: make test-unit
@@ -72,10 +73,10 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
go version | grep -q 'go1.26.6'
go version | grep -q 'go1.27.0'
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
version: v2.9
version: v2.13
args: --timeout=30m
working-directory: backend
+1 -1
View File
@@ -115,7 +115,7 @@ jobs:
- name: Verify Go version
run: |
go version | grep -q 'go1.26.6'
go version | grep -q 'go1.27.0'
# Docker setup for GoReleaser
- name: Set up QEMU
+1 -1
View File
@@ -23,7 +23,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
go version | grep -q 'go1.26.6'
go version | grep -q 'go1.27.0'
- name: Run govulncheck
working-directory: backend
run: |
+12 -1
View File
@@ -1,4 +1,5 @@
docs/claude-relay-service/
# 本地杂项文档、测试数据和外部项目副本
/docs-local/
.codex
# ===================
@@ -136,13 +137,23 @@ docs/*
!docs/PAYMENT_CN.md
!docs/ADMIN_PAYMENT_INTEGRATION_API.md
!docs/ASYNC_IMAGE_TASKS.md
!docs/BATCH_IMAGE_MVP.md
!docs/COMPOSITE_GROUPS.md
!docs/PLUGIN_DEVELOPMENT.md
!docs/channel-monitor-v2-safe-defaults.md
!docs/legal/
!docs/legal/*.md
!docs/screenshots/
docs/screenshots/*
!docs/screenshots/mobile-account-actions-menu.png
.serena/
.codex/
frontend/coverage/
aicodex
output/
.codegraph/
# Vitest / Vite cache at repo root
.vite/
# 本地闭源插件开发目录及构建产物不进入 Sub2API 仓库
/plugins/
+5 -5
View File
@@ -34,8 +34,8 @@
### 开发工具
```bash
# golangci-lint(CI 用 v2.9,本地建议装同一版以免版本差异带来的噪音)
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.9
# golangci-lint(CI 用 v2.13,本地建议装同一版以免版本差异带来的噪音)
go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.13
# pnpm (前端包管理)
npm install -g pnpm
@@ -47,13 +47,13 @@ npm install -g pnpm
| Workflow | 触发条件 | 检查内容 |
|----------|----------|----------|
| **backend-ci.yml** | push, pull_request | 单元测试 + 集成测试 + golangci-lint v2.9 |
| **backend-ci.yml** | push, pull_request | 单元测试 + 集成测试 + golangci-lint v2.13 |
| **security-scan.yml** | push, pull_request, 每周一 | govulncheck + gosec + pnpm audit |
| **release.yml** | tag `v*` | 构建发布(PR 不触发) |
### CI 要求
- Go 版本必须是 **1.26.6**:三个 workflow 都用 `go-version-file: backend/go.mod` 取版本,随后硬断言 `go version | grep -q 'go1.26.6'`。升级 Go 时要同时改 `backend/go.mod`、`backend-ci.yml`(两处)、`release.yml`、`security-scan.yml` 里的这句断言,**以及三个 Dockerfile 里的 Go 构建镜像**(`Dockerfile` / `deploy/Dockerfile` 的 `ARG GOLANG_IMAGE`、`backend/Dockerfile` 的 `FROM golang:`)。前者漏了 CI 会在版本校验步骤直接失败;**后者漏了 CI 不会报,而是等到有人用这些 Dockerfile 构建时才失败**(`go.mod requires go >= X (running Y; GOTOOLCHAIN=local)`)。
- Go 版本必须是 **1.27.0**:三个 workflow 都用 `go-version-file: backend/go.mod` 取版本,随后硬断言 `go version | grep -q 'go1.27.0'`。升级 Go 时要同时改 `backend/go.mod`、`backend-ci.yml`(两处)、`release.yml`、`security-scan.yml` 里的这句断言,**以及三个 Dockerfile 里的 Go 构建镜像**(`Dockerfile` / `deploy/Dockerfile` 的 `ARG GOLANG_IMAGE`、`backend/Dockerfile` 的 `FROM golang:`)。前者漏了 CI 会在版本校验步骤直接失败;**后者漏了 CI 不会报,而是等到有人用这些 Dockerfile 构建时才失败**(`go.mod requires go >= X (running Y; GOTOOLCHAIN=local)`)。
- 前端使用 `pnpm install --frozen-lockfile`,必须提交 `pnpm-lock.yaml`
### 本地测试命令
@@ -203,7 +203,7 @@ go test -tags=integration ./...
**解决**:
```bash
cd backend
go generate ./ent # 重新生成 ent 代码
go generate ./ent # 重新生成 ent 代码(json.RawMessage 字段会生成为同类型的 jsontext.Value,属预期)
git add ent/ # 生成的文件也要提交
```
+1 -1
View File
@@ -8,7 +8,7 @@
# =============================================================================
ARG NODE_IMAGE=node:24-alpine
ARG GOLANG_IMAGE=golang:1.26.6-alpine
ARG GOLANG_IMAGE=golang:1.27.0-alpine
ARG ALPINE_IMAGE=alpine:3.21
ARG POSTGRES_IMAGE=postgres:18-alpine
ARG GOPROXY=https://goproxy.cn,direct
+39 -15
View File
@@ -4,7 +4,7 @@
# Sub2API
[![Go](https://img.shields.io/badge/Go-1.26.5-00ADD8.svg)](https://golang.org/)
[![Go](https://img.shields.io/badge/Go-1.27.0-00ADD8.svg)](https://golang.org/)
[![Vue](https://img.shields.io/badge/Vue-3.4+-4FC08D.svg)](https://vuejs.org/)
[![PostgreSQL](https://img.shields.io/badge/PostgreSQL-15+-336791.svg)](https://www.postgresql.org/)
[![Redis](https://img.shields.io/badge/Redis-7+-DC382D.svg)](https://redis.io/)
@@ -58,6 +58,11 @@ Please read the following carefully before using this project:
<td>Thanks to AIGoCode for sponsoring this project! AIGoCode is an all-in-one platform that integrates Claude Code, Codex, and the latest Gemini models, providing you with stable, efficient, and highly cost-effective AI coding services. The platform offers flexible subscription plans, zero risk of account suspension, direct access with no VPN required, and lightning-fast responses. AIGoCode has prepared a special benefit for sub2api users: if you register via <a href="https://aigocode.com/invite/SUB2API">this link</a>, you'll receive an extra 10% bonus credit on your first top-up!</td>
</tr>
<tr>
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
<td>Real GPT-5.6 series at 3% of OpenAI pricing — <a href="https://codex-everywhere.com">CodexEverywhere</a> is democratizing access to frontier models for developers worldwide. We believe in transparency and honesty, with model quality verified by active community oversight for months. USD and crypto friendly. Start with a free $20 trial at <a href="https://codex-everywhere.com">codex-everywhere.com</a>.</td>
</tr>
<tr>
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
<td>Huge thanks to BmoPlus for sponsoring this project! BmoPlus is a highly reliable AI account provider built strictly for heavy AI users and developers. They offer rock-solid, ready-to-use accounts and official top-up services for ChatGPT Plus / ChatGPT Pro (Full Warranty) / Claude Pro / Super Grok / Gemini Pro. By registering and ordering through <a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus - Premium AI Accounts & Top-ups</a>, users can unlock the mind-blowing rate of 10% of the official GPT subscription price (90% OFF)</td>
@@ -117,18 +122,6 @@ Please read the following carefully before using this project:
</td>
</tr>
<tr>
<td width="180"><a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z"><img src="assets/partners/logos/claudeapi.jpg" alt="claudeapi" width="150"></a></td>
<td>Thanks to Claude API for sponsoring this project! <a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">Claude API</a> is an official-channel API provider focused on Claude models. Built on official Anthropic keys and the official AWS Bedrock channel, it delivers a stable integration experience for Claude Code and Agent applications, supports the full Claude model lineup, and retains official capabilities such as Tool Use and long context. The service involves no reverse engineering and no model degradation, making it a great fit for heavy Claude Code users, Agent engineers, and enterprise engineering teams. Register via the <a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">exclusive link</a> and contact customer support to receive free trial credits; invoicing and team onboarding are also supported.
</td>
</tr>
<tr>
<td width="180"><a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default"><img src="assets/partners/logos/code0.jpg" alt="code0" width="150"></a></td>
<td>Thanks to code0.ai for sponsoring this project! <a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">code0.ai</a> is an AI coding workbench for developers and engineering teams, aggregating mainstream agent coding capabilities such as Claude Code and Codex, and covering common development scenarios including code generation, project understanding, debugging and fixing, code review, and documentation generation. It suits independent developers, Agent engineers, open-source maintainers, and enterprise R&D teams, with invoicing and team onboarding supported. Register via the <a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">exclusive link</a> and contact customer support to receive free trial credits and experience a more efficient AI coding workflow.
</td>
</tr>
<tr>
<td width="180"><a href="https://nagora.ai/"><img src="assets/partners/logos/nagora.png" alt="Nagora" width="150"></a></td>
<td><a href="https://nagora.ai/">Nagora</a> is a multi-model AI API gateway built for developers and teams. With a single account and API key, you can access more than 26 leading text and image models through one unified interface. It is compatible with OpenAI, Anthropic, and Gemini protocols and integrates seamlessly with development tools such as Claude Code, Codex, and Gemini CLI. The platform provides intelligent routing, automatic failover, transparent pricing, and consolidated billing, along with budget management, rate limiting, and concurrency controls. This makes AI usage more reliable and manageable across individual development, team collaboration, and production environments. No changes to your existing application are required. Simply replace the Base URL and API key to complete the integration in as little as one minute.</td>
@@ -174,6 +167,11 @@ Please read the following carefully before using this project:
<td><a href="https://www.duckip.cn/?keyword=cu7oog6y">DuckIP</a> - 90M+ global residential network resources across 195+ countries and regions, with rotation and sticky sessions for public data collection, RAG updates, model evaluation, and multi-region data workloads. 🟢Residential Proxy - 20% Off; 🟢Static Residential Proxy - Starting at ¥50.00/IP; 🟢Unlimited Residential Proxy - Starting at ¥19.8/Hour. ✅Get 500M Free Trial.</td>
</tr>
<tr>
<td width="180"><a href="https://go.apimart.ai/gh-sub2api"><img src="assets/partners/logos/apimart.jpg" alt="APIMart" width="150"></a></td>
<td>Thanks to APIMart for sponsoring this project! <a href="https://go.apimart.ai/gh-sub2api">APIMart</a> is a low-cost API platform for AI image and video generation — GPT-Image-2 from $0.006 per image, with 160+ images per dollar. One async API covers both image and video: submit a task, get an ID, and retrieve results via polling or callback. Batch tens of thousands of images without timeouts, and switch models without changing code. Pay as you go with no monthly fee — <a href="https://go.apimart.ai/gh-sub2api">sign up here</a> to get started.</td>
</tr>
</table>
## Overview
@@ -206,7 +204,7 @@ Community projects that extend or integrate with Sub2API:
| Component | Technology |
|-----------|------------|
| Backend | Go 1.26.5, Gin, Ent |
| Backend | Go 1.27.0, Gin, Ent |
| Frontend | Vue 3.4+, Vite 5+, TailwindCSS |
| Database | PostgreSQL 15+ |
| Cache/Queue | Redis 7+ |
@@ -652,6 +650,32 @@ Or set `GATEWAY_OPENAI_WS_MODE_ROUTER_V2_ENABLED=true` in the environment.
Use `http_bridge` for client-WebSocket/upstream-HTTP operation when rolling out
or mitigating upstream WebSocket issues.
#### Force OpenAI upstream HTTP/SSE
When an egress proxy or network repeatedly reconnects OpenAI Responses
WebSockets, set the global fallback in the persisted deployment configuration:
```yaml
gateway:
openai_ws:
force_http: true
```
For Compose and Apple container deployments, the equivalent `.env` setting is:
```bash
GATEWAY_OPENAI_WS_FORCE_HTTP=true
```
This selects HTTP/SSE for OpenAI upstream Responses traffic that would
otherwise use WebSocket. It does not change the client-facing protocol or force
HTTP/1.1; configure `gateway.openai_http2.enabled` (or
`GATEWAY_OPENAI_HTTP2_ENABLED=false`) separately when a proxy is incompatible
with HTTP/2. Unlike the account-level `http_bridge` mode, this global fallback
takes effect without enabling `mode_router_v2_enabled`. Keep the setting in the
deployment's persisted `.env` or `config.yaml`, rather than inside a running
container, so it is read again after an image update or container recreation.
#### ⚠️ Important: Creating the Admin Account
The initial admin account is **only created via the setup wizard** (served at `http://<host>:8080` on first run). The `default.admin_email` / `default.admin_password` fields in `config.yaml` are **not used** to create it — they exist in the template for historical reasons.
@@ -730,7 +754,7 @@ Sub2API supports both Grok subscription accounts through xAI OAuth and standard
- Codex CLI style Responses WebSocket ingress is accepted on the Responses targets and bridged to xAI HTTP/SSE Responses upstream
- Text models: `grok-4.5`, `grok-4.3`, `grok-build-0.1`, `grok-composer-2.5-fast`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309`
- Media targets for Grok groups: `/v1/images/generations`, `/images/generations`, `/v1/images/edits`, `/images/edits`, `/v1/videos/generations`, `/videos/generations`, `/v1/videos/edits`, `/videos/edits`, `/v1/videos/extensions`, `/videos/extensions`, `/v1/videos/{request_id}`, and `/videos/{request_id}`. Generation, editing, and extension requests require the group image-generation permission.
- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5`
- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-image-2.0`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5`
- JSON image-edit and video-generation requests accept image references in `image`, `images`, `reference_images`, and `mask` objects. Use `url` for xAI-compatible payloads; the legacy `image_url` field remains accepted and is normalized to `url` before forwarding.
- Out of scope for this provider: TTS, transcription, browser automation, cookies, and Grok web scraping
+12 -14
View File
@@ -4,7 +4,7 @@
# Sub2API
[![Go](https://img.shields.io/badge/Go-1.26.5-00ADD8.svg)](https://golang.org/)
[![Go](https://img.shields.io/badge/Go-1.27.0-00ADD8.svg)](https://golang.org/)
[![Vue](https://img.shields.io/badge/Vue-3.4+-4FC08D.svg)](https://vuejs.org/)
[![PostgreSQL](https://img.shields.io/badge/PostgreSQL-15+-336791.svg)](https://www.postgresql.org/)
[![Redis](https://img.shields.io/badge/Redis-7+-DC382D.svg)](https://redis.io/)
@@ -59,6 +59,11 @@
<td>感谢 AIGoCode 赞助了本项目!AIGoCode 是一站式集成 Claude Code、Codex 以及最新 Gemini 模型的综合平台,为您提供稳定、高效、高性价比的 AI 编程服务。平台提供灵活的订阅方案,零封号风险,免 VPN 直连,响应极速。AIGoCode 为 sub2api 用户准备了专属福利:通过<a href="https://aigocode.com/invite/SUB2API">此链接</a>注册,首次充值可额外获得 10% 赠送额度!</td>
</tr>
<tr>
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
<td>Real GPT-5.6 series at 3% of OpenAI pricing — <a href="https://codex-everywhere.com">CodexEverywhere</a> is democratizing access to frontier models for developers worldwide. We believe in transparency and honesty, with model quality verified by active community oversight for months. USD and crypto friendly. Start with a free $20 trial at <a href="https://codex-everywhere.com">codex-everywhere.com</a>.</td>
</tr>
<tr>
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
<td>感谢 BmoPlus 赞助了本项目!BmoPlus 是一家专为AI订阅重度用户打造的可靠 AI 账号代充服务商,提供稳定的 ChatGPT Plus / ChatGPT Pro(全程质保) / Claude Pro / Super Grok / Gemini Pro 的官方代充&成品账号。 通过<a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus AI成品号专卖/代充</a>注册下单的用户,可享GPT 官网订阅一折 的震撼价格!</td>
@@ -120,18 +125,6 @@
</td>
</tr>
<tr>
<td width="180"><a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z"><img src="assets/partners/logos/claudeapi.jpg" alt="claudeapi" width="150"></a></td>
<td>感谢 Claude API 对本项目的赞助! <a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">Claude API</a> 是专注 Claude 模型的官方渠道 API 服务商,基于 Anthropic 官方 Key 与 AWS Bedrock 官方渠道,提供稳定的 Claude Code 与 Agent 应用接入体验,支持 Claude 全系列模型,保留 Tool Use、长上下文等官方能力。服务非逆向、非降智,适合 Claude Code 深度用户、Agent 工程师与企业技术团队使用。通过<a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">[专属链接]</a>注册后联系客服,可领取免费测试额度,并支持开票和团队对接。
</td>
</tr>
<tr>
<td width="180"><a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default"><img src="assets/partners/logos/code0.jpg" alt="code0" width="150"></a></td>
<td>感谢 code0.ai 对本项目的赞助! <a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">code0.ai</a> 是面向开发者与技术团队的 AI 编程工作台,聚合 Claude Code、Codex 等主流 Agent 编程能力,支持代码生成、项目理解、调试修复、代码审查与文档生成等常见研发场景。适合独立开发者、Agent 工程师、开源项目维护者和企业研发团队使用,支持开票和团队对接。通过<a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">[专属链接]</a>注册后联系客服,可领取免费测试额度,体验更高效的 AI 编程工作流。
</td>
</tr>
<tr>
<td width="180"><a href="https://nagora.ai/"><img src="assets/partners/logos/nagora.png" alt="Nagora" width="150"></a></td>
<td><a href="https://nagora.ai/">Nagora</a> 是专为开发者和团队打造的多模型 AI API 网关。通过一个账户和一枚 API Key,即可统一调用 26+ 款主流文本与图像模型,兼容 OpenAI、Anthropic 与 Gemini 协议,并可无缝接入 Claude Code、Codex、Gemini CLI 等开发工具。平台提供智能路由、自动故障转移、透明计费与统一账单,同时支持预算、限速、并发控制,让个人开发、团队协作和生产环境中的 AI 调用更稳定、更可控。无需改造现有应用,只需替换 Base URL 与 API Key,最快 1 分钟即可完成接入。</td>
@@ -177,6 +170,11 @@
<td><a href="https://www.duckip.cn/?keyword=cu7oog6y">DuckIP</a> - 9000 万+ 全球住宅网络资源,覆盖 195+ 国家和地区,支持轮换和粘性会话,适用于公共数据采集、RAG 更新、模型评估和多区域数据工作负载。🟢住宅代理 - 8 折优惠;🟢静态住宅代理 - ¥50.00/IP 起;🟢无限住宅代理 - ¥19.8/小时 起。✅免费领取 500M 试用流量。</td>
</tr>
<tr>
<td width="180"><a href="https://go.apimart.ai/gh-sub2api"><img src="assets/partners/logos/apimart.jpg" alt="APIMart" width="150"></a></td>
<td>感谢 APIMart 赞助了本项目!<a href="https://go.apimart.ai/gh-sub2api">APIMart</a> 是专注于 AI 图片/视频生成的低价 API 平台,GPT-Image-2 低至 $0.006/张,1 美元可生成 160+ 张图片。图片、视频一套异步 API 通吃:提交任务获取 ID,通过轮询或回调获取结果;批量生成上万张图片也不会超时,切换模型无需修改代码。按量付费、无月费,通过<a href="https://go.apimart.ai/gh-sub2api">此注册链接</a>注册即可开始使用。</td>
</tr>
</table>
## 项目概述
@@ -208,7 +206,7 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅的
| 组件 | 技术 |
|------|------|
| 后端 | Go 1.26.5, Gin, Ent |
| 后端 | Go 1.27.0, Gin, Ent |
| 前端 | Vue 3.4+, Vite 5+, TailwindCSS |
| 数据库 | PostgreSQL 15+ |
| 缓存/队列 | Redis 7+ |
+12 -14
View File
@@ -4,7 +4,7 @@
# Sub2API
[![Go](https://img.shields.io/badge/Go-1.26.5-00ADD8.svg)](https://golang.org/)
[![Go](https://img.shields.io/badge/Go-1.27.0-00ADD8.svg)](https://golang.org/)
[![Vue](https://img.shields.io/badge/Vue-3.4+-4FC08D.svg)](https://vuejs.org/)
[![PostgreSQL](https://img.shields.io/badge/PostgreSQL-15+-336791.svg)](https://www.postgresql.org/)
[![Redis](https://img.shields.io/badge/Redis-7+-DC382D.svg)](https://redis.io/)
@@ -58,6 +58,11 @@
<td>AIGoCode のご支援に感謝します!AIGoCode は Claude Code、Codex、最新の Gemini モデルを統合したオールインワンプラットフォームで、安定的かつ効率的でコストパフォーマンスに優れた AI コーディングサービスを提供します。柔軟なサブスクリプションプラン、アカウント停止リスクゼロ、VPN 不要の直接アクセス、超高速レスポンスが特長です。AIGoCode は sub2api ユーザー向けに特別特典を用意しています:<a href="https://aigocode.com/invite/SUB2API">こちらのリンク</a>から登録すると、初回チャージ時に 10% のボーナスクレジットを追加プレゼント!</td>
</tr>
<tr>
<td width="180"><a href="https://codex-everywhere.com"><img src="assets/partners/logos/codex-everywhere.jpg" alt="CodexEverywhere" width="150"></a></td>
<td>OpenAI 公式価格のわずか 3% で本物の GPT-5.6 シリーズを提供 — <a href="https://codex-everywhere.com">CodexEverywhere</a> は世界中の開発者にフロンティアモデルへのアクセスを民主化しています。私たちは透明性と誠実さを信条とし、モデル品質は数か月にわたるアクティブなコミュニティの監視によって検証されています。USD および暗号通貨に対応。<a href="https://codex-everywhere.com">codex-everywhere.com</a> で $20 の無料トライアルから始めましょう。</td>
</tr>
<tr>
<td width="180"><a href="https://shop.bmoplus.com/?utm_source=github"><img src="assets/partners/logos/bmoplus.jpg" alt="bmoplus" width="150"></a></td>
<td>本プロジェクトにご支援いただいた BmoPlus に感謝いたします!BmoPlusは、AIサブスクリプションのヘビーユーザー向けに特化した信頼性の高いAIアカウントサービスプロバイダーであり、安定した ChatGPT Plus / ChatGPT Pro (完全保証) / Claude Pro / Super Grok / Gemini Pro の公式代行チャージおよび即納アカウントを提供しています。こちらの<a href="https://shop.bmoplus.com/?utm_source=github">BmoPlus AIアカウント専門店/代行チャージ</a>経由でご登録・ご注文いただいたユーザー様は、GPTを 公式サイト価格の約1割(90% OFF) という驚異的な価格でご利用いただけます!</td>
@@ -119,18 +124,6 @@
</td>
</tr>
<tr>
<td width="180"><a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z"><img src="assets/partners/logos/claudeapi.jpg" alt="claudeapi" width="150"></a></td>
<td>Claude API のご支援に感謝します!<a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">Claude API</a> は Claude モデルに特化した公式チャネルの API サービスプロバイダーで、Anthropic 公式キーと AWS Bedrock 公式チャネルをベースに、Claude Code や Agent アプリケーションへの安定した接続体験を提供します。Claude 全シリーズのモデルに対応し、Tool Use や長文コンテキストなどの公式機能もそのまま利用可能。リバースエンジニアリングやモデル劣化のないサービスで、Claude Code のヘビーユーザー、Agent エンジニア、企業の技術チームに最適です。<a href="https://console.claudeapi.com/agent/register/drTKjyn6wGLK061Z">専用リンク</a>から登録後カスタマーサポートへご連絡いただくと、無料お試しクレジットを受け取れます。請求書発行やチーム導入にも対応しています。
</td>
</tr>
<tr>
<td width="180"><a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default"><img src="assets/partners/logos/code0.jpg" alt="code0" width="150"></a></td>
<td>code0.ai のご支援に感謝します!<a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">code0.ai</a> は開発者と技術チーム向けの AI プログラミングワークベンチで、Claude Code や Codex などの主要な Agent コーディング能力を集約し、コード生成、プロジェクト理解、デバッグと修正、コードレビュー、ドキュメント生成といった一般的な開発シーンをサポートします。個人開発者、Agent エンジニア、OSS メンテナー、企業の開発チームに最適で、請求書発行やチーム導入にも対応。<a href="https://code0.ai/agent/register/LgpIgl9JHtVG53V1?utm_source=zcf&utm_medium=partner&utm_campaign=zcf_2026&utm_content=default">専用リンク</a>から登録後カスタマーサポートへご連絡いただくと、無料お試しクレジットを受け取り、より効率的な AI プログラミングワークフローを体験できます。
</td>
</tr>
<tr>
<td width="180"><a href="https://nagora.ai/"><img src="assets/partners/logos/nagora.png" alt="Nagora" width="150"></a></td>
<td><a href="https://nagora.ai/">Nagora</a>は、開発者やチーム向けに設計されたマルチモデルAI APIゲートウェイです。1つのアカウントと1つのAPIキーだけで、26種類以上の主要なテキストモデルおよび画像モデルを一元的に利用できます。OpenAI、Anthropic、Geminiの各プロトコルに対応し、Claude Code、Codex、Gemini CLIなどの開発ツールにもシームレスに接続できます。 プラットフォームには、インテリジェントルーティング、自動フェイルオーバー、透明性の高い料金体系、請求の一元管理に加え、予算管理、レート制限、同時実行数の制御機能が備わっています。これにより、個人開発、チームでの共同作業、本番環境におけるAI APIの利用を、より安定的かつ柔軟に管理できます。 既存のアプリケーションを改修する必要はありません。Base URLとAPIキーを置き換えるだけで、最短1分で導入を完了できます。</td>
@@ -176,6 +169,11 @@
<td><a href="https://www.duckip.cn/?keyword=cu7oog6y">DuckIP</a> - 195 以上の国と地域にわたる 9,000 万以上のグローバルレジデンシャルネットワークリソース。ローテーションとスティッキーセッションに対応し、パブリックデータ収集、RAG 更新、モデル評価、マルチリージョンデータワークロードに最適。🟢レジデンシャルプロキシ - 20% オフ;🟢スタティックレジデンシャルプロキシ - ¥50.00/IP から;🟢無制限レジデンシャルプロキシ - ¥19.8/時間 から。✅500M 無料トライアルを取得。</td>
</tr>
<tr>
<td width="180"><a href="https://go.apimart.ai/gh-sub2api"><img src="assets/partners/logos/apimart.jpg" alt="APIMart" width="150"></a></td>
<td>APIMart のご支援に感謝します!<a href="https://go.apimart.ai/gh-sub2api">APIMart</a> は AI 画像・動画生成に特化した低価格 API プラットフォームです。GPT-Image-2 は 1 枚 $0.006 から、1 ドルで 160 枚以上の画像を生成できます。画像と動画の両方に対応する非同期 API を 1 つで利用でき、タスクを送信して ID を取得し、ポーリングまたはコールバックで結果を取得できます。数万枚規模のバッチ処理でもタイムアウトせず、モデルを変更してもコードの変更は不要です。月額料金なしの従量課金制で、<a href="https://go.apimart.ai/gh-sub2api">こちらの登録リンク</a>からすぐに利用を開始できます。</td>
</tr>
</table>
## 概要
@@ -207,7 +205,7 @@ Sub2API を拡張・統合するコミュニティプロジェクト:
| コンポーネント | 技術 |
|-----------|------------|
| バックエンド | Go 1.26.5, Gin, Ent |
| バックエンド | Go 1.27.0, Gin, Ent |
| フロントエンド | Vue 3.4+, Vite 5+, TailwindCSS |
| データベース | PostgreSQL 15+ |
| キャッシュ/キュー | Redis 7+ |
Binary file not shown.

After

Width:  |  Height:  |  Size: 60 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 3.6 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 3.2 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 5.9 KiB

+10
View File
@@ -121,6 +121,16 @@ linters:
# Default: true — must be true, ent generates 130K+ lines of code
generated-is-used: true
exclusions:
rules:
# G703/G704 污点分析在测试文件中只会命中 httptest mock 与测试者自设的
# 环境变量路径,不构成攻击面;且该分析跨环境结果不稳定(本地/CI 报告的
# 位置集合不同),逐点 nolint 无法收敛,故按路径豁免。生产代码不豁免,
# 必须逐点 //nolint:gosec 并写明信任边界。
- path: '_test\.go$'
linters: [ gosec ]
text: 'G70[34]'
formatters:
enable:
- gofmt
+1 -1
View File
@@ -1,4 +1,4 @@
FROM golang:1.26.6-alpine
FROM golang:1.27.0-alpine
WORKDIR /app
+1 -1
View File
@@ -1 +1 @@
0.1.178
0.1.183
+5
View File
@@ -153,6 +153,11 @@ func runMainServer() {
log.Fatalf("Failed to initialize application: %v", err)
}
defer app.Cleanup()
if app.PluginManager != nil {
if err := app.PluginManager.Start(context.Background()); err != nil {
log.Printf("Plugin manager started in degraded state: %v", err)
}
}
if app.PromptAudit != nil {
if err := app.PromptAudit.Start(context.Background()); err != nil {
// Startup continues so unrelated APIs stay up. Fail-closed (unavailable)
+33 -10
View File
@@ -25,9 +25,10 @@ import (
)
type Application struct {
Server *http.Server
PromptAudit *securityaudit.PromptService
Cleanup func()
Server *http.Server
PromptAudit *securityaudit.PromptService
PluginManager *service.PluginManager
Cleanup func()
}
func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
@@ -51,12 +52,13 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
// BuildInfo provider
provideServiceBuildInfo,
providePluginHostInfo,
// Cleanup function provider
provideCleanup,
// Application struct
wire.Struct(new(Application), "Server", "PromptAudit", "Cleanup"),
wire.Struct(new(Application), "Server", "PromptAudit", "PluginManager", "Cleanup"),
)
return nil, nil
}
@@ -72,6 +74,13 @@ func provideServiceBuildInfo(buildInfo handler.BuildInfo) service.BuildInfo {
}
}
func providePluginHostInfo(buildInfo handler.BuildInfo) service.PluginHostInfo {
return service.PluginHostInfo{
Version: buildInfo.Version,
BuildType: buildInfo.BuildType,
}
}
func provideCleanup(
entClient *ent.Client,
rdb *redis.Client,
@@ -116,7 +125,9 @@ func provideCleanup(
upstreamBillingProbe *service.UpstreamBillingProbeService,
ollamaCloudUsage *service.OllamaCloudUsageService,
auditLog *service.AuditLogService,
openAIAutoReset *service.OpenAIQuotaAutoResetService,
promptAudit *securityaudit.PromptService,
pluginManager *service.PluginManager,
) func() {
return func() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
@@ -129,6 +140,18 @@ func provideCleanup(
// 应用层清理步骤可并行执行,基础设施资源(Redis/Ent)最后按顺序关闭。
parallelSteps := []cleanupStep{
{"PluginManager", func() error {
if pluginManager != nil {
pluginManager.Stop()
}
return nil
}},
{"OpenAIQuotaAutoResetService", func() error {
if openAIAutoReset != nil {
openAIAutoReset.Stop()
}
return nil
}},
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
@@ -328,12 +351,12 @@ func provideCleanup(
return nil
}},
{"ChannelMonitorV2Aggregator", func() error {
if channelMonitorV2Aggregator != nil {
channelMonitorV2Aggregator.Stop()
}
return nil
}},
{"ChannelMonitorRunner", func() error {
if channelMonitorV2Aggregator != nil {
channelMonitorV2Aggregator.Stop()
}
return nil
}},
{"ChannelMonitorRunner", func() error {
if channelMonitorRunner != nil {
channelMonitorRunner.Stop()
}
+42 -13
View File
@@ -194,13 +194,16 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository, configConfig)
grokQuotaFetcher := service.NewGrokQuotaFetcher()
grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, configConfig, usageLogRepository, settingService)
openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory, openAIGatewayService)
usageCache := service.NewUsageCache()
accountUsageService := service.ProvideAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService, openAIGatewayService)
accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService, settingService)
pluginRepository := repository.NewPluginRepository(db)
pluginHostInfo := providePluginHostInfo(buildInfo)
pluginManager := service.NewPluginManager(pluginRepository, secretEncryptor, configConfig, pluginHostInfo)
accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService, settingService, pluginManager)
crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig)
accountHandler := admin.ProvideAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, grokQuotaService)
adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService)
@@ -251,6 +254,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
errorPassthroughService := service.NewErrorPassthroughService(errorPassthroughRepository, errorPassthroughCache)
errorPassthroughHandler := admin.NewErrorPassthroughHandler(errorPassthroughService)
tlsFingerprintProfileHandler := admin.NewTLSFingerprintProfileHandler(tlsFingerprintProfileService)
pluginHandler := admin.NewPluginHandler(pluginManager)
adminAPIKeyHandler := admin.NewAdminAPIKeyHandler(adminService)
scheduledTestPlanRepository := repository.NewScheduledTestPlanRepository(db)
scheduledTestResultRepository := repository.NewScheduledTestResultRepository(db)
@@ -280,14 +284,14 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
auditLogHandler := admin.NewAuditLogHandler(auditLogService, totpService)
upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db)
ollamaCloudUsageService := service.ProvideOllamaCloudUsageService(accountRepository, httpUpstream, settingService, secretEncryptor, configConfig, leaderLockCache, db)
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, cnProviderHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService, ollamaCloudUsageService)
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, cnProviderHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, pluginHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService, ollamaCloudUsageService)
usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig)
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig)
legacyEngine := securityaudit.NewLegacyModerationAdapter(contentModerationService)
coordinator := securityaudit.NewCoordinator(legacyEngine, promptService)
gatewayHandler := handler.ProvideGatewayHandler(gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService, coordinator)
openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, grokQuotaService, configConfig, coordinator)
openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, pluginManager, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, grokQuotaService, configConfig, coordinator)
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService)
totpHandler := handler.NewTotpHandler(totpService)
passkeyRepository := repository.NewPasskeyRepository(db)
@@ -300,7 +304,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService)
paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry)
availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService)
modelPlazaHandler := handler.NewModelPlazaHandler(channelService, apiKeyService, settingService)
modelPlazaService := service.NewModelPlazaService(channelRepository, groupRepository, pricingService, billingService, modelPricingResolver)
modelPlazaHandler := handler.NewModelPlazaHandler(modelPlazaService, apiKeyService, settingService)
imageTaskStore := repository.NewImageTaskStore(redisClient)
imageTaskService := service.ProvideImageTaskService(imageTaskStore, imageStorageSettingService)
asyncImageHandler := handler.NewAsyncImageHandler(imageTaskService, openAIGatewayHandler)
@@ -314,7 +319,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
batchImageHandler := handler.ProvideBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService, openAIGatewayHandler)
idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig)
idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, channelMonitorV2Handler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, passkeyHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, modelPlazaHandler, asyncImageHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService)
openAIQuotaAutoResetService := service.ProvideOpenAIQuotaAutoResetService(accountRepository, openAIQuotaService, rateLimitService, idempotencyCoordinator, auditLogService, settingService, leaderLockCache)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, channelMonitorV2Handler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, passkeyHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, modelPlazaHandler, asyncImageHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService, openAIQuotaAutoResetService)
jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService, settingService, auditLogService)
optionalJWTAuthMiddleware := middleware.NewOptionalJWTAuthMiddleware(authService, userService, settingService, auditLogService)
adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService, auditLogService)
@@ -341,11 +347,12 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService, channelMonitorQuotaFetcher)
channelMonitorV2Aggregator := service.ProvideChannelMonitorV2Aggregator(channelMonitorV2Repository, db, settingService)
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, cnProviderBalanceCheckService, openAICodexVersionSyncService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, channelMonitorV2Aggregator, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, ollamaCloudUsageService, auditLogService, promptService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, cnProviderBalanceCheckService, openAICodexVersionSyncService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, channelMonitorV2Aggregator, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, ollamaCloudUsageService, auditLogService, openAIQuotaAutoResetService, promptService, pluginManager)
application := &Application{
Server: httpServer,
PromptAudit: promptService,
Cleanup: v,
Server: httpServer,
PromptAudit: promptService,
PluginManager: pluginManager,
Cleanup: v,
}
return application, nil
}
@@ -353,9 +360,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
// wire.go:
type Application struct {
Server *http.Server
PromptAudit *securityaudit.PromptService
Cleanup func()
Server *http.Server
PromptAudit *securityaudit.PromptService
PluginManager *service.PluginManager
Cleanup func()
}
func providePrivacyClientFactory() service.PrivacyClientFactory {
@@ -369,6 +377,13 @@ func provideServiceBuildInfo(buildInfo handler.BuildInfo) service.BuildInfo {
}
}
func providePluginHostInfo(buildInfo handler.BuildInfo) service.PluginHostInfo {
return service.PluginHostInfo{
Version: buildInfo.Version,
BuildType: buildInfo.BuildType,
}
}
func provideCleanup(
entClient *ent.Client,
rdb *redis.Client,
@@ -413,7 +428,9 @@ func provideCleanup(
upstreamBillingProbe *service.UpstreamBillingProbeService,
ollamaCloudUsage *service.OllamaCloudUsageService,
auditLog *service.AuditLogService,
openAIAutoReset *service.OpenAIQuotaAutoResetService,
promptAudit *securityaudit.PromptService,
pluginManager *service.PluginManager,
) func() {
return func() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
@@ -425,6 +442,18 @@ func provideCleanup(
}
parallelSteps := []cleanupStep{
{"PluginManager", func() error {
if pluginManager != nil {
pluginManager.Stop()
}
return nil
}},
{"OpenAIQuotaAutoResetService", func() error {
if openAIAutoReset != nil {
openAIAutoReset.Stop()
}
return nil
}},
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
+2
View File
@@ -94,7 +94,9 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
nil, // upstreamBillingProbe
nil, // ollamaCloudUsage
nil, // auditLog
nil, // openAIAutoReset
nil, // promptAudit
nil, // pluginManager
)
require.NotPanics(t, func() {
+2 -1
View File
@@ -4,6 +4,7 @@ package ent
import (
"encoding/json"
"encoding/json/jsontext"
"fmt"
"strings"
"time"
@@ -100,7 +101,7 @@ type Group struct {
// 是否按上下文长度应用模型阶梯价格;默认开启以保持官方/渠道长上下文价
LongContextPricingEnabled bool `json:"long_context_pricing_enabled,omitempty"`
// 分组逐模型定价;优先级高于渠道和内置定价
ModelPricing json.RawMessage `json:"model_pricing,omitempty"`
ModelPricing jsontext.Value `json:"model_pricing,omitempty"`
// 是否仅允许 Claude Code 客户端
ClaudeCodeOnly bool `json:"claude_code_only,omitempty"`
// 非 Claude Code 请求降级使用的分组 ID
+5 -5
View File
@@ -4,7 +4,7 @@ package ent
import (
"context"
"encoding/json"
"encoding/json/jsontext"
"errors"
"fmt"
"time"
@@ -575,7 +575,7 @@ func (_c *GroupCreate) SetNillableLongContextPricingEnabled(v *bool) *GroupCreat
}
// SetModelPricing sets the "model_pricing" field.
func (_c *GroupCreate) SetModelPricing(v json.RawMessage) *GroupCreate {
func (_c *GroupCreate) SetModelPricing(v jsontext.Value) *GroupCreate {
_c.mutation.SetModelPricing(v)
return _c
}
@@ -2445,7 +2445,7 @@ func (u *GroupUpsert) UpdateLongContextPricingEnabled() *GroupUpsert {
}
// SetModelPricing sets the "model_pricing" field.
func (u *GroupUpsert) SetModelPricing(v json.RawMessage) *GroupUpsert {
func (u *GroupUpsert) SetModelPricing(v jsontext.Value) *GroupUpsert {
u.Set(group.FieldModelPricing, v)
return u
}
@@ -3615,7 +3615,7 @@ func (u *GroupUpsertOne) UpdateLongContextPricingEnabled() *GroupUpsertOne {
}
// SetModelPricing sets the "model_pricing" field.
func (u *GroupUpsertOne) SetModelPricing(v json.RawMessage) *GroupUpsertOne {
func (u *GroupUpsertOne) SetModelPricing(v jsontext.Value) *GroupUpsertOne {
return u.Update(func(s *GroupUpsert) {
s.SetModelPricing(v)
})
@@ -5005,7 +5005,7 @@ func (u *GroupUpsertBulk) UpdateLongContextPricingEnabled() *GroupUpsertBulk {
}
// SetModelPricing sets the "model_pricing" field.
func (u *GroupUpsertBulk) SetModelPricing(v json.RawMessage) *GroupUpsertBulk {
func (u *GroupUpsertBulk) SetModelPricing(v jsontext.Value) *GroupUpsertBulk {
return u.Update(func(s *GroupUpsert) {
s.SetModelPricing(v)
})
+5 -5
View File
@@ -4,7 +4,7 @@ package ent
import (
"context"
"encoding/json"
"encoding/json/jsontext"
"errors"
"fmt"
"time"
@@ -803,13 +803,13 @@ func (_u *GroupUpdate) SetNillableLongContextPricingEnabled(v *bool) *GroupUpdat
}
// SetModelPricing sets the "model_pricing" field.
func (_u *GroupUpdate) SetModelPricing(v json.RawMessage) *GroupUpdate {
func (_u *GroupUpdate) SetModelPricing(v jsontext.Value) *GroupUpdate {
_u.mutation.SetModelPricing(v)
return _u
}
// AppendModelPricing appends value to the "model_pricing" field.
func (_u *GroupUpdate) AppendModelPricing(v json.RawMessage) *GroupUpdate {
func (_u *GroupUpdate) AppendModelPricing(v jsontext.Value) *GroupUpdate {
_u.mutation.AppendModelPricing(v)
return _u
}
@@ -2924,13 +2924,13 @@ func (_u *GroupUpdateOne) SetNillableLongContextPricingEnabled(v *bool) *GroupUp
}
// SetModelPricing sets the "model_pricing" field.
func (_u *GroupUpdateOne) SetModelPricing(v json.RawMessage) *GroupUpdateOne {
func (_u *GroupUpdateOne) SetModelPricing(v jsontext.Value) *GroupUpdateOne {
_u.mutation.SetModelPricing(v)
return _u
}
// AppendModelPricing appends value to the "model_pricing" field.
func (_u *GroupUpdateOne) AppendModelPricing(v json.RawMessage) *GroupUpdateOne {
func (_u *GroupUpdateOne) AppendModelPricing(v jsontext.Value) *GroupUpdateOne {
_u.mutation.AppendModelPricing(v)
return _u
}
+1
View File
@@ -1801,6 +1801,7 @@ var (
{Name: "signup_source", Type: field.TypeString, Default: "email"},
{Name: "last_login_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "last_active_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "restrict_public_groups", Type: field.TypeBool, Default: false},
{Name: "balance_notify_enabled", Type: field.TypeBool, Default: true},
{Name: "balance_notify_threshold_type", Type: field.TypeString, Default: "fixed"},
{Name: "balance_notify_threshold", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}},
+78 -24
View File
@@ -4,7 +4,7 @@ package ent
import (
"context"
"encoding/json"
"encoding/json/jsontext"
"errors"
"fmt"
"sync"
@@ -22142,8 +22142,8 @@ type GroupMutation struct {
audio_stt_price_per_hour *float64
addaudio_stt_price_per_hour *float64
long_context_pricing_enabled *bool
model_pricing *json.RawMessage
appendmodel_pricing json.RawMessage
model_pricing *jsontext.Value
appendmodel_pricing jsontext.Value
claude_code_only *bool
fallback_group_id *int64
addfallback_group_id *int64
@@ -24404,13 +24404,13 @@ func (m *GroupMutation) ResetLongContextPricingEnabled() {
}
// SetModelPricing sets the "model_pricing" field.
func (m *GroupMutation) SetModelPricing(jm json.RawMessage) {
m.model_pricing = &jm
func (m *GroupMutation) SetModelPricing(j jsontext.Value) {
m.model_pricing = &j
m.appendmodel_pricing = nil
}
// ModelPricing returns the value of the "model_pricing" field in the mutation.
func (m *GroupMutation) ModelPricing() (r json.RawMessage, exists bool) {
func (m *GroupMutation) ModelPricing() (r jsontext.Value, exists bool) {
v := m.model_pricing
if v == nil {
return
@@ -24421,7 +24421,7 @@ func (m *GroupMutation) ModelPricing() (r json.RawMessage, exists bool) {
// OldModelPricing returns the old "model_pricing" field's value of the Group entity.
// If the Group object wasn't provided to the builder, the object is fetched from the database.
// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
func (m *GroupMutation) OldModelPricing(ctx context.Context) (v json.RawMessage, err error) {
func (m *GroupMutation) OldModelPricing(ctx context.Context) (v jsontext.Value, err error) {
if !m.op.Is(OpUpdateOne) {
return v, errors.New("OldModelPricing is only allowed on UpdateOne operations")
}
@@ -24435,13 +24435,13 @@ func (m *GroupMutation) OldModelPricing(ctx context.Context) (v json.RawMessage,
return oldValue.ModelPricing, nil
}
// AppendModelPricing adds jm to the "model_pricing" field.
func (m *GroupMutation) AppendModelPricing(jm json.RawMessage) {
m.appendmodel_pricing = append(m.appendmodel_pricing, jm...)
// AppendModelPricing adds j to the "model_pricing" field.
func (m *GroupMutation) AppendModelPricing(j jsontext.Value) {
m.appendmodel_pricing = append(m.appendmodel_pricing, j...)
}
// AppendedModelPricing returns the list of values that were appended to the "model_pricing" field in this mutation.
func (m *GroupMutation) AppendedModelPricing() (json.RawMessage, bool) {
func (m *GroupMutation) AppendedModelPricing() (jsontext.Value, bool) {
if len(m.appendmodel_pricing) == 0 {
return nil, false
}
@@ -26515,7 +26515,7 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error {
m.SetLongContextPricingEnabled(v)
return nil
case group.FieldModelPricing:
v, ok := value.(json.RawMessage)
v, ok := value.(jsontext.Value)
if !ok {
return fmt.Errorf("unexpected type %T for field %s", value, name)
}
@@ -43151,8 +43151,8 @@ type UsageCleanupTaskMutation struct {
created_at *time.Time
updated_at *time.Time
status *string
filters *json.RawMessage
appendfilters json.RawMessage
filters *jsontext.Value
appendfilters jsontext.Value
created_by *int64
addcreated_by *int64
deleted_rows *int64
@@ -43376,13 +43376,13 @@ func (m *UsageCleanupTaskMutation) ResetStatus() {
}
// SetFilters sets the "filters" field.
func (m *UsageCleanupTaskMutation) SetFilters(jm json.RawMessage) {
m.filters = &jm
func (m *UsageCleanupTaskMutation) SetFilters(j jsontext.Value) {
m.filters = &j
m.appendfilters = nil
}
// Filters returns the value of the "filters" field in the mutation.
func (m *UsageCleanupTaskMutation) Filters() (r json.RawMessage, exists bool) {
func (m *UsageCleanupTaskMutation) Filters() (r jsontext.Value, exists bool) {
v := m.filters
if v == nil {
return
@@ -43393,7 +43393,7 @@ func (m *UsageCleanupTaskMutation) Filters() (r json.RawMessage, exists bool) {
// OldFilters returns the old "filters" field's value of the UsageCleanupTask entity.
// If the UsageCleanupTask object wasn't provided to the builder, the object is fetched from the database.
// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
func (m *UsageCleanupTaskMutation) OldFilters(ctx context.Context) (v json.RawMessage, err error) {
func (m *UsageCleanupTaskMutation) OldFilters(ctx context.Context) (v jsontext.Value, err error) {
if !m.op.Is(OpUpdateOne) {
return v, errors.New("OldFilters is only allowed on UpdateOne operations")
}
@@ -43407,13 +43407,13 @@ func (m *UsageCleanupTaskMutation) OldFilters(ctx context.Context) (v json.RawMe
return oldValue.Filters, nil
}
// AppendFilters adds jm to the "filters" field.
func (m *UsageCleanupTaskMutation) AppendFilters(jm json.RawMessage) {
m.appendfilters = append(m.appendfilters, jm...)
// AppendFilters adds j to the "filters" field.
func (m *UsageCleanupTaskMutation) AppendFilters(j jsontext.Value) {
m.appendfilters = append(m.appendfilters, j...)
}
// AppendedFilters returns the list of values that were appended to the "filters" field in this mutation.
func (m *UsageCleanupTaskMutation) AppendedFilters() (json.RawMessage, bool) {
func (m *UsageCleanupTaskMutation) AppendedFilters() (jsontext.Value, bool) {
if len(m.appendfilters) == 0 {
return nil, false
}
@@ -43964,7 +43964,7 @@ func (m *UsageCleanupTaskMutation) SetField(name string, value ent.Value) error
m.SetStatus(v)
return nil
case usagecleanuptask.FieldFilters:
v, ok := value.(json.RawMessage)
v, ok := value.(jsontext.Value)
if !ok {
return fmt.Errorf("unexpected type %T for field %s", value, name)
}
@@ -48422,6 +48422,7 @@ type UserMutation struct {
signup_source *string
last_login_at *time.Time
last_active_at *time.Time
restrict_public_groups *bool
balance_notify_enabled *bool
balance_notify_threshold_type *string
balance_notify_threshold *float64
@@ -49347,6 +49348,42 @@ func (m *UserMutation) ResetLastActiveAt() {
delete(m.clearedFields, user.FieldLastActiveAt)
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (m *UserMutation) SetRestrictPublicGroups(b bool) {
m.restrict_public_groups = &b
}
// RestrictPublicGroups returns the value of the "restrict_public_groups" field in the mutation.
func (m *UserMutation) RestrictPublicGroups() (r bool, exists bool) {
v := m.restrict_public_groups
if v == nil {
return
}
return *v, true
}
// OldRestrictPublicGroups returns the old "restrict_public_groups" field's value of the User entity.
// If the User object wasn't provided to the builder, the object is fetched from the database.
// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
func (m *UserMutation) OldRestrictPublicGroups(ctx context.Context) (v bool, err error) {
if !m.op.Is(OpUpdateOne) {
return v, errors.New("OldRestrictPublicGroups is only allowed on UpdateOne operations")
}
if m.id == nil || m.oldValue == nil {
return v, errors.New("OldRestrictPublicGroups requires an ID field in the mutation")
}
oldValue, err := m.oldValue(ctx)
if err != nil {
return v, fmt.Errorf("querying old value for OldRestrictPublicGroups: %w", err)
}
return oldValue.RestrictPublicGroups, nil
}
// ResetRestrictPublicGroups resets all changes to the "restrict_public_groups" field.
func (m *UserMutation) ResetRestrictPublicGroups() {
m.restrict_public_groups = nil
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (m *UserMutation) SetBalanceNotifyEnabled(b bool) {
m.balance_notify_enabled = &b
@@ -50373,7 +50410,7 @@ func (m *UserMutation) Type() string {
// order to get all numeric fields that were incremented/decremented, call
// AddedFields().
func (m *UserMutation) Fields() []string {
fields := make([]string, 0, 24)
fields := make([]string, 0, 25)
if m.created_at != nil {
fields = append(fields, user.FieldCreatedAt)
}
@@ -50428,6 +50465,9 @@ func (m *UserMutation) Fields() []string {
if m.last_active_at != nil {
fields = append(fields, user.FieldLastActiveAt)
}
if m.restrict_public_groups != nil {
fields = append(fields, user.FieldRestrictPublicGroups)
}
if m.balance_notify_enabled != nil {
fields = append(fields, user.FieldBalanceNotifyEnabled)
}
@@ -50490,6 +50530,8 @@ func (m *UserMutation) Field(name string) (ent.Value, bool) {
return m.LastLoginAt()
case user.FieldLastActiveAt:
return m.LastActiveAt()
case user.FieldRestrictPublicGroups:
return m.RestrictPublicGroups()
case user.FieldBalanceNotifyEnabled:
return m.BalanceNotifyEnabled()
case user.FieldBalanceNotifyThresholdType:
@@ -50547,6 +50589,8 @@ func (m *UserMutation) OldField(ctx context.Context, name string) (ent.Value, er
return m.OldLastLoginAt(ctx)
case user.FieldLastActiveAt:
return m.OldLastActiveAt(ctx)
case user.FieldRestrictPublicGroups:
return m.OldRestrictPublicGroups(ctx)
case user.FieldBalanceNotifyEnabled:
return m.OldBalanceNotifyEnabled(ctx)
case user.FieldBalanceNotifyThresholdType:
@@ -50694,6 +50738,13 @@ func (m *UserMutation) SetField(name string, value ent.Value) error {
}
m.SetLastActiveAt(v)
return nil
case user.FieldRestrictPublicGroups:
v, ok := value.(bool)
if !ok {
return fmt.Errorf("unexpected type %T for field %s", value, name)
}
m.SetRestrictPublicGroups(v)
return nil
case user.FieldBalanceNotifyEnabled:
v, ok := value.(bool)
if !ok {
@@ -50953,6 +51004,9 @@ func (m *UserMutation) ResetField(name string) error {
case user.FieldLastActiveAt:
m.ResetLastActiveAt()
return nil
case user.FieldRestrictPublicGroups:
m.ResetRestrictPublicGroups()
return nil
case user.FieldBalanceNotifyEnabled:
m.ResetBalanceNotifyEnabled()
return nil
+9 -5
View File
@@ -2217,24 +2217,28 @@ func init() {
user.DefaultSignupSource = userDescSignupSource.Default.(string)
// user.SignupSourceValidator is a validator for the "signup_source" field. It is called by the builders before save.
user.SignupSourceValidator = userDescSignupSource.Validators[0].(func(string) error)
// userDescRestrictPublicGroups is the schema descriptor for restrict_public_groups field.
userDescRestrictPublicGroups := userFields[15].Descriptor()
// user.DefaultRestrictPublicGroups holds the default value on creation for the restrict_public_groups field.
user.DefaultRestrictPublicGroups = userDescRestrictPublicGroups.Default.(bool)
// userDescBalanceNotifyEnabled is the schema descriptor for balance_notify_enabled field.
userDescBalanceNotifyEnabled := userFields[15].Descriptor()
userDescBalanceNotifyEnabled := userFields[16].Descriptor()
// user.DefaultBalanceNotifyEnabled holds the default value on creation for the balance_notify_enabled field.
user.DefaultBalanceNotifyEnabled = userDescBalanceNotifyEnabled.Default.(bool)
// userDescBalanceNotifyThresholdType is the schema descriptor for balance_notify_threshold_type field.
userDescBalanceNotifyThresholdType := userFields[16].Descriptor()
userDescBalanceNotifyThresholdType := userFields[17].Descriptor()
// user.DefaultBalanceNotifyThresholdType holds the default value on creation for the balance_notify_threshold_type field.
user.DefaultBalanceNotifyThresholdType = userDescBalanceNotifyThresholdType.Default.(string)
// userDescBalanceNotifyExtraEmails is the schema descriptor for balance_notify_extra_emails field.
userDescBalanceNotifyExtraEmails := userFields[18].Descriptor()
userDescBalanceNotifyExtraEmails := userFields[19].Descriptor()
// user.DefaultBalanceNotifyExtraEmails holds the default value on creation for the balance_notify_extra_emails field.
user.DefaultBalanceNotifyExtraEmails = userDescBalanceNotifyExtraEmails.Default.(string)
// userDescTotalRecharged is the schema descriptor for total_recharged field.
userDescTotalRecharged := userFields[19].Descriptor()
userDescTotalRecharged := userFields[20].Descriptor()
// user.DefaultTotalRecharged holds the default value on creation for the total_recharged field.
user.DefaultTotalRecharged = userDescTotalRecharged.Default.(float64)
// userDescRpmLimit is the schema descriptor for rpm_limit field.
userDescRpmLimit := userFields[20].Descriptor()
userDescRpmLimit := userFields[21].Descriptor()
// user.DefaultRpmLimit holds the default value on creation for the rpm_limit field.
user.DefaultRpmLimit = userDescRpmLimit.Default.(int)
userallowedgroupFields := schema.UserAllowedGroup{}.Fields()
+5
View File
@@ -96,6 +96,11 @@ func (User) Fields() []ent.Field {
Nillable().
SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
// 公开分组访问限制:为 false 时用户可绑定任意非专属分组(默认行为),
// 为 true 时仅可绑定 user_allowed_groups 中列出的公开分组。
field.Bool("restrict_public_groups").
Default(false),
// 余额不足通知
field.Bool("balance_notify_enabled").
Default(true),
+2 -1
View File
@@ -4,6 +4,7 @@ package ent
import (
"encoding/json"
"encoding/json/jsontext"
"fmt"
"strings"
"time"
@@ -25,7 +26,7 @@ type UsageCleanupTask struct {
// Status holds the value of the "status" field.
Status string `json:"status,omitempty"`
// Filters holds the value of the "filters" field.
Filters json.RawMessage `json:"filters,omitempty"`
Filters jsontext.Value `json:"filters,omitempty"`
// CreatedBy holds the value of the "created_by" field.
CreatedBy int64 `json:"created_by,omitempty"`
// DeletedRows holds the value of the "deleted_rows" field.
+5 -5
View File
@@ -4,7 +4,7 @@ package ent
import (
"context"
"encoding/json"
"encoding/json/jsontext"
"errors"
"fmt"
"time"
@@ -58,7 +58,7 @@ func (_c *UsageCleanupTaskCreate) SetStatus(v string) *UsageCleanupTaskCreate {
}
// SetFilters sets the "filters" field.
func (_c *UsageCleanupTaskCreate) SetFilters(v json.RawMessage) *UsageCleanupTaskCreate {
func (_c *UsageCleanupTaskCreate) SetFilters(v jsontext.Value) *UsageCleanupTaskCreate {
_c.mutation.SetFilters(v)
return _c
}
@@ -375,7 +375,7 @@ func (u *UsageCleanupTaskUpsert) UpdateStatus() *UsageCleanupTaskUpsert {
}
// SetFilters sets the "filters" field.
func (u *UsageCleanupTaskUpsert) SetFilters(v json.RawMessage) *UsageCleanupTaskUpsert {
func (u *UsageCleanupTaskUpsert) SetFilters(v jsontext.Value) *UsageCleanupTaskUpsert {
u.Set(usagecleanuptask.FieldFilters, v)
return u
}
@@ -592,7 +592,7 @@ func (u *UsageCleanupTaskUpsertOne) UpdateStatus() *UsageCleanupTaskUpsertOne {
}
// SetFilters sets the "filters" field.
func (u *UsageCleanupTaskUpsertOne) SetFilters(v json.RawMessage) *UsageCleanupTaskUpsertOne {
func (u *UsageCleanupTaskUpsertOne) SetFilters(v jsontext.Value) *UsageCleanupTaskUpsertOne {
return u.Update(func(s *UsageCleanupTaskUpsert) {
s.SetFilters(v)
})
@@ -999,7 +999,7 @@ func (u *UsageCleanupTaskUpsertBulk) UpdateStatus() *UsageCleanupTaskUpsertBulk
}
// SetFilters sets the "filters" field.
func (u *UsageCleanupTaskUpsertBulk) SetFilters(v json.RawMessage) *UsageCleanupTaskUpsertBulk {
func (u *UsageCleanupTaskUpsertBulk) SetFilters(v jsontext.Value) *UsageCleanupTaskUpsertBulk {
return u.Update(func(s *UsageCleanupTaskUpsert) {
s.SetFilters(v)
})
+5 -5
View File
@@ -4,7 +4,7 @@ package ent
import (
"context"
"encoding/json"
"encoding/json/jsontext"
"errors"
"fmt"
"time"
@@ -51,13 +51,13 @@ func (_u *UsageCleanupTaskUpdate) SetNillableStatus(v *string) *UsageCleanupTask
}
// SetFilters sets the "filters" field.
func (_u *UsageCleanupTaskUpdate) SetFilters(v json.RawMessage) *UsageCleanupTaskUpdate {
func (_u *UsageCleanupTaskUpdate) SetFilters(v jsontext.Value) *UsageCleanupTaskUpdate {
_u.mutation.SetFilters(v)
return _u
}
// AppendFilters appends value to the "filters" field.
func (_u *UsageCleanupTaskUpdate) AppendFilters(v json.RawMessage) *UsageCleanupTaskUpdate {
func (_u *UsageCleanupTaskUpdate) AppendFilters(v jsontext.Value) *UsageCleanupTaskUpdate {
_u.mutation.AppendFilters(v)
return _u
}
@@ -374,13 +374,13 @@ func (_u *UsageCleanupTaskUpdateOne) SetNillableStatus(v *string) *UsageCleanupT
}
// SetFilters sets the "filters" field.
func (_u *UsageCleanupTaskUpdateOne) SetFilters(v json.RawMessage) *UsageCleanupTaskUpdateOne {
func (_u *UsageCleanupTaskUpdateOne) SetFilters(v jsontext.Value) *UsageCleanupTaskUpdateOne {
_u.mutation.SetFilters(v)
return _u
}
// AppendFilters appends value to the "filters" field.
func (_u *UsageCleanupTaskUpdateOne) AppendFilters(v json.RawMessage) *UsageCleanupTaskUpdateOne {
func (_u *UsageCleanupTaskUpdateOne) AppendFilters(v jsontext.Value) *UsageCleanupTaskUpdateOne {
_u.mutation.AppendFilters(v)
return _u
}
+12 -1
View File
@@ -53,6 +53,8 @@ type User struct {
LastLoginAt *time.Time `json:"last_login_at,omitempty"`
// LastActiveAt holds the value of the "last_active_at" field.
LastActiveAt *time.Time `json:"last_active_at,omitempty"`
// RestrictPublicGroups holds the value of the "restrict_public_groups" field.
RestrictPublicGroups bool `json:"restrict_public_groups,omitempty"`
// BalanceNotifyEnabled holds the value of the "balance_notify_enabled" field.
BalanceNotifyEnabled bool `json:"balance_notify_enabled,omitempty"`
// BalanceNotifyThresholdType holds the value of the "balance_notify_threshold_type" field.
@@ -237,7 +239,7 @@ func (*User) scanValues(columns []string) ([]any, error) {
values := make([]any, len(columns))
for i := range columns {
switch columns[i] {
case user.FieldTotpEnabled, user.FieldBalanceNotifyEnabled:
case user.FieldTotpEnabled, user.FieldRestrictPublicGroups, user.FieldBalanceNotifyEnabled:
values[i] = new(sql.NullBool)
case user.FieldBalance, user.FieldFrozenBalance, user.FieldBalanceNotifyThreshold, user.FieldTotalRecharged:
values[i] = new(sql.NullFloat64)
@@ -381,6 +383,12 @@ func (_m *User) assignValues(columns []string, values []any) error {
_m.LastActiveAt = new(time.Time)
*_m.LastActiveAt = value.Time
}
case user.FieldRestrictPublicGroups:
if value, ok := values[i].(*sql.NullBool); !ok {
return fmt.Errorf("unexpected type %T for field restrict_public_groups", values[i])
} else if value.Valid {
_m.RestrictPublicGroups = value.Bool
}
case user.FieldBalanceNotifyEnabled:
if value, ok := values[i].(*sql.NullBool); !ok {
return fmt.Errorf("unexpected type %T for field balance_notify_enabled", values[i])
@@ -588,6 +596,9 @@ func (_m *User) String() string {
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
builder.WriteString("restrict_public_groups=")
builder.WriteString(fmt.Sprintf("%v", _m.RestrictPublicGroups))
builder.WriteString(", ")
builder.WriteString("balance_notify_enabled=")
builder.WriteString(fmt.Sprintf("%v", _m.BalanceNotifyEnabled))
builder.WriteString(", ")
+10
View File
@@ -51,6 +51,8 @@ const (
FieldLastLoginAt = "last_login_at"
// FieldLastActiveAt holds the string denoting the last_active_at field in the database.
FieldLastActiveAt = "last_active_at"
// FieldRestrictPublicGroups holds the string denoting the restrict_public_groups field in the database.
FieldRestrictPublicGroups = "restrict_public_groups"
// FieldBalanceNotifyEnabled holds the string denoting the balance_notify_enabled field in the database.
FieldBalanceNotifyEnabled = "balance_notify_enabled"
// FieldBalanceNotifyThresholdType holds the string denoting the balance_notify_threshold_type field in the database.
@@ -212,6 +214,7 @@ var Columns = []string{
FieldSignupSource,
FieldLastLoginAt,
FieldLastActiveAt,
FieldRestrictPublicGroups,
FieldBalanceNotifyEnabled,
FieldBalanceNotifyThresholdType,
FieldBalanceNotifyThreshold,
@@ -280,6 +283,8 @@ var (
DefaultSignupSource string
// SignupSourceValidator is a validator for the "signup_source" field. It is called by the builders before save.
SignupSourceValidator func(string) error
// DefaultRestrictPublicGroups holds the default value on creation for the "restrict_public_groups" field.
DefaultRestrictPublicGroups bool
// DefaultBalanceNotifyEnabled holds the default value on creation for the "balance_notify_enabled" field.
DefaultBalanceNotifyEnabled bool
// DefaultBalanceNotifyThresholdType holds the default value on creation for the "balance_notify_threshold_type" field.
@@ -390,6 +395,11 @@ func ByLastActiveAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldLastActiveAt, opts...).ToFunc()
}
// ByRestrictPublicGroups orders the results by the restrict_public_groups field.
func ByRestrictPublicGroups(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldRestrictPublicGroups, opts...).ToFunc()
}
// ByBalanceNotifyEnabled orders the results by the balance_notify_enabled field.
func ByBalanceNotifyEnabled(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldBalanceNotifyEnabled, opts...).ToFunc()
+15
View File
@@ -145,6 +145,11 @@ func LastActiveAt(v time.Time) predicate.User {
return predicate.User(sql.FieldEQ(FieldLastActiveAt, v))
}
// RestrictPublicGroups applies equality check predicate on the "restrict_public_groups" field. It's identical to RestrictPublicGroupsEQ.
func RestrictPublicGroups(v bool) predicate.User {
return predicate.User(sql.FieldEQ(FieldRestrictPublicGroups, v))
}
// BalanceNotifyEnabled applies equality check predicate on the "balance_notify_enabled" field. It's identical to BalanceNotifyEnabledEQ.
func BalanceNotifyEnabled(v bool) predicate.User {
return predicate.User(sql.FieldEQ(FieldBalanceNotifyEnabled, v))
@@ -1115,6 +1120,16 @@ func LastActiveAtNotNil() predicate.User {
return predicate.User(sql.FieldNotNull(FieldLastActiveAt))
}
// RestrictPublicGroupsEQ applies the EQ predicate on the "restrict_public_groups" field.
func RestrictPublicGroupsEQ(v bool) predicate.User {
return predicate.User(sql.FieldEQ(FieldRestrictPublicGroups, v))
}
// RestrictPublicGroupsNEQ applies the NEQ predicate on the "restrict_public_groups" field.
func RestrictPublicGroupsNEQ(v bool) predicate.User {
return predicate.User(sql.FieldNEQ(FieldRestrictPublicGroups, v))
}
// BalanceNotifyEnabledEQ applies the EQ predicate on the "balance_notify_enabled" field.
func BalanceNotifyEnabledEQ(v bool) predicate.User {
return predicate.User(sql.FieldEQ(FieldBalanceNotifyEnabled, v))
+65
View File
@@ -270,6 +270,20 @@ func (_c *UserCreate) SetNillableLastActiveAt(v *time.Time) *UserCreate {
return _c
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (_c *UserCreate) SetRestrictPublicGroups(v bool) *UserCreate {
_c.mutation.SetRestrictPublicGroups(v)
return _c
}
// SetNillableRestrictPublicGroups sets the "restrict_public_groups" field if the given value is not nil.
func (_c *UserCreate) SetNillableRestrictPublicGroups(v *bool) *UserCreate {
if v != nil {
_c.SetRestrictPublicGroups(*v)
}
return _c
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (_c *UserCreate) SetBalanceNotifyEnabled(v bool) *UserCreate {
_c.mutation.SetBalanceNotifyEnabled(v)
@@ -636,6 +650,10 @@ func (_c *UserCreate) defaults() error {
v := user.DefaultSignupSource
_c.mutation.SetSignupSource(v)
}
if _, ok := _c.mutation.RestrictPublicGroups(); !ok {
v := user.DefaultRestrictPublicGroups
_c.mutation.SetRestrictPublicGroups(v)
}
if _, ok := _c.mutation.BalanceNotifyEnabled(); !ok {
v := user.DefaultBalanceNotifyEnabled
_c.mutation.SetBalanceNotifyEnabled(v)
@@ -730,6 +748,9 @@ func (_c *UserCreate) check() error {
return &ValidationError{Name: "signup_source", err: fmt.Errorf(`ent: validator failed for field "User.signup_source": %w`, err)}
}
}
if _, ok := _c.mutation.RestrictPublicGroups(); !ok {
return &ValidationError{Name: "restrict_public_groups", err: errors.New(`ent: missing required field "User.restrict_public_groups"`)}
}
if _, ok := _c.mutation.BalanceNotifyEnabled(); !ok {
return &ValidationError{Name: "balance_notify_enabled", err: errors.New(`ent: missing required field "User.balance_notify_enabled"`)}
}
@@ -844,6 +865,10 @@ func (_c *UserCreate) createSpec() (*User, *sqlgraph.CreateSpec) {
_spec.SetField(user.FieldLastActiveAt, field.TypeTime, value)
_node.LastActiveAt = &value
}
if value, ok := _c.mutation.RestrictPublicGroups(); ok {
_spec.SetField(user.FieldRestrictPublicGroups, field.TypeBool, value)
_node.RestrictPublicGroups = value
}
if value, ok := _c.mutation.BalanceNotifyEnabled(); ok {
_spec.SetField(user.FieldBalanceNotifyEnabled, field.TypeBool, value)
_node.BalanceNotifyEnabled = value
@@ -1384,6 +1409,18 @@ func (u *UserUpsert) ClearLastActiveAt() *UserUpsert {
return u
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (u *UserUpsert) SetRestrictPublicGroups(v bool) *UserUpsert {
u.Set(user.FieldRestrictPublicGroups, v)
return u
}
// UpdateRestrictPublicGroups sets the "restrict_public_groups" field to the value that was provided on create.
func (u *UserUpsert) UpdateRestrictPublicGroups() *UserUpsert {
u.SetExcluded(user.FieldRestrictPublicGroups)
return u
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (u *UserUpsert) SetBalanceNotifyEnabled(v bool) *UserUpsert {
u.Set(user.FieldBalanceNotifyEnabled, v)
@@ -1819,6 +1856,20 @@ func (u *UserUpsertOne) ClearLastActiveAt() *UserUpsertOne {
})
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (u *UserUpsertOne) SetRestrictPublicGroups(v bool) *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
s.SetRestrictPublicGroups(v)
})
}
// UpdateRestrictPublicGroups sets the "restrict_public_groups" field to the value that was provided on create.
func (u *UserUpsertOne) UpdateRestrictPublicGroups() *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
s.UpdateRestrictPublicGroups()
})
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (u *UserUpsertOne) SetBalanceNotifyEnabled(v bool) *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
@@ -2436,6 +2487,20 @@ func (u *UserUpsertBulk) ClearLastActiveAt() *UserUpsertBulk {
})
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (u *UserUpsertBulk) SetRestrictPublicGroups(v bool) *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
s.SetRestrictPublicGroups(v)
})
}
// UpdateRestrictPublicGroups sets the "restrict_public_groups" field to the value that was provided on create.
func (u *UserUpsertBulk) UpdateRestrictPublicGroups() *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
s.UpdateRestrictPublicGroups()
})
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (u *UserUpsertBulk) SetBalanceNotifyEnabled(v bool) *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
+34
View File
@@ -321,6 +321,20 @@ func (_u *UserUpdate) ClearLastActiveAt() *UserUpdate {
return _u
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (_u *UserUpdate) SetRestrictPublicGroups(v bool) *UserUpdate {
_u.mutation.SetRestrictPublicGroups(v)
return _u
}
// SetNillableRestrictPublicGroups sets the "restrict_public_groups" field if the given value is not nil.
func (_u *UserUpdate) SetNillableRestrictPublicGroups(v *bool) *UserUpdate {
if v != nil {
_u.SetRestrictPublicGroups(*v)
}
return _u
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (_u *UserUpdate) SetBalanceNotifyEnabled(v bool) *UserUpdate {
_u.mutation.SetBalanceNotifyEnabled(v)
@@ -1069,6 +1083,9 @@ func (_u *UserUpdate) sqlSave(ctx context.Context) (_node int, err error) {
if _u.mutation.LastActiveAtCleared() {
_spec.ClearField(user.FieldLastActiveAt, field.TypeTime)
}
if value, ok := _u.mutation.RestrictPublicGroups(); ok {
_spec.SetField(user.FieldRestrictPublicGroups, field.TypeBool, value)
}
if value, ok := _u.mutation.BalanceNotifyEnabled(); ok {
_spec.SetField(user.FieldBalanceNotifyEnabled, field.TypeBool, value)
}
@@ -1997,6 +2014,20 @@ func (_u *UserUpdateOne) ClearLastActiveAt() *UserUpdateOne {
return _u
}
// SetRestrictPublicGroups sets the "restrict_public_groups" field.
func (_u *UserUpdateOne) SetRestrictPublicGroups(v bool) *UserUpdateOne {
_u.mutation.SetRestrictPublicGroups(v)
return _u
}
// SetNillableRestrictPublicGroups sets the "restrict_public_groups" field if the given value is not nil.
func (_u *UserUpdateOne) SetNillableRestrictPublicGroups(v *bool) *UserUpdateOne {
if v != nil {
_u.SetRestrictPublicGroups(*v)
}
return _u
}
// SetBalanceNotifyEnabled sets the "balance_notify_enabled" field.
func (_u *UserUpdateOne) SetBalanceNotifyEnabled(v bool) *UserUpdateOne {
_u.mutation.SetBalanceNotifyEnabled(v)
@@ -2775,6 +2806,9 @@ func (_u *UserUpdateOne) sqlSave(ctx context.Context) (_node *User, err error) {
if _u.mutation.LastActiveAtCleared() {
_spec.ClearField(user.FieldLastActiveAt, field.TypeTime)
}
if value, ok := _u.mutation.RestrictPublicGroups(); ok {
_spec.SetField(user.FieldRestrictPublicGroups, field.TypeBool, value)
}
if value, ok := _u.mutation.BalanceNotifyEnabled(); ok {
_spec.SetField(user.FieldBalanceNotifyEnabled, field.TypeBool, value)
}
+12 -7
View File
@@ -1,6 +1,6 @@
module github.com/Wei-Shaw/sub2api
go 1.26.6
go 1.27.0
require (
entgo.io/ent v0.14.5
@@ -24,6 +24,8 @@ require (
github.com/google/uuid v1.6.0
github.com/google/wire v0.7.0
github.com/gorilla/websocket v1.5.3
github.com/hashicorp/go-hclog v1.6.3
github.com/hashicorp/go-plugin v1.8.0
github.com/imroc/req/v3 v3.59.0
github.com/klauspost/compress v1.18.2
github.com/lib/pq v1.10.9
@@ -54,6 +56,8 @@ require (
golang.org/x/net v0.56.0
golang.org/x/sync v0.21.0
golang.org/x/term v0.44.0
google.golang.org/grpc v1.82.1
google.golang.org/protobuf v1.36.11
gopkg.in/natefinch/lumberjack.v2 v2.2.1
gopkg.in/yaml.v3 v3.0.1
modernc.org/sqlite v1.44.3
@@ -121,12 +125,14 @@ require (
github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
github.com/go-webauthn/x v0.2.6 // indirect
github.com/goccy/go-json v0.10.2 // indirect
github.com/golang/protobuf v1.5.4 // indirect
github.com/google/go-cmp v0.7.0 // indirect
github.com/google/go-querystring v1.1.0 // indirect
github.com/google/go-tpm v0.9.8 // indirect
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3 // indirect
github.com/hashicorp/hcl v1.0.0 // indirect
github.com/hashicorp/hcl/v2 v2.18.1 // indirect
github.com/hashicorp/yamux v0.1.2 // indirect
github.com/icholy/digest v1.1.0 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/cpuid/v2 v2.2.4 // indirect
@@ -149,6 +155,7 @@ require (
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/morikuni/aec v1.0.0 // indirect
github.com/ncruces/go-strftime v1.0.0 // indirect
github.com/oklog/run v1.1.0 // indirect
github.com/opencontainers/go-digest v1.0.0 // indirect
github.com/opencontainers/image-spec v1.1.1 // indirect
github.com/pelletier/go-toml/v2 v2.2.2 // indirect
@@ -187,10 +194,9 @@ require (
github.com/zclconf/go-cty-yaml v1.1.0 // indirect
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
go.opentelemetry.io/otel v1.41.0 // indirect
go.opentelemetry.io/otel/metric v1.41.0 // indirect
go.opentelemetry.io/otel/sdk v1.41.0 // indirect
go.opentelemetry.io/otel/trace v1.41.0 // indirect
go.opentelemetry.io/otel v1.43.0 // indirect
go.opentelemetry.io/otel/metric v1.43.0 // indirect
go.opentelemetry.io/otel/trace v1.43.0 // indirect
go.uber.org/atomic v1.10.0 // indirect
go.uber.org/automaxprocs v1.6.0 // indirect
go.uber.org/multierr v1.9.0 // indirect
@@ -200,8 +206,7 @@ require (
golang.org/x/text v0.39.0 // indirect
golang.org/x/time v0.12.0 // indirect
golang.org/x/tools v0.47.0 // indirect
google.golang.org/grpc v1.75.1 // indirect
google.golang.org/protobuf v1.36.10 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 // indirect
gopkg.in/ini.v1 v1.67.0 // indirect
modernc.org/libc v1.67.6 // indirect
modernc.org/mathutil v1.7.1 // indirect
+45 -26
View File
@@ -120,6 +120,8 @@ github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw=
github.com/bufbuild/protocompile v0.14.1/go.mod h1:ppVdAIhbr2H8asPk6k4pY7t9zB1OU5DoEw9xY/FUi1c=
github.com/bytedance/sonic v1.5.0/go.mod h1:ED5hyg4y6t3/9Ku1R6dU/4KyJ48DZ4jPhfY1O2AihPM=
github.com/bytedance/sonic v1.9.1 h1:6iJ6NqdoxCDr6mbY8h18oSO+cShGSMRGCEo7F2h0x8s=
github.com/bytedance/sonic v1.9.1/go.mod h1:i736AoUSYt75HyZLoJW9ERYxcy6eaN6h4BZXU064P/U=
@@ -176,6 +178,7 @@ github.com/ebitengine/purego v0.8.4/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk=
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=
github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU=
github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2Wg=
@@ -232,6 +235,8 @@ github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrU
github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w=
github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0=
github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M=
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
@@ -250,8 +255,6 @@ github.com/google/go-tpm-tools v0.3.13-0.20230620182252-4639ecce2aba/go.mod h1:E
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE=
github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4=
@@ -262,6 +265,10 @@ github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aN
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3 h1:NmZ1PKzSTQbuGHw9DGPFomqkkLWMC+vZCkfs+FHv1Vg=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3/go.mod h1:zQrxl1YP88HQlA6i9c63DSVPFklWpGX4OWAc9bFuaH4=
github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB11/k=
github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M=
github.com/hashicorp/go-plugin v1.8.0 h1:ie8S6RRY8RvB2usYZv+AAZ/wBvx2AU5p5QeP5j/FORs=
github.com/hashicorp/go-plugin v1.8.0/go.mod h1:BExt6KEaIYx804z8k4gRzRLEvxKVb+kn0NMcihqOqb8=
github.com/hashicorp/golang-lru v0.5.4 h1:YDjusn29QI/Das2iO9M0BHnIbxPeyuCHsjMW+lJfyTc=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
@@ -269,6 +276,8 @@ github.com/hashicorp/hcl v1.0.0 h1:0Anlzjpi4vEasTeNFn2mLJgTSwt0+6sfsiTG8qcWGx4=
github.com/hashicorp/hcl v1.0.0/go.mod h1:E5yfLk+7swimpb2L/Alb/PJmXilQ/rhwaUYs4T20WEQ=
github.com/hashicorp/hcl/v2 v2.18.1 h1:6nxnOJFku1EuSawSD81fuviYUV8DxFr3fp2dUi3ZYSo=
github.com/hashicorp/hcl/v2 v2.18.1/go.mod h1:ThLC89FV4p9MPW804KVbe/cEXoQ8NZEh+JtMeeGErHE=
github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8=
github.com/hashicorp/yamux v0.1.2/go.mod h1:C+zze2n6e/7wshOZep2A70/aQU6QBRWJO/G6FT1wIns=
github.com/icholy/digest v1.1.0 h1:HfGg9Irj7i+IX1o1QAmPfIBNu/Q5A5Tu3n/MED9k9H4=
github.com/icholy/digest v1.1.0/go.mod h1:QNrsSGQ5v7v9cReDI0+eyjsXGUoRSUZQHeQ5C4XLa0Y=
github.com/imroc/req/v3 v3.59.0 h1:PqKhJHyBmJYob47LVuTHwRZE00ZO6icbLHe5Zra13jo=
@@ -281,6 +290,8 @@ github.com/jackc/pgx/v5 v5.7.4 h1:9wKznZrhWa2QiHL+NjTSPP6yjl3451BX3imWDnokYlg=
github.com/jackc/pgx/v5 v5.7.4/go.mod h1:ncY89UGWxg82EykZUwSpUKEfccBGGYq1xjrOpsbsfGQ=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jhump/protoreflect v1.17.0 h1:qOEr613fac2lOuTgWN4tPAtLL7fUSbuJL5X5XumQh94=
github.com/jhump/protoreflect v1.17.0/go.mod h1:h9+vUUL38jiBzck8ck+6G/aeMX8Z4QUY/NiJPwPNi+8=
github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
@@ -307,13 +318,15 @@ github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0 h1:6E+4a0GO5zZEnZ
github.com/lufia/plan9stats v0.0.0-20211012122336-39d0f177ccd0/go.mod h1:zJYVVT2jmtg6P3p1VtQj7WsuWi/y4VnjVBn7F8KPB3I=
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
github.com/mattn/go-colorable v0.1.9/go.mod h1:u6P/XSegPjTcexA+o6vUJrdnUu04hMope9wVRipJSqc=
github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4=
github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA=
github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg=
github.com/mattn/go-isatty v0.0.12/go.mod h1:cbi8OIDigv2wuxKPP5vlRcQ1OAZbq2CE4Kysco4FUpU=
github.com/mattn/go-isatty v0.0.14/go.mod h1:7GGIvUiUoEMVVmxf/4nioHXj79iQHKdU27kJ6hsGG94=
github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-runewidth v0.0.15 h1:UNAjwbU9l54TA3KzvqLGxwWjHmMgBUVhBiTjelZgg3U=
github.com/mattn/go-runewidth v0.0.15/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w=
github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM=
github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg=
github.com/mdelapenya/tlscert v0.2.0 h1:7H81W6Z/4weDvZBNOfQte5GpIMo0lGYEeWbkGp5LJHI=
@@ -350,8 +363,8 @@ github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7P
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e/go.mod h1:zD1mROLANZcx1PVRCS0qkT7pwLkGfwJo4zjcN/Tysno=
github.com/olekukonko/tablewriter v0.0.5 h1:P2Ga83D34wi1o9J6Wh1mRuqd4mF/x/lgBS7N7AbDhec=
github.com/olekukonko/tablewriter v0.0.5/go.mod h1:hPp6KlRPjbx+hW8ykQs1w3UBbZlj6HuIJcUGPhkA7kY=
github.com/oklog/run v1.1.0 h1:GEenZ1cK0+q0+wsJew9qUg/DyD8k3JzYsZAi5gYi2mA=
github.com/oklog/run v1.1.0/go.mod h1:sVPdnTZT1zYwAJeCMu2Th4T21pA3FPOQRfWjQlk7DVU=
github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U=
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
@@ -386,8 +399,6 @@ github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEv
github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY=
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
@@ -423,8 +434,6 @@ github.com/spf13/afero v1.11.0 h1:WJQKhtpdm3v2IzqG8VMqrr6Rf3UYpEF239Jy9wNepM8=
github.com/spf13/afero v1.11.0/go.mod h1:GH9Y3pIexgf1MTIWtNGyogA5MwRIDXGUr+hbWNoBjkY=
github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0=
github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo=
github.com/spf13/cobra v1.7.0 h1:hyqWnYt1ZQShIddO5kBpj3vu05/++x6tJ6dg8EC572I=
github.com/spf13/cobra v1.7.0/go.mod h1:uLxZILRyS/50WlhOIKD7W6V5bgeIt+4sICxh6uRMrb0=
github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA=
github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ=
@@ -439,6 +448,7 @@ github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UV
github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals=
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
@@ -507,18 +517,20 @@ go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0/go.mod h1:p8pYQP+m5XfbZm9fxtSKAbM6oIllS7s2AfxrChvc7iw=
go.opentelemetry.io/otel v1.41.0 h1:YlEwVsGAlCvczDILpUXpIpPSL/VPugt7zHThEMLce1c=
go.opentelemetry.io/otel v1.41.0/go.mod h1:Yt4UwgEKeT05QbLwbyHXEwhnjxNO6D8L5PQP51/46dE=
go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I=
go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.24.0 h1:t6wl9SPayj+c7lEIFgm4ooDBZVb01IhLB4InpomhRw8=
go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.24.0/go.mod h1:iSDOcsnSA5INXzZtwaBPrKp/lWu/V14Dd+llD0oI2EA=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.24.0 h1:Xw8U6u2f8DK2XAkGRFV7BBLENgnTGX9i4rQRxJf+/vs=
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.24.0/go.mod h1:6KW1Fm6R/s6Z3PGXwSJN2K4eT6wQB3vXX6CVnYX9NmM=
go.opentelemetry.io/otel/metric v1.41.0 h1:rFnDcs4gRzBcsO9tS8LCpgR0dxg4aaxWlJxCno7JlTQ=
go.opentelemetry.io/otel/metric v1.41.0/go.mod h1:xPvCwd9pU0VN8tPZYzDZV/BMj9CM9vs00GuBjeKhJps=
go.opentelemetry.io/otel/sdk v1.41.0 h1:YPIEXKmiAwkGl3Gu1huk1aYWwtpRLeskpV+wPisxBp8=
go.opentelemetry.io/otel/sdk v1.41.0/go.mod h1:ahFdU0G5y8IxglBf0QBJXgSe7agzjE4GiTJ6HT9ud90=
go.opentelemetry.io/otel/trace v1.41.0 h1:Vbk2co6bhj8L59ZJ6/xFTskY+tGAbOnCtQGVVa9TIN0=
go.opentelemetry.io/otel/trace v1.41.0/go.mod h1:U1NU4ULCoxeDKc09yCWdWe+3QoyweJcISEVa1RBzOis=
go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM=
go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY=
go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg=
go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg=
go.opentelemetry.io/otel/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw=
go.opentelemetry.io/otel/sdk/metric v1.43.0/go.mod h1:C/RJtwSEJ5hzTiUz5pXF1kILHStzb9zFlIEe85bhj6A=
go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A=
go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0=
go.opentelemetry.io/proto/otlp v1.3.1 h1:TrMUixzpM0yuc/znrFTP9MMRh8trP93mkCiDVeXrui0=
go.opentelemetry.io/proto/otlp v1.3.1/go.mod h1:0X1WI4de4ZsLrrJNLAQbFeLCm3T7yBkR0XqQ7niQU+8=
go.uber.org/atomic v1.10.0 h1:9qC72Qh0+3MqyJbAn8YU5xVq1frD8bn3JtD2oXtafVQ=
@@ -608,6 +620,8 @@ golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5h
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200116001909-b77594299b42/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200323222414-85ca7c5b95cd/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200509044756-6aff5f38e54f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -615,6 +629,9 @@ golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20210927094055-39ccf1dd6fa6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220503163025-988cb79eb6c6/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220704084225-05e143d24a9e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
@@ -678,29 +695,31 @@ golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc=
google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc=
google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 h1:wpZ8pe2x1Q3f2KyT5f8oP/fa9rHAKgFPr/HZdNuS+PQ=
google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4 h1:8XJ4pajGwOlasW+L13MnEGA8W4115jJySQtVfS2/IBU=
google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4/go.mod h1:NnuHhy+bxcg30o7FnVAZbXsPHUDQ9qKWAQKCD7VxFtk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 h1:i8QOKZfYg6AbGVZzUAY3LrNWCKF8O6zFisU9Wl9RER4=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4/go.mod h1:HSkG/KdJWusxU1F6CNrwNDjBMgisKxGnc5dAZfT0mjQ=
google.golang.org/genproto/googleapis/api v0.0.0-20260414002931-afd174a4e478 h1:yQugLulqltosq0B/f8l4w9VryjV+N/5gcW0jQ3N8Qec=
google.golang.org/genproto/googleapis/api v0.0.0-20260414002931-afd174a4e478/go.mod h1:C6ADNqOxbgdUUeRTU+LCHDPB9ttAMCTff6auwCVa4uc=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 h1:RmoJA1ujG+/lRGNfUnOMfhCy5EipVMyvUE+KNbPbTlw=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c=
google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg=
google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY=
google.golang.org/grpc v1.31.0/go.mod h1:N36X2cJ7JwdamYAgDz+s+rVMFjt3numwzf/HckM8pak=
google.golang.org/grpc v1.75.1 h1:/ODCNEuf9VghjgO3rqLcfg8fiOP0nSluljWFlDxELLI=
google.golang.org/grpc v1.75.1/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
google.golang.org/grpc v1.82.1 h1:NnAxzGRA0677vCa4BUkOAnO5+FfQqVl9iUXeD0IqcGE=
google.golang.org/grpc v1.82.1/go.mod h1:yzTZ1TB1Z3SG+LIYaI+WiE8D5+PZ3ArnrSp8zF3+/ZA=
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM=
google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE=
google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo=
google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE=
google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20200227125254-8fa46927fb4f/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
+99 -1
View File
@@ -32,7 +32,7 @@ const (
// DefaultCSPPolicy is the default Content-Security-Policy with nonce support
// __CSP_NONCE__ will be replaced with actual nonce at request time by the SecurityHeaders middleware
const DefaultCSPPolicy = "default-src 'self'; worker-src 'self' blob:; script-src 'self' __CSP_NONCE__ https://challenges.cloudflare.com https://*.alicdn.com https://static.cloudflareinsights.com https://turing.captcha.qcloud.com https://turing.captcha.gtimg.com https://ca.turing.captcha.qcloud.com https://global.turing.captcha.gtimg.com https://www.tycaptcha.com https://cloudcache.tencentcs.com https://*.stripe.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; style-src 'self' 'unsafe-inline' https://*.captcha.gtimg.com https://fonts.googleapis.com https://*.alicdn.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; img-src 'self' data: blob: https:; font-src 'self' data: https://fonts.gstatic.com; connect-src 'self' https://turing.captcha.qcloud.com https://www.tycaptcha.com https://rce.tencentrio.com https:; frame-src https://challenges.cloudflare.com https://turing.captcha.qcloud.com https://ca.turing.captcha.qcloud.com https://www.tycaptcha.com https://*.stripe.com https://checkout.airwallex.com https://checkout-demo.airwallex.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'"
const DefaultCSPPolicy = "default-src 'self'; worker-src 'self' blob:; script-src 'self' __CSP_NONCE__ https://challenges.cloudflare.com https://*.alicdn.com https://static.cloudflareinsights.com https://turing.captcha.qcloud.com https://turing.captcha.gtimg.com https://ca.turing.captcha.qcloud.com https://global.turing.captcha.gtimg.com https://www.tycaptcha.com https://cloudcache.tencentcs.com https://*.stripe.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; style-src 'self' 'unsafe-inline' https://*.captcha.gtimg.com https://fonts.googleapis.com https://*.alicdn.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; img-src 'self' data: blob: https:; font-src 'self' data: https://fonts.gstatic.com; connect-src 'self' https://turing.captcha.qcloud.com https://www.tycaptcha.com https://rce.tencentrio.com https:; frame-src 'self' https://challenges.cloudflare.com https://turing.captcha.qcloud.com https://ca.turing.captcha.qcloud.com https://www.tycaptcha.com https://*.stripe.com https://checkout.airwallex.com https://checkout-demo.airwallex.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'"
// UMQ(用户消息队列)模式常量
const (
@@ -61,6 +61,10 @@ const (
// 可通过 gateway.upstream_response_read_max_bytes 配置项覆盖。
const DefaultUpstreamResponseReadMaxBytes int64 = 128 * 1024 * 1024
// DefaultModelsListReadMaxBytes 上游模型列表响应体的默认读取上限。
// 可通过 gateway.models_list_read_max_bytes 配置项覆盖。
const DefaultModelsListReadMaxBytes int64 = 8 * 1024 * 1024
type Config struct {
Server ServerConfig `mapstructure:"server"`
Log LogConfig `mapstructure:"log"`
@@ -99,6 +103,18 @@ type Config struct {
Idempotency IdempotencyConfig `mapstructure:"idempotency"`
BatchImage BatchImageConfig `mapstructure:"batch_image"`
ImageStorage ImageStorageConfig `mapstructure:"image_storage"`
Plugins PluginConfig `mapstructure:"plugins"`
}
// PluginConfig 控制管理员手动上传的本地进程插件。
// 默认不包含插件,也不允许安装未签名插件;TrustedPublishers 用于追加第三方发布者。
type PluginConfig struct {
DataDir string `mapstructure:"data_dir"`
AllowUnsigned bool `mapstructure:"allow_unsigned"`
TrustedPublishers map[string]string `mapstructure:"trusted_publishers"`
MaxUploadBytes int64 `mapstructure:"max_upload_bytes"`
MaxUncompressedBytes int64 `mapstructure:"max_uncompressed_bytes"`
StartTimeoutSeconds int `mapstructure:"start_timeout_seconds"`
}
type LogConfig struct {
@@ -830,6 +846,53 @@ type ProxyFallbackConfig struct {
type ProxyProbeConfig struct {
InsecureSkipVerify bool `mapstructure:"insecure_skip_verify"` // 已禁用:禁止跳过 TLS 证书验证
// URLs 按优先级排列的自定义探测 URL 列表。
// 留空时使用内置默认列表(ip-api → ipify)。
// 某些 AI API 专用代理只允许访问特定域名,配置多个备选可提高探测成功率。
URLs []ProbeURLConfig `mapstructure:"urls"`
}
// ProbeURLConfig 描述一个探测端点及其响应解析方式。
type ProbeURLConfig struct {
URL string `mapstructure:"url"`
Parser string `mapstructure:"parser"` // "ip-api" / "ipify" / "chatgpt-trace"
}
func normalizeProxyProbeURLs(targets []ProbeURLConfig) ([]ProbeURLConfig, error) {
if len(targets) == 0 {
return nil, nil
}
normalized := make([]ProbeURLConfig, 0, len(targets))
for i, target := range targets {
rawURL := strings.TrimSpace(target.URL)
parser := strings.ToLower(strings.TrimSpace(target.Parser))
if rawURL == "" {
return nil, fmt.Errorf("entry %d: url is required", i)
}
if parser == "" {
return nil, fmt.Errorf("entry %d: parser is required", i)
}
switch parser {
case "ip-api", "ipify", "chatgpt-trace":
default:
return nil, fmt.Errorf("entry %d: unsupported parser %q", i, target.Parser)
}
parsed, err := url.Parse(rawURL)
if err != nil || parsed.Host == "" {
return nil, fmt.Errorf("entry %d: invalid url %q", i, target.URL)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return nil, fmt.Errorf("entry %d: url scheme must be http or https", i)
}
normalized = append(normalized, ProbeURLConfig{
URL: rawURL,
Parser: parser,
})
}
return normalized, nil
}
type BillingConfig struct {
@@ -887,6 +950,9 @@ type GatewayConfig struct {
// OpenAIResponseHeaderTimeout: OpenAI/Codex 上游等待响应头的超时时间(秒),0表示无超时
// OpenAI/Codex 请求可能在上游排队较久;默认不使用通用响应头超时截断。
OpenAIResponseHeaderTimeout int `mapstructure:"openai_response_header_timeout"`
// GrokResponseHeaderTimeout bounds the pre-first-byte wait for xAI/Grok.
// A zero value uses the provider-safe default instead of the generic gateway timeout.
GrokResponseHeaderTimeout int `mapstructure:"grok_response_header_timeout"`
// OpenAIFirstOutputTimeoutSeconds: native HTTP Responses 首个语义输出超时(秒),0表示禁用。
OpenAIFirstOutputTimeoutSeconds int `mapstructure:"openai_first_output_timeout_seconds"`
// OpenAIHighEffortFirstOutputTimeoutSeconds: high/xhigh/max 推理的首个语义输出超时(秒)。
@@ -898,6 +964,8 @@ type GatewayConfig struct {
TextMaxBodySize int64 `mapstructure:"text_max_body_size"`
// 非流式上游响应体读取上限(字节),用于防止无界读取导致内存放大
UpstreamResponseReadMaxBytes int64 `mapstructure:"upstream_response_read_max_bytes"`
// 上游模型列表响应体读取上限(字节)
ModelsListReadMaxBytes int64 `mapstructure:"models_list_read_max_bytes"`
// 代理探测响应体读取上限(字节)
ProxyProbeResponseReadMaxBytes int64 `mapstructure:"proxy_probe_response_read_max_bytes"`
// Gemini 上游响应头调试日志开关(默认关闭,避免高频日志开销)
@@ -2219,6 +2287,14 @@ func setDefaults() {
viper.SetDefault("pricing.update_interval_hours", 24)
viper.SetDefault("pricing.hash_check_interval_minutes", 10)
// 本地进程插件。插件必须由管理员手动上传,项目默认不携带任何插件能力。
viper.SetDefault("plugins.data_dir", "")
viper.SetDefault("plugins.allow_unsigned", false)
viper.SetDefault("plugins.trusted_publishers", map[string]string{})
viper.SetDefault("plugins.max_upload_bytes", int64(128*1024*1024))
viper.SetDefault("plugins.max_uncompressed_bytes", int64(256*1024*1024))
viper.SetDefault("plugins.start_timeout_seconds", 15)
// Timezone (default to Asia/Shanghai for Chinese users)
viper.SetDefault("timezone", "Asia/Shanghai")
@@ -2280,6 +2356,7 @@ func setDefaults() {
// Gateway
viper.SetDefault("gateway.response_header_timeout", 600) // 600秒(10分钟)等待上游响应头,LLM高负载时可能排队较久
viper.SetDefault("gateway.openai_response_header_timeout", 0)
viper.SetDefault("gateway.grok_response_header_timeout", 120)
viper.SetDefault("gateway.openai_first_output_timeout_seconds", 0)
viper.SetDefault("gateway.openai_high_effort_first_output_timeout_seconds", 0)
viper.SetDefault("gateway.log_upstream_error_body", true)
@@ -2385,6 +2462,7 @@ func setDefaults() {
viper.SetDefault("gateway.max_body_size", int64(256*1024*1024))
viper.SetDefault("gateway.text_max_body_size", int64(32*1024*1024))
viper.SetDefault("gateway.upstream_response_read_max_bytes", DefaultUpstreamResponseReadMaxBytes)
viper.SetDefault("gateway.models_list_read_max_bytes", DefaultModelsListReadMaxBytes)
viper.SetDefault("gateway.proxy_probe_response_read_max_bytes", int64(1024*1024))
viper.SetDefault("gateway.gemini_debug_response_headers", false)
viper.SetDefault("gateway.connection_pool_isolation", ConnectionPoolIsolationAccountProxy)
@@ -2569,6 +2647,20 @@ func (c *Config) Validate() error {
}
c.Security.ForwardedClientIPHeaders = forwardedClientIPHeaders
c.SetForwardedClientIPSettings(c.Security.TrustForwardedIPForAPIKeyACL, forwardedClientIPHeaders)
proxyProbeURLs, err := normalizeProxyProbeURLs(c.Security.ProxyProbe.URLs)
if err != nil {
return fmt.Errorf("security.proxy_probe.urls: %w", err)
}
c.Security.ProxyProbe.URLs = proxyProbeURLs
if c.Plugins.MaxUploadBytes <= 0 || c.Plugins.MaxUploadBytes > 1024*1024*1024 {
return fmt.Errorf("plugins.max_upload_bytes must be between 1 and 1073741824")
}
if c.Plugins.MaxUncompressedBytes < c.Plugins.MaxUploadBytes || c.Plugins.MaxUncompressedBytes > 2*1024*1024*1024 {
return fmt.Errorf("plugins.max_uncompressed_bytes must be between max_upload_bytes and 2147483648")
}
if c.Plugins.StartTimeoutSeconds < 1 || c.Plugins.StartTimeoutSeconds > 120 {
return fmt.Errorf("plugins.start_timeout_seconds must be between 1 and 120")
}
if c.Server.ReadHeaderTimeout < 1 || c.Server.ReadHeaderTimeout > 60 {
return fmt.Errorf("server.read_header_timeout must be between 1 and 60 seconds")
}
@@ -3180,6 +3272,9 @@ func (c *Config) Validate() error {
if c.Gateway.UpstreamResponseReadMaxBytes <= 0 {
return fmt.Errorf("gateway.upstream_response_read_max_bytes must be positive")
}
if c.Gateway.ModelsListReadMaxBytes <= 0 {
return fmt.Errorf("gateway.models_list_read_max_bytes must be positive")
}
if c.Gateway.ProxyProbeResponseReadMaxBytes <= 0 {
return fmt.Errorf("gateway.proxy_probe_response_read_max_bytes must be positive")
}
@@ -3189,6 +3284,9 @@ func (c *Config) Validate() error {
if c.Gateway.OpenAIResponseHeaderTimeout < 0 {
return fmt.Errorf("gateway.openai_response_header_timeout must be non-negative")
}
if c.Gateway.GrokResponseHeaderTimeout < 0 || c.Gateway.GrokResponseHeaderTimeout > 1800 {
return fmt.Errorf("gateway.grok_response_header_timeout must be between 0-1800 seconds")
}
if c.Gateway.OpenAIFirstOutputTimeoutSeconds < 0 || c.Gateway.OpenAIFirstOutputTimeoutSeconds > 600 ||
(c.Gateway.OpenAIFirstOutputTimeoutSeconds > 0 && c.Gateway.OpenAIFirstOutputTimeoutSeconds < 30) {
return fmt.Errorf("gateway.openai_first_output_timeout_seconds must be 0 or between 30-600 seconds")
+21
View File
@@ -23,6 +23,13 @@ func resetViperWithJWTSecret(t *testing.T) {
t.Setenv("JWT_SECRET", strings.Repeat("x", 32))
}
func TestLoadDefaultModelsListReadMaxBytes(t *testing.T) {
resetViperWithJWTSecret(t)
cfg, err := Load()
require.NoError(t, err)
require.Equal(t, DefaultModelsListReadMaxBytes, cfg.Gateway.ModelsListReadMaxBytes)
}
func TestLoadTimezonePrecedence(t *testing.T) {
tests := []struct {
name string
@@ -553,6 +560,15 @@ func TestLoadOpenAIWSClientFirstMessageTimeoutFromEnv(t *testing.T) {
require.Equal(t, 120, cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds)
}
func TestLoadOpenAIWSForceHTTPFromEnv(t *testing.T) {
resetViperWithJWTSecret(t)
t.Setenv("GATEWAY_OPENAI_WS_FORCE_HTTP", "true")
cfg, err := Load()
require.NoError(t, err)
require.True(t, cfg.Gateway.OpenAIWS.ForceHTTP)
}
func TestLoadDefaultOpenAICompactModel(t *testing.T) {
resetViperWithJWTSecret(t)
@@ -1794,6 +1810,11 @@ func TestValidateConfigErrors(t *testing.T) {
mutate: func(c *Config) { c.Gateway.TextMaxBodySize = c.Gateway.MaxBodySize + 1 },
wantErr: "gateway.text_max_body_size",
},
{
name: "gateway models list read limit",
mutate: func(c *Config) { c.Gateway.ModelsListReadMaxBytes = 0 },
wantErr: "gateway.models_list_read_max_bytes",
},
{
name: "gateway response header timeout",
mutate: func(c *Config) { c.Gateway.ResponseHeaderTimeout = -1 },
@@ -45,6 +45,20 @@ func collectMapstructureKeys(t reflect.Type, prefix string, out map[string]strin
// is out of scope here — such settings need a config file either way.
continue
}
if ft.Kind() == reflect.Slice {
elem := ft.Elem()
for elem.Kind() == reflect.Ptr {
elem = elem.Elem()
}
if elem.Kind() == reflect.Struct {
// AutomaticEnv exposes one string value. Viper's string-to-slice
// hook can populate scalar slices, but it cannot decode a string
// into []struct. Registering a default would turn silent ignore
// into a startup unmarshal error, so structured slices remain
// config-file-only just like maps.
continue
}
}
out[strings.ToLower(key)] = ft.String()
}
}
@@ -62,7 +76,7 @@ func collectMapstructureKeys(t reflect.Type, prefix string, out map[string]strin
// were lost, silently disabling async image tasks for env-driven deployments.
//
// When this fails, register a zero-valued default in setEnvReachableDefaults
// for each reported key.
// for each reported scalar key. Maps and slices of structs are config-file-only.
func TestConfigKeysAreEnvReachable(t *testing.T) {
bound := map[string]string{}
collectMapstructureKeys(reflect.TypeOf(Config{}), "", bound)
@@ -0,0 +1,47 @@
//go:build unit
package config
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestNormalizeProxyProbeURLs(t *testing.T) {
t.Parallel()
got, err := normalizeProxyProbeURLs([]ProbeURLConfig{
{URL: " https://chatgpt.com/cdn-cgi/trace ", Parser: " CHATGPT-TRACE "},
{URL: "https://api64.ipify.org?format=json", Parser: "ipify"},
})
require.NoError(t, err)
require.Equal(t, []ProbeURLConfig{
{URL: "https://chatgpt.com/cdn-cgi/trace", Parser: "chatgpt-trace"},
{URL: "https://api64.ipify.org?format=json", Parser: "ipify"},
}, got)
}
func TestNormalizeProxyProbeURLsRejectsInvalidEntries(t *testing.T) {
t.Parallel()
tests := []struct {
name string
target ProbeURLConfig
wantErr string
}{
{name: "missing URL", target: ProbeURLConfig{Parser: "ipify"}, wantErr: "url is required"},
{name: "missing parser", target: ProbeURLConfig{URL: "https://example.com"}, wantErr: "parser is required"},
{name: "unknown parser", target: ProbeURLConfig{URL: "https://example.com", Parser: "ip_api"}, wantErr: "unsupported parser"},
{name: "relative URL", target: ProbeURLConfig{URL: "/cdn-cgi/trace", Parser: "chatgpt-trace"}, wantErr: "invalid url"},
{name: "unsupported scheme", target: ProbeURLConfig{URL: "ftp://example.com/file", Parser: "ipify"}, wantErr: "scheme must be http or https"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
_, err := normalizeProxyProbeURLs([]ProbeURLConfig{tt.target})
require.ErrorContains(t, err, tt.wantErr)
})
}
}
+4 -3
View File
@@ -44,6 +44,7 @@ const (
APIProtocolChatCompletions = "chat_completions" // OpenAI Chat Completions(默认)
APIProtocolAnthropic = "anthropic" // 原生 Anthropic /v1/messages(适配 Claude Code)
APIProtocolResponses = "responses" // OpenAI Responses(仅 deepseek,适配 Codex)
APIProtocolAdaptive = "adaptive" // 按入站协议优先选择供应商原生端点
)
// Account type constants
@@ -104,11 +105,11 @@ var DefaultAntigravityModelMapping = map[string]string{
"claude-opus-4-6": "claude-opus-4-6-thinking", // 简称映射
"claude-opus-4-5-thinking": "claude-opus-4-6-thinking", // 迁移旧模型
"claude-sonnet-4-6": "claude-sonnet-4-6",
"claude-sonnet-4-5": "claude-sonnet-4-5",
"claude-sonnet-4-5-thinking": "claude-sonnet-4-5-thinking",
"claude-sonnet-4-5": "claude-sonnet-4-5", // 显式 canonical 选择透传
"claude-sonnet-4-5-thinking": "claude-sonnet-4-6", // 迁移旧兼容别名
// Claude 详细版本 ID 映射
"claude-opus-4-5-20251101": "claude-opus-4-6-thinking", // 迁移旧模型
"claude-sonnet-4-5-20250929": "claude-sonnet-4-5",
"claude-sonnet-4-5-20250929": "claude-sonnet-4-6", // 迁移旧模型
// Claude Haiku → Sonnet(无 Haiku 支持)
"claude-haiku-4-5": "claude-sonnet-4-6",
"claude-haiku-4-5-20251001": "claude-sonnet-4-6",
+15
View File
@@ -43,6 +43,21 @@ func TestDefaultAntigravityModelMapping_ContainsNewClaudeModels(t *testing.T) {
}
}
func TestDefaultAntigravityModelMapping_PreservesExplicitSonnet45AndMigratesLegacyAliases(t *testing.T) {
t.Parallel()
cases := map[string]string{
"claude-sonnet-4-5": "claude-sonnet-4-5",
"claude-sonnet-4-5-thinking": "claude-sonnet-4-6",
"claude-sonnet-4-5-20250929": "claude-sonnet-4-6",
}
for model, want := range cases {
if got := DefaultAntigravityModelMapping[model]; got != want {
t.Fatalf("expected model %q to map to %q, got %q", model, want, got)
}
}
}
func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) {
t.Parallel()
@@ -2614,8 +2614,14 @@ func (h *AccountHandler) GetAvailableModels(c *gin.Context) {
// Handle Gemini accounts
if account.IsGemini() {
// For OAuth accounts: return default Gemini models
// Consumer Google One OAuth still uses the legacy Gemini CLI / Code
// Assist channel. Do not advertise newer 3.x or image models that the
// channel cannot serve.
if account.IsOAuth() {
if account.IsGeminiGoogleOne() {
response.Success(c, geminicli.GoogleOneModels)
return
}
response.Success(c, geminicli.DefaultModels)
return
}
@@ -2768,13 +2774,15 @@ func (h *AccountHandler) SyncUpstreamModels(c *gin.Context) {
return
}
models, err := h.accountTestService.FetchUpstreamSupportedModels(c.Request.Context(), account)
catalog, err := h.accountTestService.SyncUpstreamModelCatalog(c.Request.Context(), account)
if err != nil {
var syncErr *service.UpstreamModelSyncError
if errors.As(err, &syncErr) {
switch syncErr.Kind {
case service.UpstreamModelSyncErrorConfiguration, service.UpstreamModelSyncErrorUnsupported:
response.BadRequest(c, syncErr.SafeMessage())
case service.UpstreamModelSyncErrorInternal:
response.InternalError(c, syncErr.SafeMessage())
default:
slog.Warn("sync_upstream_models_failed", "account_id", accountID, "kind", syncErr.Kind)
response.Error(c, http.StatusBadGateway, syncErr.SafeMessage())
@@ -2787,29 +2795,35 @@ func (h *AccountHandler) SyncUpstreamModels(c *gin.Context) {
return
}
response.Success(c, gin.H{"models": models})
response.Success(c, catalog)
}
// SyncUpstreamModelsPreview handles syncing live supported models using provided credentials (no account ID needed).
// POST /api/v1/admin/accounts/models/sync-upstream-preview
func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) {
var req struct {
Platform string `json:"platform" binding:"required"`
Type string `json:"type" binding:"required"`
BaseURL string `json:"base_url"`
APIKey string `json:"api_key" binding:"required"`
Platform string `json:"platform" binding:"required"`
Type string `json:"type" binding:"required"`
BaseURL string `json:"base_url"`
APIKey string `json:"api_key" binding:"required"`
ModelMapping map[string]string `json:"model_mapping"`
}
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
modelMapping := make(map[string]any, len(req.ModelMapping))
for sourceModel, upstreamModel := range req.ModelMapping {
modelMapping[sourceModel] = upstreamModel
}
tempAccount := &service.Account{
Platform: req.Platform,
Type: req.Type,
Credentials: map[string]any{
"api_key": req.APIKey,
"base_url": req.BaseURL,
"api_key": req.APIKey,
"base_url": req.BaseURL,
"model_mapping": modelMapping,
},
}
@@ -2818,13 +2832,15 @@ func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) {
return
}
models, err := h.accountTestService.FetchUpstreamSupportedModels(c.Request.Context(), tempAccount)
catalog, err := h.accountTestService.SyncUpstreamModelCatalog(c.Request.Context(), tempAccount)
if err != nil {
var syncErr *service.UpstreamModelSyncError
if errors.As(err, &syncErr) {
switch syncErr.Kind {
case service.UpstreamModelSyncErrorConfiguration, service.UpstreamModelSyncErrorUnsupported:
response.BadRequest(c, syncErr.SafeMessage())
case service.UpstreamModelSyncErrorInternal:
response.InternalError(c, syncErr.SafeMessage())
default:
slog.Warn("sync_upstream_models_preview_failed", "platform", req.Platform, "kind", syncErr.Kind)
response.Error(c, http.StatusBadGateway, syncErr.SafeMessage())
@@ -2837,7 +2853,7 @@ func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) {
return
}
response.Success(c, gin.H{"models": models})
response.Success(c, catalog)
}
// SetPrivacy handles setting privacy for a single OpenAI/Antigravity OAuth account
@@ -38,14 +38,20 @@ func setupAvailableModelsRouter(adminSvc service.AdminService) *gin.Engine {
}
type syncUpstreamHTTPUpstream struct {
resp *http.Response
err error
resp *http.Response
responses []*http.Response
err error
}
func (u *syncUpstreamHTTPUpstream) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) {
if u.err != nil {
return nil, u.err
}
if len(u.responses) > 0 {
resp := u.responses[0]
u.responses = u.responses[1:]
return resp, nil
}
return u.resp, nil
}
@@ -68,6 +74,7 @@ func setupSyncUpstreamModelsRouter(adminSvc service.AdminService, upstream servi
)
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, accountTestSvc, nil, nil, nil, nil, nil)
router.POST("/api/v1/admin/accounts/:id/models/sync-upstream", handler.SyncUpstreamModels)
router.POST("/api/v1/admin/accounts/models/sync-upstream-preview", handler.SyncUpstreamModelsPreview)
return router
}
@@ -287,6 +294,42 @@ func TestAccountHandlerGetAvailableModels_OpenAISparkShadowReturnsMappingModels(
}, ids, "影子可用模型由 model_mapping 派生(非写死)")
}
func TestAccountHandlerGetAvailableModels_GeminiGoogleOneUsesConservativeCatalog(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 45,
Name: "google-one",
Platform: service.PlatformGemini,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
Credentials: map[string]any{
"oauth_type": "google_one",
},
},
}
router := setupAvailableModelsRouter(svc)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/45/models", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
ids := make([]string, 0, len(resp.Data))
for _, model := range resp.Data {
ids = append(ids, model.ID)
}
require.ElementsMatch(t, []string{"gemini-2.0-flash", "gemini-2.5-flash", "gemini-2.5-pro"}, ids)
require.NotContains(t, ids, "gemini-3.5-flash")
require.NotContains(t, ids, "gemini-2.5-flash-image")
}
func TestAccountHandlerSyncUpstreamModels_ConfigErrorReturnsBadRequest(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
@@ -311,6 +354,99 @@ func TestAccountHandlerSyncUpstreamModels_ConfigErrorReturnsBadRequest(t *testin
require.Contains(t, rec.Body.String(), "No OpenAI API key is available")
}
func TestAccountHandlerSyncUpstreamModelsReturnsCapabilityMetadata(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 48, Name: "custom-openai", Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey, Status: service.StatusActive,
Credentials: map[string]any{"api_key": "key", "base_url": "https://provider.example/v1"},
},
}
upstream := &syncUpstreamHTTPUpstream{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"models":[{
"id":"custom-thinking-model",
"reasoning":true,
"default_reasoning_level":"high",
"supported_reasoning_levels":["low","high"],
"input_modalities":["text","image"],
"context_window":256000
}]}`)),
}}
router := setupSyncUpstreamModelsRouter(svc, upstream)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/48/models/sync-upstream", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data service.UpstreamModelCatalog `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, []string{"custom-thinking-model"}, resp.Data.Models)
metadata := resp.Data.Metadata["custom-thinking-model"]
require.NotNil(t, metadata.Reasoning)
require.True(t, *metadata.Reasoning)
require.Equal(t, []string{"low", "high"}, metadata.SupportedReasoningLevels)
require.Equal(t, []string{"text", "image"}, metadata.InputModalities)
}
// Scenario: 创建账号 preview 将具体 mapping 传给 404/405 配置回退。
func TestAccountHandlerSyncUpstreamModelsPreviewUsesProvidedModelMapping(t *testing.T) {
upstream := &syncUpstreamHTTPUpstream{responses: []*http.Response{
{
StatusCode: http.StatusNotFound,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"error":"not found"}`)),
},
{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{
"configured-provider": {
"api": "https://provider.example/v1",
"models": {
"glm-5.3": {
"id": "glm-5.3",
"reasoning": true,
"reasoning_options": [{"type":"effort","values":["low","high"]}],
"modalities": {"input":["text"],"output":["text"]},
"limit": {"context":1000000,"output":131072}
}
}
}
}`)),
},
}}
router := setupSyncUpstreamModelsRouter(newStubAdminService(), upstream)
rec := httptest.NewRecorder()
req := httptest.NewRequest(
http.MethodPost,
"/api/v1/admin/accounts/models/sync-upstream-preview",
strings.NewReader(`{
"platform":"openai",
"type":"apikey",
"base_url":"https://provider.example/v1",
"api_key":"key",
"model_mapping":{"public-glm":"glm-5.3"}
}`),
)
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data service.UpstreamModelCatalog `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, []string{"glm-5.3"}, resp.Data.Models)
require.Equal(t, []string{"low", "high"}, resp.Data.Metadata["glm-5.3"].SupportedReasoningLevels)
}
func TestAccountHandlerSyncUpstreamModels_UpstreamErrorDoesNotExposeBody(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
@@ -341,3 +477,52 @@ func TestAccountHandlerSyncUpstreamModels_UpstreamErrorDoesNotExposeBody(t *test
require.Contains(t, rec.Body.String(), "Upstream model list request failed with HTTP 502")
require.NotContains(t, rec.Body.String(), "SECRET_TOKEN")
}
// Scenario: 能力补全失败显示部分成功。
func TestAccountHandlerSyncUpstreamModels_MetadataEnrichmentFailureReturnsWarning(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 46,
Name: "opencode-id-only-model-list",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Credentials: map[string]any{
"api_key": "opencode-key",
"base_url": "https://opencode.ai/zen/v1",
},
},
}
upstream := &syncUpstreamHTTPUpstream{responses: []*http.Response{
{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"x-preview-f-free"}]}`)),
},
{
StatusCode: http.StatusBadGateway,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"error":"registry unavailable"}`)),
},
}}
router := setupSyncUpstreamModelsRouter(svc, upstream)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/46/models/sync-upstream", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data struct {
Models []string `json:"models"`
Warnings []struct {
Code string `json:"code"`
} `json:"warnings"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, []string{"x-preview-f-free"}, resp.Data.Models)
require.Len(t, resp.Data.Warnings, 1)
require.Equal(t, "upstream_model_metadata_incomplete", resp.Data.Warnings[0].Code)
}
@@ -64,6 +64,8 @@ type channelModelPricingRequest struct {
OutputPrice *float64 `json:"output_price" binding:"omitempty,min=0"`
CacheWritePrice *float64 `json:"cache_write_price" binding:"omitempty,min=0"`
CacheReadPrice *float64 `json:"cache_read_price" binding:"omitempty,min=0"`
FastMultiplier *float64 `json:"fast_multiplier" binding:"omitempty,gt=0"`
FlexMultiplier *float64 `json:"flex_multiplier" binding:"omitempty,gt=0"`
ImageInputPrice *float64 `json:"image_input_price" binding:"omitempty,min=0"`
ImageOutputPrice *float64 `json:"image_output_price" binding:"omitempty,min=0"`
PerRequestPrice *float64 `json:"per_request_price" binding:"omitempty,min=0"`
@@ -72,8 +74,9 @@ type channelModelPricingRequest struct {
}
type channelTimePricingRequest struct {
Timezone string `json:"timezone"`
Periods []channelTimePricingPeriodRequest `json:"periods"`
Timezone string `json:"timezone"`
WeekdaysOnly bool `json:"weekdays_only"`
Periods []channelTimePricingPeriodRequest `json:"periods"`
}
type channelTimePricingPeriodRequest struct {
@@ -83,15 +86,19 @@ type channelTimePricingPeriodRequest struct {
}
type pricingIntervalRequest struct {
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
InputMultiplier *float64 `json:"input_multiplier" binding:"omitempty,gt=0"`
OutputMultiplier *float64 `json:"output_multiplier" binding:"omitempty,gt=0"`
CacheWriteMultiplier *float64 `json:"cache_write_multiplier" binding:"omitempty,gt=0"`
CacheReadMultiplier *float64 `json:"cache_read_multiplier" binding:"omitempty,gt=0"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
}
type accountStatsPricingRuleRequest struct {
@@ -128,6 +135,8 @@ type channelModelPricingResponse struct {
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
FastMultiplier *float64 `json:"fast_multiplier"`
FlexMultiplier *float64 `json:"flex_multiplier"`
ImageInputPrice *float64 `json:"image_input_price"`
ImageOutputPrice *float64 `json:"image_output_price"`
PerRequestPrice *float64 `json:"per_request_price"`
@@ -136,8 +145,9 @@ type channelModelPricingResponse struct {
}
type channelTimePricingResponse struct {
Timezone string `json:"timezone"`
Periods []channelTimePricingPeriodResponse `json:"periods"`
Timezone string `json:"timezone"`
WeekdaysOnly bool `json:"weekdays_only"`
Periods []channelTimePricingPeriodResponse `json:"periods"`
}
type channelTimePricingPeriodResponse struct {
@@ -147,16 +157,20 @@ type channelTimePricingPeriodResponse struct {
}
type pricingIntervalResponse struct {
ID int64 `json:"id"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label,omitempty"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
ID int64 `json:"id"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label,omitempty"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
InputMultiplier *float64 `json:"input_multiplier"`
OutputMultiplier *float64 `json:"output_multiplier"`
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
}
type accountStatsPricingRuleResponse struct {
@@ -248,6 +262,8 @@ func pricingToResponse(p *service.ChannelModelPricing) channelModelPricingRespon
OutputPrice: p.OutputPrice,
CacheWritePrice: p.CacheWritePrice,
CacheReadPrice: p.CacheReadPrice,
FastMultiplier: p.FastMultiplier,
FlexMultiplier: p.FlexMultiplier,
ImageInputPrice: p.ImageInputPrice,
ImageOutputPrice: p.ImageOutputPrice,
PerRequestPrice: p.PerRequestPrice,
@@ -268,25 +284,33 @@ func timePricingToResponse(value *service.ChannelTimePricing) *channelTimePricin
Multiplier: period.Multiplier,
})
}
return &channelTimePricingResponse{Timezone: value.Timezone, Periods: periods}
return &channelTimePricingResponse{
Timezone: value.Timezone,
WeekdaysOnly: value.WeekdaysOnly,
Periods: periods,
}
}
func intervalToResponse(iv service.PricingInterval) pricingIntervalResponse {
return pricingIntervalResponse{
ID: iv.ID,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
ID: iv.ID,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
InputMultiplier: iv.InputMultiplier,
OutputMultiplier: iv.OutputMultiplier,
CacheWriteMultiplier: iv.CacheWriteMultiplier,
CacheReadMultiplier: iv.CacheReadMultiplier,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
}
}
func pricingRequestToService(reqs []channelModelPricingRequest) []service.ChannelModelPricing {
func pricingRequestToService(reqs []channelModelPricingRequest, allowChannelMultipliers bool) []service.ChannelModelPricing {
result := make([]service.ChannelModelPricing, 0, len(reqs))
for _, r := range reqs {
billingMode := service.BillingMode(r.BillingMode)
@@ -296,18 +320,34 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
platform := r.Platform
intervals := make([]service.PricingInterval, 0, len(r.Intervals))
for _, iv := range r.Intervals {
var inputMultiplier, outputMultiplier, cacheWriteMultiplier, cacheReadMultiplier *float64
if allowChannelMultipliers {
inputMultiplier = iv.InputMultiplier
outputMultiplier = iv.OutputMultiplier
cacheWriteMultiplier = iv.CacheWriteMultiplier
cacheReadMultiplier = iv.CacheReadMultiplier
}
intervals = append(intervals, service.PricingInterval{
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
InputMultiplier: inputMultiplier,
OutputMultiplier: outputMultiplier,
CacheWriteMultiplier: cacheWriteMultiplier,
CacheReadMultiplier: cacheReadMultiplier,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
})
}
var fastMultiplier, flexMultiplier *float64
if allowChannelMultipliers {
fastMultiplier = r.FastMultiplier
flexMultiplier = r.FlexMultiplier
}
result = append(result, service.ChannelModelPricing{
Platform: platform,
Models: r.Models,
@@ -316,6 +356,8 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
OutputPrice: r.OutputPrice,
CacheWritePrice: r.CacheWritePrice,
CacheReadPrice: r.CacheReadPrice,
FastMultiplier: fastMultiplier,
FlexMultiplier: flexMultiplier,
ImageInputPrice: r.ImageInputPrice,
ImageOutputPrice: r.ImageOutputPrice,
PerRequestPrice: r.PerRequestPrice,
@@ -338,7 +380,11 @@ func timePricingRequestToService(value *channelTimePricingRequest) *service.Chan
Multiplier: period.Multiplier,
})
}
return &service.ChannelTimePricing{Timezone: value.Timezone, Periods: periods}
return &service.ChannelTimePricing{
Timezone: value.Timezone,
WeekdaysOnly: value.WeekdaysOnly,
Periods: periods,
}
}
func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) service.AccountStatsPricingRule {
@@ -346,7 +392,7 @@ func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) s
Name: r.Name,
GroupIDs: r.GroupIDs,
AccountIDs: r.AccountIDs,
Pricing: pricingRequestToService(r.Pricing),
Pricing: pricingRequestToService(r.Pricing, false),
}
}
@@ -407,7 +453,7 @@ func (h *ChannelHandler) Create(c *gin.Context) {
return
}
pricing := pricingRequestToService(req.ModelPricing)
pricing := pricingRequestToService(req.ModelPricing, true)
// Main model_pricing requires a platform; default to anthropic for backward compatibility.
for i := range pricing {
if pricing[i].Platform == "" {
@@ -481,7 +527,7 @@ func (h *ChannelHandler) Update(c *gin.Context) {
ApplyPricingToAccountStats: req.ApplyPricingToAccountStats,
}
if req.ModelPricing != nil {
pricing := pricingRequestToService(*req.ModelPricing)
pricing := pricingRequestToService(*req.ModelPricing, true)
for i := range pricing {
if pricing[i].Platform == "" {
pricing[i].Platform = service.PlatformAnthropic
@@ -305,7 +305,7 @@ func TestPricingRequestToService_Defaults(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := pricingRequestToService([]channelModelPricingRequest{tt.req})
result := pricingRequestToService([]channelModelPricingRequest{tt.req}, true)
require.Len(t, result, 1)
switch tt.wantField {
case "BillingMode":
@@ -332,7 +332,7 @@ func TestPricingRequestToService_WithAllFields(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
r := result[0]
require.Equal(t, "openai", r.Platform)
@@ -373,7 +373,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
require.Len(t, result[0].Intervals, 2)
@@ -396,7 +396,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
}
func TestPricingRequestToService_EmptySlice(t *testing.T) {
result := pricingRequestToService([]channelModelPricingRequest{})
result := pricingRequestToService([]channelModelPricingRequest{}, true)
require.NotNil(t, result)
require.Empty(t, result)
}
@@ -410,7 +410,7 @@ func TestPricingRequestToService_NilPriceFields(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
r := result[0]
require.Nil(t, r.InputPrice)
@@ -426,28 +426,67 @@ func TestPricingRequestToService_TimePricing(t *testing.T) {
Models: []string{"gpt-5"},
BillingMode: "token",
TimePricing: &channelTimePricingRequest{
Timezone: "Asia/Shanghai",
Timezone: "Asia/Shanghai",
WeekdaysOnly: true,
Periods: []channelTimePricingPeriodRequest{{
StartTime: "09:00", EndTime: "12:00", Multiplier: 2,
}},
},
}
got := pricingRequestToService([]channelModelPricingRequest{req})
got := pricingRequestToService([]channelModelPricingRequest{req}, true)
require.Equal(t, "Asia/Shanghai", got[0].TimePricing.Timezone)
require.True(t, got[0].TimePricing.WeekdaysOnly)
require.Equal(t, 2.0, got[0].TimePricing.Periods[0].Multiplier)
}
func TestPricingRequestToService_TimePricingNil(t *testing.T) {
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}})
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}}, true)
require.Nil(t, got[0].TimePricing)
}
// 账号成本统计规则不支持倍率:allowChannelMultipliers=false 时必须丢弃,
// 避免渠道倍率意外污染账号成本口径。
func TestPricingRequestToService_MultipliersGatedByFlag(t *testing.T) {
req := channelModelPricingRequest{
Models: []string{"gpt-5"},
BillingMode: "token",
FastMultiplier: float64Ptr(2.5),
FlexMultiplier: float64Ptr(0.5),
Intervals: []pricingIntervalRequest{{
MinTokens: 272000,
InputMultiplier: float64Ptr(2),
OutputMultiplier: float64Ptr(1.5),
CacheWriteMultiplier: float64Ptr(2),
CacheReadMultiplier: float64Ptr(2),
}},
}
allowed := pricingRequestToService([]channelModelPricingRequest{req}, true)
require.Equal(t, float64Ptr(2.5), allowed[0].FastMultiplier)
require.Equal(t, float64Ptr(0.5), allowed[0].FlexMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].InputMultiplier)
require.Equal(t, float64Ptr(1.5), allowed[0].Intervals[0].OutputMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheWriteMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheReadMultiplier)
dropped := pricingRequestToService([]channelModelPricingRequest{req}, false)
require.Nil(t, dropped[0].FastMultiplier)
require.Nil(t, dropped[0].FlexMultiplier)
require.Nil(t, dropped[0].Intervals[0].InputMultiplier)
require.Nil(t, dropped[0].Intervals[0].OutputMultiplier)
require.Nil(t, dropped[0].Intervals[0].CacheWriteMultiplier)
require.Nil(t, dropped[0].Intervals[0].CacheReadMultiplier)
// 非倍率字段不受开关影响
require.Equal(t, 272000, dropped[0].Intervals[0].MinTokens)
}
func TestPricingToResponse_TimePricing(t *testing.T) {
got := pricingToResponse(&service.ChannelModelPricing{
BillingMode: service.BillingModeToken,
TimePricing: &service.ChannelTimePricing{
Timezone: "Asia/Shanghai",
Timezone: "Asia/Shanghai",
WeekdaysOnly: true,
Periods: []service.ChannelTimePricingPeriod{{
StartTime: "14:00", EndTime: "18:00", Multiplier: 1.25,
}},
@@ -456,6 +495,7 @@ func TestPricingToResponse_TimePricing(t *testing.T) {
require.NotNil(t, got.TimePricing)
require.Equal(t, "Asia/Shanghai", got.TimePricing.Timezone)
require.True(t, got.TimePricing.WeekdaysOnly)
require.Equal(t, 1.25, got.TimePricing.Periods[0].Multiplier)
}
@@ -615,8 +615,10 @@ func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) {
response.BadRequest(c, "grok quota service is not enabled")
return
}
// ResetQuota 恒返回 GROK_QUOTA_RESET_UNSUPPORTED(xAI 无 OAuth 配额重置接口),err != nil 恒真为预期。
//nolint:staticcheck // SA4023
result, err := h.quotaService.ResetQuota(c.Request.Context(), accountID)
if err != nil {
if err != nil { //nolint:staticcheck // SA4023
response.ErrorFrom(c, err)
return
}
@@ -27,6 +27,7 @@ type OpenAIOAuthHandler struct {
type openAIQuotaService interface {
QueryUsage(ctx context.Context, accountID int64) (*service.OpenAIQuotaUsage, error)
CacheResetCreditsSnapshot(ctx context.Context, accountID int64, credits *service.OpenAIRateLimitResetCredits) error
CachePostResetSnapshot(ctx context.Context, accountID int64, usage *service.OpenAIQuotaUsage) error
ResetCredit(ctx context.Context, accountID int64) (*service.OpenAIQuotaResetResult, error)
}
@@ -34,12 +35,6 @@ type openAIAccountStateRecoverer interface {
RecoverAccountState(ctx context.Context, accountID int64, options service.AccountRecoveryOptions) (*service.SuccessfulTestRecoveryResult, error)
}
const (
openAIQuotaResetWarningCacheRefreshFailed = "reset_credit_cache_refresh_failed"
openAIQuotaResetWarningAccountRecoveryFailed = "account_state_recovery_failed"
openAIQuotaResetWarningAccountRefreshFailed = "account_state_refresh_failed"
)
// openAIQuotaResetPostProcessTimeout bounds the work performed AFTER the
// (non-refundable) reset credit has already been consumed upstream. The whole
// request must stay comfortably inside the panel HTTP client timeout, otherwise
@@ -493,6 +488,7 @@ func (h *OpenAIOAuthHandler) QueryQuota(c *gin.Context) {
response.ErrorFrom(c, err)
return
}
service.NotifyOpenAIAutoResetCredit(accountID)
response.Success(c, usage)
}
@@ -523,6 +519,7 @@ func (h *OpenAIOAuthHandler) RefreshQuota(c *gin.Context) {
response.Error(c, http.StatusInternalServerError, "openai quota query returned an empty result")
return
}
service.NotifyOpenAIAutoResetCredit(accountID)
refreshResponse := openAIQuotaRefreshResponse{OpenAIQuotaUsage: *usage}
// A failed snapshot write leaves the previous cache intact — report it as a
@@ -600,54 +597,19 @@ func (h *OpenAIOAuthHandler) ResetQuota(c *gin.Context) {
postCtx, cancelPost := openAIQuotaResetPostProcessContext(c.Request.Context())
defer cancelPost()
// Step 1 — unblocking the account is the whole point of consuming a credit
// (#3672 / #3740), so it runs FIRST and is never gated on the display cache.
// Recovery is DB-only and leaves the manual `schedulable` switch untouched.
if h.rateLimitService == nil {
resetResponse.WarningCode = openAIQuotaResetWarningAccountRecoveryFailed
response.Success(c, resetResponse)
return
postResult := service.RunOpenAIQuotaResetPostProcess(
postCtx,
accountID,
h.quotaService,
h.rateLimitService,
h.adminService.GetAccount,
)
resetResponse.Quota = postResult.Quota
resetResponse.CacheRefreshed = postResult.CacheRefreshed
resetResponse.AccountStateRecovered = postResult.AccountStateRecovered
resetResponse.WarningCode = postResult.WarningCode
if postResult.Account != nil {
resetResponse.Account = dto.AccountFromService(postResult.Account)
}
if _, err := h.rateLimitService.RecoverAccountState(postCtx, accountID, service.AccountRecoveryOptions{
InvalidateToken: true,
}); err != nil {
// Recovery failures are almost always storage-level; the remaining steps
// share that dependency, so stop here instead of compounding the failure.
slog.Warn("openai_quota_reset_account_recovery_failed", "account_id", accountID, "error", err)
resetResponse.WarningCode = openAIQuotaResetWarningAccountRecoveryFailed
response.Success(c, resetResponse)
return
}
resetResponse.AccountStateRecovered = true
// Step 2 — refresh the reset-credit display cache. A failure here is reported
// but must not hide the recovered account row produced by step 3.
usage, usageErr := h.quotaService.QueryUsage(postCtx, accountID)
switch {
case usageErr != nil || usage == nil:
slog.Warn("openai_quota_reset_cache_refresh_failed", "account_id", accountID, "error", usageErr)
resetResponse.WarningCode = openAIQuotaResetWarningCacheRefreshFailed
default:
if err := h.quotaService.CacheResetCreditsSnapshot(postCtx, accountID, usage.RateLimitResetCredits); err != nil {
slog.Warn("openai_quota_reset_cache_refresh_failed", "account_id", accountID, "error", err)
resetResponse.WarningCode = openAIQuotaResetWarningCacheRefreshFailed
} else {
resetResponse.Quota = usage
resetResponse.CacheRefreshed = true
}
}
// Step 3 — hand back the post-recovery account row so the list drops the
// stale rate-limit badge without waiting for the next poll.
account, err := h.adminService.GetAccount(postCtx, accountID)
if err != nil {
slog.Warn("openai_quota_reset_account_refresh_failed", "account_id", accountID, "error", err)
if resetResponse.WarningCode == "" {
resetResponse.WarningCode = openAIQuotaResetWarningAccountRefreshFailed
}
response.Success(c, resetResponse)
return
}
resetResponse.Account = dto.AccountFromService(account)
response.Success(c, resetResponse)
}
@@ -48,6 +48,12 @@ func (s *openAIQuotaWorkflowStub) CacheResetCreditsSnapshot(ctx context.Context,
return s.cacheErr
}
func (s *openAIQuotaWorkflowStub) CachePostResetSnapshot(ctx context.Context, _ int64, _ *service.OpenAIQuotaUsage) error {
s.cacheCalls++
s.cacheCtxErr = ctx.Err()
return s.cacheErr
}
type openAIAccountStateRecovererStub struct {
err error
calls int
@@ -215,7 +221,7 @@ func TestOpenAIResetQuota_RecoveryFailureStopsWorkflow(t *testing.T) {
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningAccountRecoveryFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningAccountRecoveryFailed, envelope.Data.WarningCode)
require.False(t, envelope.Data.AccountStateRecovered)
require.False(t, envelope.Data.CacheRefreshed)
require.Nil(t, envelope.Data.Quota)
@@ -237,7 +243,7 @@ func TestOpenAIResetQuota_MissingRecovererReportsRecoveryFailure(t *testing.T) {
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningAccountRecoveryFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningAccountRecoveryFailed, envelope.Data.WarningCode)
require.False(t, envelope.Data.AccountStateRecovered)
require.Zero(t, quota.queryCalls)
require.Zero(t, adminService.calls)
@@ -260,7 +266,7 @@ func TestOpenAIResetQuota_QueryFailureStillRecoversAndReturnsAccount(t *testing.
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.True(t, envelope.Data.AccountStateRecovered)
require.False(t, envelope.Data.CacheRefreshed)
require.Nil(t, envelope.Data.Quota)
@@ -285,7 +291,7 @@ func TestOpenAIResetQuota_CacheFailureStillRecoversAndReturnsAccount(t *testing.
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.True(t, envelope.Data.AccountStateRecovered)
require.False(t, envelope.Data.CacheRefreshed)
require.Nil(t, envelope.Data.Quota)
@@ -307,7 +313,7 @@ func TestOpenAIResetQuota_AccountRefreshFailureReportsRecoveredState(t *testing.
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningAccountRefreshFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningAccountRefreshFailed, envelope.Data.WarningCode)
require.True(t, envelope.Data.CacheRefreshed)
require.True(t, envelope.Data.AccountStateRecovered)
require.NotNil(t, envelope.Data.Quota)
@@ -331,7 +337,7 @@ func TestOpenAIResetQuota_CacheAndAccountFailureKeepsFirstWarning(t *testing.T)
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.True(t, envelope.Data.AccountStateRecovered)
require.Nil(t, envelope.Data.Account)
}
@@ -435,7 +441,7 @@ func TestOpenAIQuotaEmptyUsageIsHandledWithoutPanic(t *testing.T) {
status, envelope := performOpenAIQuotaResetRequest(t, handler)
require.Equal(t, http.StatusOK, status)
require.Equal(t, openAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.Equal(t, service.OpenAIQuotaResetWarningCacheRefreshFailed, envelope.Data.WarningCode)
require.True(t, envelope.Data.AccountStateRecovered)
require.NotNil(t, envelope.Data.Account)
require.Zero(t, quota.cacheCalls)
@@ -0,0 +1,257 @@
package admin
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"mime"
"net/http"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
const pluginUISessionTTL = 30 * time.Minute
// PluginHandler 提供插件安装、生命周期、配置和隔离 UI 资源接口。
type PluginHandler struct {
manager *service.PluginManager
}
func NewPluginHandler(manager *service.PluginManager) *PluginHandler {
return &PluginHandler{manager: manager}
}
func (h *PluginHandler) List(c *gin.Context) {
plugins, err := h.manager.List(c.Request.Context())
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, plugins)
}
func (h *PluginHandler) Get(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
plugin, err := h.manager.Get(c.Request.Context(), id)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, plugin)
}
func (h *PluginHandler) Upload(c *gin.Context) {
maxBytes := h.manager.MaxUploadBytes()
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, maxBytes+(1<<20))
file, header, err := c.Request.FormFile("plugin")
if err != nil {
response.BadRequest(c, "请选择有效的 .s2plugin 文件")
return
}
defer func() { _ = file.Close() }()
if !strings.HasSuffix(strings.ToLower(header.Filename), ".s2plugin") {
response.BadRequest(c, "插件包扩展名必须是 .s2plugin")
return
}
var installedBy *int64
if subject, ok := middleware.GetAuthSubjectFromContext(c); ok && subject.UserID > 0 {
userID := subject.UserID
installedBy = &userID
}
plugin, err := h.manager.Install(c.Request.Context(), file, installedBy)
if err != nil {
response.BadRequest(c, err.Error())
return
}
response.Created(c, plugin)
}
type pluginEnableRequest struct {
AcceptUntested bool `json:"accept_untested"`
RolloutPercent int `json:"rollout_percent"`
}
func (h *PluginHandler) Enable(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
request := pluginEnableRequest{RolloutPercent: 100}
if err := c.ShouldBindJSON(&request); err != nil {
response.BadRequest(c, "启用参数无效")
return
}
plugin, err := h.manager.Enable(c.Request.Context(), id, request.AcceptUntested, request.RolloutPercent)
if err != nil {
response.BadRequest(c, err.Error())
return
}
response.Success(c, plugin)
}
func (h *PluginHandler) Disable(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
plugin, err := h.manager.Disable(c.Request.Context(), id)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, plugin)
}
func (h *PluginHandler) Delete(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
if err := h.manager.Delete(c.Request.Context(), id); err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, gin.H{"message": "插件已卸载"})
}
func (h *PluginHandler) GetConfig(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
configJSON, err := h.manager.GetConfig(c.Request.Context(), id)
if err != nil {
response.ErrorFrom(c, err)
return
}
c.Data(http.StatusOK, "application/json; charset=utf-8", configJSON)
}
func (h *PluginHandler) SaveConfig(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
decoder := json.NewDecoder(http.MaxBytesReader(c.Writer, c.Request.Body, 4*1024*1024))
decoder.UseNumber()
var value any
if err := decoder.Decode(&value); err != nil {
response.BadRequest(c, "插件配置必须是有效 JSON")
return
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
response.BadRequest(c, "插件配置只能包含一个 JSON 值")
return
}
raw, err := json.Marshal(value)
if err != nil {
response.BadRequest(c, "插件配置无法序列化")
return
}
saved, err := h.manager.SaveConfig(c.Request.Context(), id, raw)
if err != nil {
response.BadRequest(c, err.Error())
return
}
c.Data(http.StatusOK, "application/json; charset=utf-8", saved)
}
func (h *PluginHandler) Test(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
result, err := h.manager.Test(c.Request.Context(), id)
if err != nil {
response.BadRequest(c, err.Error())
return
}
response.Success(c, result)
}
func (h *PluginHandler) CreateUISession(c *gin.Context) {
id, ok := pluginIDParam(c)
if !ok {
return
}
assetToken, expires, err := h.manager.CreateUIAssetToken(c.Request.Context(), id, pluginUISessionTTL)
if err != nil {
response.ErrorFrom(c, err)
return
}
bridgeToken, err := randomPluginToken()
if err != nil {
response.InternalError(c, "创建插件 UI Bridge 失败")
return
}
response.Success(c, gin.H{
"url": fmt.Sprintf("/api/v1/plugin-ui/%s/index.html#bridge_token=%s", assetToken, bridgeToken),
"bridge_token": bridgeToken,
"ui_bridge_version": 1,
"expires_at": expires,
})
}
// ServeUIAsset 使用短时随机能力 URL 提供插件静态资源,不向 iframe 暴露管理员凭据。
func (h *PluginHandler) ServeUIAsset(c *gin.Context) {
token := strings.TrimSpace(c.Param("token"))
pluginID, err := h.manager.ResolveUIAssetToken(token)
if err != nil {
c.Status(http.StatusGone)
return
}
relative := strings.TrimPrefix(c.Param("path"), "/")
data, logicalPath, err := h.manager.ReadUIAsset(c.Request.Context(), pluginID, relative)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
c.Status(http.StatusNotFound)
return
}
c.Status(http.StatusNotFound)
return
}
contentType := mime.TypeByExtension(filepath.Ext(logicalPath))
if contentType == "" {
contentType = "application/octet-stream"
}
c.Header("Cache-Control", "private, no-store")
c.Header("Referrer-Policy", "no-referrer")
c.Header("X-Content-Type-Options", "nosniff")
c.Header("X-Frame-Options", "SAMEORIGIN")
// sandbox iframe 没有 allow-same-origin,会以不透明来源加载自己的 CSS/JS。
// 资源 URL 由短时随机能力 Token 保护,Bridge Token 只存在于 fragment 中。
c.Header("Cross-Origin-Resource-Policy", "cross-origin")
c.Header("Content-Security-Policy", "default-src 'none'; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; img-src 'self' data: blob:; font-src 'self' data:; connect-src 'none'; base-uri 'none'; form-action 'none'; frame-ancestors 'self'; navigate-to 'none'")
c.Data(http.StatusOK, contentType, data)
}
func pluginIDParam(c *gin.Context) (int64, bool) {
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
response.BadRequest(c, "插件 ID 无效")
return 0, false
}
return id, true
}
func randomPluginToken() (string, error) {
buffer := make([]byte, 32)
if _, err := rand.Read(buffer); err != nil {
return "", err
}
return hex.EncodeToString(buffer), nil
}
@@ -382,9 +382,10 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
ModelPlazaEnabled: settings.ModelPlazaEnabled,
ModelPlazaRequireAuth: settings.ModelPlazaRequireAuth,
ModelPlazaDescription: settings.ModelPlazaDescription,
ModelPlazaEnabled: settings.ModelPlazaEnabled,
ModelPlazaRequireAuth: settings.ModelPlazaRequireAuth,
PluginManagementEnabled: settings.PluginManagementEnabled,
ModelPlazaDescription: settings.ModelPlazaDescription,
AffiliateEnabled: settings.AffiliateEnabled,
@@ -347,6 +347,9 @@ type UpdateSettingsRequest struct {
ModelPlazaRequireAuth *bool `json:"model_plaza_require_auth"`
ModelPlazaDescription *string `json:"model_plaza_description"`
// Plugin management menu visibility switch; plugin runtime is unaffected.
PluginManagementEnabled *bool `json:"plugin_management_enabled"`
// Affiliate (邀请返利) feature switch
AffiliateEnabled *bool `json:"affiliate_enabled"`
@@ -438,7 +441,7 @@ func buildSettingKeyByJSONName() map[string]string {
out := make(map[string]string, t.NumField())
for i := 0; i < t.NumField(); i++ {
field := t.Field(i)
if field.Type.Kind() == reflect.Ptr {
if field.Type.Kind() == reflect.Pointer {
continue
}
name, _, _ := strings.Cut(field.Tag.Get("json"), ",")
@@ -1938,6 +1941,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
return previousSettings.ModelPlazaDescription
}(),
PluginManagementEnabled: func() bool {
if req.PluginManagementEnabled != nil {
return *req.PluginManagementEnabled
}
return previousSettings.PluginManagementEnabled
}(),
AffiliateEnabled: func() bool {
if req.AffiliateEnabled != nil {
return *req.AffiliateEnabled
@@ -2357,9 +2366,10 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
AvailableChannelsEnabled: updatedSettings.AvailableChannelsEnabled,
ModelPlazaEnabled: updatedSettings.ModelPlazaEnabled,
ModelPlazaRequireAuth: updatedSettings.ModelPlazaRequireAuth,
ModelPlazaDescription: updatedSettings.ModelPlazaDescription,
ModelPlazaEnabled: updatedSettings.ModelPlazaEnabled,
ModelPlazaRequireAuth: updatedSettings.ModelPlazaRequireAuth,
ModelPlazaDescription: updatedSettings.ModelPlazaDescription,
PluginManagementEnabled: updatedSettings.PluginManagementEnabled,
AffiliateEnabled: updatedSettings.AffiliateEnabled,
+45 -41
View File
@@ -59,30 +59,32 @@ func NewUserHandler(
// CreateUserRequest represents admin create user request
type CreateUserRequest struct {
Email string `json:"email" binding:"required,email"`
Password string `json:"password" binding:"required,min=6"`
Username string `json:"username"`
Notes string `json:"notes"`
Role string `json:"role" binding:"omitempty,oneof=admin user"`
Balance *float64 `json:"balance"`
Concurrency int `json:"concurrency"`
RPMLimit int `json:"rpm_limit"`
AllowedGroups []int64 `json:"allowed_groups"`
Email string `json:"email" binding:"required,email"`
Password string `json:"password" binding:"required,min=6"`
Username string `json:"username"`
Notes string `json:"notes"`
Role string `json:"role" binding:"omitempty,oneof=admin user"`
Balance *float64 `json:"balance"`
Concurrency int `json:"concurrency"`
RPMLimit int `json:"rpm_limit"`
AllowedGroups []int64 `json:"allowed_groups"`
RestrictPublicGroups bool `json:"restrict_public_groups"`
}
// UpdateUserRequest represents admin update user request
// 使用指针类型来区分"未提供"和"设置为0"
type UpdateUserRequest struct {
Email string `json:"email" binding:"omitempty,email"`
Password string `json:"password" binding:"omitempty,min=6"`
Username *string `json:"username"`
Notes *string `json:"notes"`
Role string `json:"role" binding:"omitempty,oneof=admin user"`
Balance *float64 `json:"balance"`
Concurrency *int `json:"concurrency"`
RPMLimit *int `json:"rpm_limit"`
Status string `json:"status" binding:"omitempty,oneof=active disabled"`
AllowedGroups *[]int64 `json:"allowed_groups"`
Email string `json:"email" binding:"omitempty,email"`
Password string `json:"password" binding:"omitempty,min=6"`
Username *string `json:"username"`
Notes *string `json:"notes"`
Role string `json:"role" binding:"omitempty,oneof=admin user"`
Balance *float64 `json:"balance"`
Concurrency *int `json:"concurrency"`
RPMLimit *int `json:"rpm_limit"`
Status string `json:"status" binding:"omitempty,oneof=active disabled"`
AllowedGroups *[]int64 `json:"allowed_groups"`
RestrictPublicGroups *bool `json:"restrict_public_groups"`
// GroupRates 用户专属分组倍率配置
// map[groupID]*rate,nil 表示删除该分组的专属倍率
GroupRates map[int64]*float64 `json:"group_rates"`
@@ -284,16 +286,17 @@ func (h *UserHandler) Create(c *gin.Context) {
}
user, err := h.adminService.CreateUser(c.Request.Context(), &service.CreateUserInput{
Email: req.Email,
Password: req.Password,
Username: req.Username,
Notes: req.Notes,
Role: req.Role,
Balance: req.Balance,
Concurrency: req.Concurrency,
RPMLimit: req.RPMLimit,
AllowedGroups: req.AllowedGroups,
ActorAdminID: getAdminIDFromContext(c),
Email: req.Email,
Password: req.Password,
Username: req.Username,
Notes: req.Notes,
Role: req.Role,
Balance: req.Balance,
Concurrency: req.Concurrency,
RPMLimit: req.RPMLimit,
AllowedGroups: req.AllowedGroups,
RestrictPublicGroups: req.RestrictPublicGroups,
ActorAdminID: getAdminIDFromContext(c),
})
if err != nil {
response.ErrorFrom(c, err)
@@ -342,18 +345,19 @@ func (h *UserHandler) Update(c *gin.Context) {
// 使用指针类型直接传递,nil 表示未提供该字段
user, err := h.adminService.UpdateUser(c.Request.Context(), userID, &service.UpdateUserInput{
Email: req.Email,
Password: req.Password,
Username: req.Username,
Notes: req.Notes,
Role: req.Role,
Balance: req.Balance,
Concurrency: req.Concurrency,
RPMLimit: req.RPMLimit,
Status: req.Status,
AllowedGroups: req.AllowedGroups,
GroupRates: req.GroupRates,
ActorAdminID: getAdminIDFromContext(c),
Email: req.Email,
Password: req.Password,
Username: req.Username,
Notes: req.Notes,
Role: req.Role,
Balance: req.Balance,
Concurrency: req.Concurrency,
RPMLimit: req.RPMLimit,
Status: req.Status,
AllowedGroups: req.AllowedGroups,
RestrictPublicGroups: req.RestrictPublicGroups,
GroupRates: req.GroupRates,
ActorAdminID: getAdminIDFromContext(c),
})
if err != nil {
response.ErrorFrom(c, err)
+2 -2
View File
@@ -1144,10 +1144,10 @@ func (k oidcJWK) publicKey() (any, error) {
if err != nil {
return nil, fmt.Errorf("decode ec y: %w", err)
}
if !curve.IsOnCurve(x, y) {
if !curve.IsOnCurve(x, y) { //nolint:staticcheck // JWK 以裸坐标给出公钥;替换为 ecdsa.ParseUncompressedPublicKey 需改变点编码,待单独迁移
return nil, errors.New("ec point is not on curve")
}
return &ecdsa.PublicKey{Curve: curve, X: x, Y: y}, nil
return &ecdsa.PublicKey{Curve: curve, X: x, Y: y}, nil //nolint:staticcheck // 同上
default:
return nil, fmt.Errorf("unsupported jwk kty: %s", k.Kty)
}
@@ -284,13 +284,13 @@ func toUserSupportedModels(
return out
}
// toUserPricing 将 service 层定价转换为用户 DTO;入参为 nil 时返回 nil。
func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
if p == nil {
// toUserPricingIntervals 将定价区间转换为用户 DTO 白名单形态;nil 入参返回 nil(JSON omitempty 可省略)。
func toUserPricingIntervals(src []service.PricingInterval) []userPricingIntervalDTO {
if src == nil {
return nil
}
intervals := make([]userPricingIntervalDTO, 0, len(p.Intervals))
for _, iv := range p.Intervals {
intervals := make([]userPricingIntervalDTO, 0, len(src))
for _, iv := range src {
intervals = append(intervals, userPricingIntervalDTO{
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
@@ -302,6 +302,19 @@ func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
PerRequestPrice: iv.PerRequestPrice,
})
}
return intervals
}
// toUserPricing 将 service 层定价转换为用户 DTO;入参为 nil 时返回 nil。
func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
if p == nil {
return nil
}
intervals := toUserPricingIntervals(p.Intervals)
if intervals == nil {
// 用户侧定价的 intervals 固定输出数组(空配置为 []),保持既有契约。
intervals = []userPricingIntervalDTO{}
}
billingMode := string(p.BillingMode)
if billingMode == "" {
billingMode = string(service.BillingModeToken)
@@ -72,7 +72,36 @@ func openAIReasoningEffortPolicyForRequest(c *gin.Context, apiKey *service.APIKe
return apiKey.Group.MaxReasoningEffort, apiKey.Group.ReasoningEffortMappings, true
}
func bindRequestedReasoningEffort(c *gin.Context, body []byte, model string) {
if c == nil || c.Request == nil {
return
}
effort := service.CanonicalRequestedReasoningEffort(body, model)
if effort == nil {
return
}
c.Request = c.Request.WithContext(service.WithRequestedReasoningEffort(c.Request.Context(), *effort))
}
func stampOpenAIRequestedReasoningEffort(result *service.OpenAIForwardResult, c *gin.Context) {
if result == nil || result.RequestedReasoningEffort != nil {
return
}
if c == nil || c.Request == nil {
return
}
result.RequestedReasoningEffort = service.RequestedReasoningEffortFromContext(c.Request.Context())
}
func stampForwardRequestedReasoningEffort(result *service.ForwardResult, requested *string) {
if result == nil || result.RequestedReasoningEffort != nil {
return
}
result.RequestedReasoningEffort = requested
}
func applyOpenAIReasoningEffortPolicyForRequest(c *gin.Context, apiKey *service.APIKey, body []byte) ([]byte, bool) {
bindRequestedReasoningEffort(c, body, strings.TrimSpace(gjson.GetBytes(body, "model").String()))
maxEffort, mappings, ok := openAIReasoningEffortPolicyForRequest(c, apiKey)
if !ok {
return body, false
@@ -84,6 +113,7 @@ func bindOpenAIReasoningEffortPolicyForMessagesRequest(c *gin.Context, apiKey *s
if c == nil || c.Request == nil {
return
}
bindRequestedReasoningEffort(c, body, strings.TrimSpace(gjson.GetBytes(body, "model").String()))
// The Messages bridge synthesizes a default OpenAI effort when
// output_config.effort is omitted. Bind the group policy only for an
// explicit client value so the ceiling does not alter that default.
@@ -31,6 +31,7 @@ func TestOpenAICompatibleTextTargetAllowsCompositeProviders(t *testing.T) {
}{
{model: "grok-4.3", platform: service.PlatformGrok},
{model: "kimi-k2-thinking", platform: service.PlatformKimi},
{model: "k3", platform: service.PlatformKimi},
{model: "glm-5.2", platform: service.PlatformZhipu},
{model: "deepseek-v3.2", platform: service.PlatformDeepseek},
}
@@ -119,6 +120,9 @@ func TestOpenAIReasoningEffortPolicyForCompositeTarget(t *testing.T) {
got, changed := applyOpenAIReasoningEffortPolicyForRequest(openAICtx, apiKey, body)
require.True(t, changed)
require.JSONEq(t, `{"reasoning":{"effort":"medium"}}`, string(got))
requested := service.RequestedReasoningEffortFromContext(openAICtx.Request.Context())
require.NotNil(t, requested)
require.Equal(t, "max", *requested)
bindOpenAIReasoningEffortPolicyForMessagesRequest(openAICtx, apiKey, []byte(`{"output_config":{"effort":"max"}}`))
bound, changed := service.ApplyOpenAIReasoningEffortPolicyFromContext(openAICtx.Request.Context(), body)
+51 -16
View File
@@ -3,6 +3,7 @@ package dto
import (
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
@@ -68,10 +69,11 @@ func UserFromServiceAdmin(u *service.User) *AdminUser {
return nil
}
return &AdminUser{
User: *base,
Notes: u.Notes,
LastUsedAt: u.LastUsedAt,
GroupRates: u.GroupRates,
User: *base,
Notes: u.Notes,
LastUsedAt: u.LastUsedAt,
GroupRates: u.GroupRates,
RestrictPublicGroups: u.RestrictPublicGroups,
}
}
@@ -643,7 +645,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
RequestID: l.RequestID,
Model: requestedModel,
ServiceTier: l.ServiceTier,
ReasoningEffort: l.ReasoningEffort,
ReasoningEffort: userFacingReasoningEffort(l),
InboundEndpoint: l.InboundEndpoint,
GroupID: l.GroupID,
SubscriptionID: l.SubscriptionID,
@@ -710,20 +712,53 @@ func UsageLogFromServiceAdmin(l *service.UsageLog) *AdminUsageLog {
usageLog := usageLogFromServiceUser(l)
usageLog.UpstreamEndpoint = l.UpstreamEndpoint
return &AdminUsageLog{
UsageLog: usageLog,
UpstreamModel: l.UpstreamModel,
UpstreamResponseModel: l.UpstreamResponseModel,
UpstreamModelMismatch: l.UpstreamModelMismatch,
ChannelID: l.ChannelID,
ModelMappingChain: l.ModelMappingChain,
BillingTier: l.BillingTier,
AccountRateMultiplier: l.AccountRateMultiplier,
AccountStatsCost: l.AccountStatsCost,
IPAddress: l.IPAddress,
Account: AccountSummaryFromService(l.Account),
UsageLog: usageLog,
UpstreamModel: l.UpstreamModel,
UpstreamReasoningEffort: adminUpstreamReasoningEffort(l),
UpstreamResponseModel: l.UpstreamResponseModel,
UpstreamModelMismatch: l.UpstreamModelMismatch,
ChannelID: l.ChannelID,
ModelMappingChain: l.ModelMappingChain,
BillingTier: l.BillingTier,
AccountRateMultiplier: l.AccountRateMultiplier,
AccountStatsCost: l.AccountStatsCost,
IPAddress: l.IPAddress,
Account: AccountSummaryFromService(l.Account),
}
}
func userFacingReasoningEffort(l *service.UsageLog) *string {
if l == nil {
return nil
}
if requested := strings.TrimSpace(derefString(l.RequestedReasoningEffort)); requested != "" {
return &requested
}
return l.ReasoningEffort
}
func adminUpstreamReasoningEffort(l *service.UsageLog) *string {
if l == nil {
return nil
}
forwarded := strings.TrimSpace(derefString(l.ReasoningEffort))
if forwarded == "" {
return nil
}
requested := userFacingReasoningEffort(l)
if requested != nil && service.NormalizeMaxReasoningEffort(*requested) == service.NormalizeMaxReasoningEffort(forwarded) {
return nil
}
return &forwarded
}
func derefString(value *string) string {
if value == nil {
return ""
}
return *value
}
func UsageCleanupTaskFromService(task *service.UsageCleanupTask) *UsageCleanupTask {
if task == nil {
return nil
@@ -179,6 +179,61 @@ func TestUsageLogFromService_KeepsUserBillingAndIPWithoutAdminCostFields(t *test
require.NotContains(t, string(userJSON), "account_cost")
}
func TestUsageLogFromService_UsersSeeRequestedReasoningEffortOnly(t *testing.T) {
t.Parallel()
requested := "max"
forwarded := "xhigh"
log := &service.UsageLog{
RequestID: "req_effort",
Model: "gpt-5.4",
ReasoningEffort: &forwarded,
RequestedReasoningEffort: &requested,
}
userDTO := UsageLogFromService(log)
adminDTO := UsageLogFromServiceAdmin(log)
require.NotNil(t, userDTO.ReasoningEffort)
require.Equal(t, requested, *userDTO.ReasoningEffort)
require.NotNil(t, adminDTO.ReasoningEffort)
require.Equal(t, requested, *adminDTO.ReasoningEffort)
require.NotNil(t, adminDTO.UpstreamReasoningEffort)
require.Equal(t, forwarded, *adminDTO.UpstreamReasoningEffort)
userJSON, err := json.Marshal(userDTO)
require.NoError(t, err)
require.Contains(t, string(userJSON), `"reasoning_effort":"max"`)
require.NotContains(t, string(userJSON), "upstream_reasoning_effort")
require.NotContains(t, string(userJSON), "requested_reasoning_effort")
adminJSON, err := json.Marshal(adminDTO)
require.NoError(t, err)
require.Contains(t, string(adminJSON), `"reasoning_effort":"max"`)
require.Contains(t, string(adminJSON), `"upstream_reasoning_effort":"xhigh"`)
}
func TestUsageLogFromService_OmitsUpstreamReasoningEffortWhenUnmapped(t *testing.T) {
t.Parallel()
effort := "high"
log := &service.UsageLog{
RequestID: "req_effort_same",
Model: "gpt-5.4",
ReasoningEffort: &effort,
RequestedReasoningEffort: &effort,
}
adminDTO := UsageLogFromServiceAdmin(log)
require.NotNil(t, adminDTO.ReasoningEffort)
require.Equal(t, effort, *adminDTO.ReasoningEffort)
require.Nil(t, adminDTO.UpstreamReasoningEffort)
adminJSON, err := json.Marshal(adminDTO)
require.NoError(t, err)
require.NotContains(t, string(adminJSON), "upstream_reasoning_effort")
}
func TestUsageLogFromService_FallsBackToLegacyModelWhenRequestedModelMissing(t *testing.T) {
t.Parallel()
+7 -5
View File
@@ -316,9 +316,10 @@ type SystemSettings struct {
AvailableChannelsEnabled bool `json:"available_channels_enabled"`
// Model Plaza feature (public group/model pricing showcase)
ModelPlazaEnabled bool `json:"model_plaza_enabled"`
ModelPlazaRequireAuth bool `json:"model_plaza_require_auth"`
ModelPlazaDescription string `json:"model_plaza_description"`
ModelPlazaEnabled bool `json:"model_plaza_enabled"`
ModelPlazaRequireAuth bool `json:"model_plaza_require_auth"`
ModelPlazaDescription string `json:"model_plaza_description"`
PluginManagementEnabled bool `json:"plugin_management_enabled"`
// 风控中心功能开关
RiskControlEnabled bool `json:"risk_control_enabled"`
@@ -418,8 +419,9 @@ type PublicSettings struct {
AvailableChannelsEnabled bool `json:"available_channels_enabled"`
ModelPlazaEnabled bool `json:"model_plaza_enabled"`
ModelPlazaRequireAuth bool `json:"model_plaza_require_auth"`
ModelPlazaEnabled bool `json:"model_plaza_enabled"`
ModelPlazaRequireAuth bool `json:"model_plaza_require_auth"`
PluginManagementEnabled bool `json:"plugin_management_enabled"`
AffiliateEnabled bool `json:"affiliate_enabled"`
+8 -1
View File
@@ -48,6 +48,9 @@ type AdminUser struct {
// GroupRates 用户专属分组倍率配置
// map[groupID]rateMultiplier
GroupRates map[int64]float64 `json:"group_rates,omitempty"`
// RestrictPublicGroups 为 true 时,该用户仅可使用 allowed_groups 中列出的
// 公开分组。这是管理侧的权限开关,不下发给用户自身的接口。
RestrictPublicGroups bool `json:"restrict_public_groups"`
}
type APIKey struct {
@@ -486,8 +489,9 @@ type UsageLog struct {
Model string `json:"model"`
// ServiceTier records the OpenAI service tier used for billing, e.g. "priority" / "flex".
ServiceTier *string `json:"service_tier,omitempty"`
// ReasoningEffort is the request's reasoning effort level.
// ReasoningEffort is the client-requested effort (mapping-hidden, like Model).
// OpenAI: "low"/"medium"/"high"/"xhigh"; Claude: "low"/"medium"/"high"/"max".
// Historical rows without requested_reasoning_effort fall back to the stored effective value.
ReasoningEffort *string `json:"reasoning_effort,omitempty"`
// InboundEndpoint is the client-facing API endpoint path, e.g. /v1/chat/completions.
InboundEndpoint *string `json:"inbound_endpoint,omitempty"`
@@ -563,6 +567,9 @@ type AdminUsageLog struct {
// UpstreamModel is the actual model sent to the upstream provider after mapping.
// Omitted when no mapping was applied (requested model was used as-is).
UpstreamModel *string `json:"upstream_model,omitempty"`
// UpstreamReasoningEffort is the effort actually forwarded after group policy /
// model-family remapping. Omitted when it matches the client-requested value.
UpstreamReasoningEffort *string `json:"upstream_reasoning_effort,omitempty"`
// UpstreamResponseModel is the raw model declared by the upstream response.
UpstreamResponseModel *string `json:"upstream_response_model,omitempty"`
// UpstreamModelMismatch is nil when the upstream did not declare a model.
+8
View File
@@ -312,6 +312,14 @@ func GetInboundEndpoint(c *gin.Context) string {
// and the account platform. Handlers call this after scheduling an
// account, passing account.Platform.
func GetUpstreamEndpoint(c *gin.Context, platform string) string {
// OpenAI 转发服务维护独立的运行时端点上下文,覆盖普通入站推导。
// 这对 force_chat_completions 的错误路径尤为重要:此时可能没有
// ForwardResult,不能把入站 /v1/responses 误报成上游端点。
if platform == service.PlatformOpenAI || platform == service.PlatformGrok || service.IsCNProvider(platform) {
if endpoint := service.GetActualOpenAIUpstreamEndpoint(c); endpoint != "" {
return endpoint
}
}
if c != nil {
if value, ok := c.Get(ctxKeyActualUpstreamEndpoint); ok {
if endpoint, ok := value.(string); ok && endpoint != "" {
+10
View File
@@ -184,6 +184,16 @@ func TestGetUpstreamEndpointPrefersRuntimeOverride(t *testing.T) {
require.Equal(t, EndpointMessages, GetUpstreamEndpoint(c, service.PlatformAntigravity))
}
func TestGetUpstreamEndpointUsesOpenAIRuntimeOverride(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil)
c.Set(ctxKeyInboundEndpoint, EndpointResponses)
service.SetActualOpenAIUpstreamEndpoint(c, EndpointChatCompletions)
require.Equal(t, EndpointChatCompletions, GetUpstreamEndpoint(c, service.PlatformOpenAI))
}
func TestResolveOpenAIUpstreamEndpointPrefersForwardResult(t *testing.T) {
tests := []struct {
name string
+55 -3
View File
@@ -56,7 +56,13 @@ const (
const profitVetoExhaustedMessage = "No available accounts: all candidates rejected by group profit control"
func sameAccountRetryDelayFor(failoverErr *service.UpstreamFailoverError, retryCount int) time.Duration {
if failoverErr == nil || !failoverErr.RequestScopedTransient || retryCount <= 1 {
if failoverErr == nil {
return sameAccountRetryDelay
}
if failoverErr.SameAccountRetryDelay > 0 {
return failoverErr.SameAccountRetryDelay
}
if !failoverErr.RequestScopedTransient || retryCount <= 1 {
return sameAccountRetryDelay
}
@@ -70,6 +76,51 @@ func sameAccountRetryDelayFor(failoverErr *service.UpstreamFailoverError, retryC
return delay
}
func sameAccountRetryAllowed(failoverErr *service.UpstreamFailoverError, retryCount, retryLimit int) bool {
if failoverErr == nil || !failoverErr.RetryableOnSameAccount {
return false
}
if !sameAccountRetryDeadlineAllows(failoverErr) {
return false
}
// Error-specific caps (Grok capacity/stream-idle) remain hard limits even
// when the error also carries a freshly reconstructed deadline.
if failoverErr.SameAccountRetryMax > 0 {
if retryLimit <= 0 {
return false
}
if failoverErr.SameAccountRetryMax < retryLimit {
retryLimit = failoverErr.SameAccountRetryMax
}
return retryCount < retryLimit
}
// OAuth 429 explicitly opts into a deadline window. It is intentionally not
// bounded by the ordinary/default pool retry count.
if !failoverErr.SameAccountRetryDeadline.IsZero() {
return true
}
return retryLimit > 0 && retryCount < retryLimit
}
// sameAccountRetryDeadlineAllows prevents a retry from starting after the
// service-provided same-account retry window has elapsed.
func sameAccountRetryDeadlineAllows(failoverErr *service.UpstreamFailoverError) bool {
return failoverErr == nil || failoverErr.SameAccountRetryDeadline.IsZero() || time.Now().Before(failoverErr.SameAccountRetryDeadline)
}
// effectiveSameAccountRetryLimit applies an error-specific cap without
// overriding an explicit account setting of zero (which disables retries).
func effectiveSameAccountRetryLimit(failoverErr *service.UpstreamFailoverError, account *service.Account) int {
if account == nil {
return 0
}
limit := account.GetPoolModeRetryCount()
if limit > 0 && failoverErr != nil && failoverErr.SameAccountRetryMax > 0 && failoverErr.SameAccountRetryMax < limit {
return failoverErr.SameAccountRetryMax
}
return limit
}
// FailoverState 跨循环迭代共享的 failover 状态
type FailoverState struct {
SwitchCount int
@@ -158,14 +209,15 @@ func (s *FailoverState) HandleFailoverError(
}
// 同账号重试不算切换账号,粘性会话仅在实际切换时强制缓存计费。
sameAccountRetry := failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < retryLimit
retryCount := s.SameAccountRetryCount[accountID]
sameAccountRetry := sameAccountRetryAllowed(failoverErr, retryCount, retryLimit)
if needForceCacheBilling(s.hasBoundSession, failoverErr, sameAccountRetry) {
s.ForceCacheBilling = true
}
// 同账号重试:对 RetryableOnSameAccount 的临时性错误,先在同一账号上重试。
// 重试次数上限 retryLimit 由调用方传入(账号级 pool_mode_retry_count 配置)。
if failoverErr.RetryableOnSameAccount && s.SameAccountRetryCount[accountID] < retryLimit {
if sameAccountRetry {
s.SameAccountRetryCount[accountID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, s.SameAccountRetryCount[accountID])
logger.FromContext(ctx).Warn("gateway.failover_same_account_retry",
@@ -58,6 +58,60 @@ func TestSameAccountRetryDelayFor(t *testing.T) {
t.Run("nil error keeps fixed delay", func(t *testing.T) {
require.Equal(t, 500*time.Millisecond, sameAccountRetryDelayFor(nil, 10))
})
t.Run("explicit oauth delay wins", func(t *testing.T) {
err := &service.UpstreamFailoverError{SameAccountRetryDelay: 3 * time.Second}
require.Equal(t, 3*time.Second, sameAccountRetryDelayFor(err, 1))
})
}
func TestSameAccountRetryAllowedUsesDeadlineInsteadOfPoolCount(t *testing.T) {
err := &service.UpstreamFailoverError{
RetryableOnSameAccount: true,
SameAccountRetryDeadline: time.Now().Add(time.Minute),
}
require.True(t, sameAccountRetryAllowed(err, 100, 0))
require.True(t, sameAccountRetryAllowed(err, 100, maxSameAccountRetries))
err.SameAccountRetryDeadline = time.Now().Add(-time.Second)
require.False(t, sameAccountRetryAllowed(err, 0, 100))
}
func TestSameAccountRetryAllowedRequiresOptInAndDefaultsToCountLimit(t *testing.T) {
err := &service.UpstreamFailoverError{SameAccountRetryDeadline: time.Now().Add(time.Minute)}
require.False(t, sameAccountRetryAllowed(err, 0, maxSameAccountRetries))
err.RetryableOnSameAccount = true
err.SameAccountRetryDeadline = time.Time{}
require.True(t, sameAccountRetryAllowed(err, maxSameAccountRetries-1, maxSameAccountRetries))
require.False(t, sameAccountRetryAllowed(err, maxSameAccountRetries, maxSameAccountRetries))
}
func TestSameAccountRetryAllowedHonorsErrorMaxBeforeDeadline(t *testing.T) {
err := &service.UpstreamFailoverError{
RetryableOnSameAccount: true,
SameAccountRetryDeadline: time.Now().Add(time.Minute),
SameAccountRetryMax: 1,
}
require.True(t, sameAccountRetryAllowed(err, 0, maxSameAccountRetries))
require.False(t, sameAccountRetryAllowed(err, 1, maxSameAccountRetries))
require.False(t, sameAccountRetryAllowed(err, 0, 0), "an explicit zero retry budget remains disabled")
}
func TestSameAccountRetryDeadlineAllows(t *testing.T) {
require.True(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{}))
require.True(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{
SameAccountRetryDeadline: time.Now().Add(time.Second),
}))
require.False(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{
SameAccountRetryDeadline: time.Now().Add(-time.Second),
}))
}
func TestEffectiveSameAccountRetryLimitHonorsErrorCapAndDisabledAccount(t *testing.T) {
account := &service.Account{Type: service.AccountTypeAPIKey, Credentials: map[string]any{"pool_mode": true, "pool_mode_retry_count": float64(3)}}
require.Equal(t, 1, effectiveSameAccountRetryLimit(&service.UpstreamFailoverError{SameAccountRetryMax: 1}, account))
account.Credentials["pool_mode_retry_count"] = float64(0)
require.Equal(t, 0, effectiveSameAccountRetryLimit(&service.UpstreamFailoverError{SameAccountRetryMax: 1}, account))
}
// ---------------------------------------------------------------------------
@@ -320,6 +374,22 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) {
require.Zero(t, fs.SwitchCount)
})
t.Run("OAuth deadline存在时不按普通计数切换", func(t *testing.T) {
mock := &mockTempUnscheduler{}
fs := NewFailoverState(3, true)
fs.SameAccountRetryCount[100] = maxSameAccountRetries
err := newTestFailoverErr(http.StatusTooManyRequests, true, false)
err.SameAccountRetryDeadline = time.Now().Add(time.Minute)
err.SameAccountRetryDelay = time.Nanosecond
fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err)
require.False(t, fs.ForceCacheBilling)
require.Zero(t, fs.SwitchCount)
require.Equal(t, maxSameAccountRetries+1, fs.SameAccountRetryCount[100])
require.Empty(t, mock.calls)
})
t.Run("同账号重试耗尽并实际切换时设置ForceCacheBilling", func(t *testing.T) {
mock := &mockTempUnscheduler{}
fs := NewFailoverState(3, true)
+99 -9
View File
@@ -169,6 +169,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
body = parsedReq.Body.Bytes()
reqModel := parsedReq.Model
reqStream := parsedReq.Stream
bindRequestedReasoningEffort(c, body, reqModel)
ensureCompositeTargetPlatform(c, apiKey, reqModel)
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
@@ -541,6 +542,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
stampForwardRequestedReasoningEffort(result, service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort))
if result.ReasoningEffort == nil {
result.ReasoningEffort = service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort)
}
@@ -882,6 +884,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
stampForwardRequestedReasoningEffort(result, service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort))
if result.ReasoningEffort == nil {
result.ReasoningEffort = service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort)
}
@@ -1140,6 +1143,79 @@ func (h *GatewayHandler) Models(c *gin.Context) {
})
}
// CodexModels returns the effective group model list using the manifest shape
// expected by Codex custom providers. Official OpenAI groups continue to use
// OpenAIGatewayHandler.CodexModels so their live upstream metadata is preserved.
func (h *GatewayHandler) CodexModels(c *gin.Context) {
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok || apiKey == nil || apiKey.Group == nil {
h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required")
return
}
forcedPlatform := ""
if value, exists := middleware2.GetForcePlatformFromContext(c); exists {
forcedPlatform = strings.TrimSpace(value)
}
modelIDs := h.codexModelIDsForGroup(c.Request.Context(), apiKey.Group, forcedPlatform)
modelIDs = service.FilterCodexModelIDsForGroup(modelIDs, apiKey.Group)
body, err := h.gatewayService.BuildCodexModelsManifestForGroup(
c.Request.Context(),
apiKey.Group,
forcedPlatform,
modelIDs,
)
if err != nil {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest")
return
}
etag := service.CodexModelsManifestETag(body)
c.Header("ETag", etag)
if service.CodexModelsManifestETagMatches(c.GetHeader("If-None-Match"), etag) {
c.Status(http.StatusNotModified)
c.Writer.WriteHeaderNow()
return
}
c.Data(http.StatusOK, "application/json", body)
}
func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *service.Group, platformOverride string) []string {
if h == nil || h.gatewayService == nil || group == nil {
return nil
}
groupID := &group.ID
platform := strings.TrimSpace(platformOverride)
if platform == "" {
platform = group.Platform
}
if platform == service.PlatformComposite {
availableModels := h.compositeAvailableModels(ctx, groupID)
fallbackModels := defaultCodexModelIDsForPlatform(service.PlatformComposite)
if group.CustomModelsListEnabled() {
return filterModelsByCustomList(availableModels, fallbackModels, group.ModelsListConfig.Models)
}
if len(availableModels) > 0 {
return availableModels
}
return fallbackModels
}
availableModels := h.gatewayService.GetAvailableModels(ctx, groupID, platform)
fallbackModels := defaultCodexModelIDsForPlatform(platform)
if group.CustomModelsListEnabled() {
return filterModelsByCustomList(
customModelsListSource(platform, availableModels, fallbackModels),
fallbackModels,
group.ModelsListConfig.Models,
)
}
if len(availableModels) > 0 {
return availableModels
}
return fallbackModels
}
func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64) []string {
if h == nil || h.gatewayService == nil {
return nil
@@ -1234,11 +1310,15 @@ func writeGrokModelsList(c *gin.Context, modelIDs []string) {
if grokModelSupportsConfigurableReasoning(modelID) {
item.SupportsReasoningEffort = true
item.ReasoningEffort = "high"
item.ReasoningEfforts = []grokReasoningEffortOption{
efforts := []grokReasoningEffortOption{
{Value: "low", Label: "Low"},
{Value: "medium", Label: "Medium"},
{Value: "high", Label: "High", Default: true},
}
if service.GrokSupportsXHighReasoningEffort(modelID) {
efforts = append(efforts, grokReasoningEffortOption{Value: "xhigh", Label: "xHigh"})
}
item.ReasoningEfforts = efforts
}
models = append(models, item)
}
@@ -1340,9 +1420,26 @@ func customModelsListAllowsModel(availablePatterns []string, model string) bool
return true
}
}
normalizedClaudeModel := claude.NormalizeModelID(strings.TrimSuffix(model, "-thinking"))
if normalizedClaudeModel != model {
for _, pattern := range availablePatterns {
if pattern == normalizedClaudeModel {
return true
}
}
}
return false
}
func defaultCodexModelIDsForPlatform(platform string) []string {
switch platform {
case service.PlatformDeepseek:
return []string{"deepseek-v4-pro", "deepseek-v4-flash"}
default:
return defaultModelIDsForPlatform(platform)
}
}
func defaultModelIDsForPlatform(platform string) []string {
switch platform {
case service.PlatformOpenAI:
@@ -1361,14 +1458,7 @@ func defaultModelIDsForPlatform(platform string) []string {
}
return ids
case service.PlatformAnthropic:
ids := make([]string, 0, len(claude.DefaultModels)+len(antigravity.DefaultModels()))
for _, model := range claude.DefaultModels {
ids = append(ids, model.ID)
}
for _, model := range antigravity.DefaultModels() {
ids = append(ids, model.ID)
}
return mergeModelIDs(ids, nil)
return claude.DefaultModelIDs()
case service.PlatformGrok:
return xai.DefaultModelIDs()
case service.PlatformComposite:
@@ -5,6 +5,7 @@ import (
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
@@ -75,6 +76,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
return
}
reqModel := modelResult.String()
bindRequestedReasoningEffort(c, body, reqModel)
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformResolved(c, apiKey, reqModel) {
h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups")
@@ -173,6 +175,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
if err != nil {
if len(fs.FailedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, groupPlatform)
cls = classifySelectionFailureError(err, cls)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -333,6 +336,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
stampForwardRequestedReasoningEffort(result, service.RequestedReasoningEffortFromContext(c.Request.Context()))
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -384,6 +388,14 @@ func (h *GatewayHandler) handleCCFailoverExhausted(c *gin.Context, lastErr *serv
h.chatCompletionsErrorResponse(c, status, "server_error", message)
return
}
if lastErr != nil && lastErr.IsOpenAICapacityShed() && strings.TrimSpace(lastErr.ClientMessage) != "" {
status := lastErr.ClientStatusCode
if status <= 0 {
status = http.StatusServiceUnavailable
}
h.chatCompletionsErrorResponse(c, status, "server_error", lastErr.ClientMessage)
return
}
statusCode := http.StatusBadGateway
if lastErr != nil && lastErr.StatusCode > 0 {
statusCode = lastErr.StatusCode
@@ -5,6 +5,7 @@ import (
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
@@ -75,6 +76,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
return
}
reqModel := modelResult.String()
bindRequestedReasoningEffort(c, body, reqModel)
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformResolved(c, apiKey, reqModel) {
h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups")
@@ -175,6 +177,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
if err != nil {
if len(fs.FailedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, effectiveAPIKeyPlatform(c, apiKey))
cls = classifySelectionFailureError(err, cls)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -321,6 +324,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
stampForwardRequestedReasoningEffort(result, service.RequestedReasoningEffortFromContext(c.Request.Context()))
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -361,25 +365,38 @@ func (h *GatewayHandler) responsesErrorResponse(c *gin.Context, status int, code
// handleResponsesFailoverExhausted writes a failover-exhausted error in Responses format.
func (h *GatewayHandler) handleResponsesFailoverExhausted(c *gin.Context, lastErr *service.UpstreamFailoverError, streamStarted bool) {
if streamStarted {
return // Can't write error after stream started
}
if lastErr != nil {
copyFailoverRetryAfter(c, lastErr.ResponseHeaders)
}
if lastErr != nil && lastErr.IsCredentialFailure() {
status, message := credentialFailoverClientResponse(lastErr)
h.responsesErrorResponse(c, status, "server_error", message)
return
}
statusCode := http.StatusBadGateway
if lastErr != nil && lastErr.StatusCode > 0 {
statusCode = lastErr.StatusCode
}
if lastErr != nil && service.IsOpenAISilentRefusalErrorBody(lastErr.ResponseBody) {
status, code, message := statusCode, "server_error", "All available accounts exhausted"
if lastErr != nil && lastErr.IsCredentialFailure() {
status, message = credentialFailoverClientResponse(lastErr)
} else if lastErr != nil && lastErr.IsOpenAICapacityShed() && strings.TrimSpace(lastErr.ClientMessage) != "" {
status = lastErr.ClientStatusCode
if status <= 0 {
status = http.StatusServiceUnavailable
}
message = lastErr.ClientMessage
} else if lastErr != nil && service.IsOpenAISilentRefusalErrorBody(lastErr.ResponseBody) {
service.SetOpsUpstreamError(c, statusCode, service.OpenAISilentRefusalClientMessage(), "")
h.responsesErrorResponse(c, http.StatusBadGateway, "upstream_error", service.OpenAISilentRefusalClientMessage())
status, code, message = http.StatusBadGateway, "upstream_error", service.OpenAISilentRefusalClientMessage()
} else if lastErr != nil && statusCode == http.StatusTooManyRequests {
status, code, message = http.StatusTooManyRequests, "rate_limit_error", "All available accounts are currently rate-limited. Please retry later."
}
if streamStarted {
// A slot-wait heartbeat commits HTTP 200 before any upstream response.
// In that case a terminal frame is still required; once any semantic or
// official terminal bytes exist, preserve them without appending a second
// generic response.failed.
service.MarkOpsStreamError(c, code, message, status)
if c != nil && c.Writer != nil && (c.Writer.Size() <= 0 || gatewayStreamHasOnlyHeartbeats(c)) {
writeResponsesFailedSSE(c, code, message)
}
return
}
h.responsesErrorResponse(c, statusCode, "server_error", "All available accounts exhausted")
h.responsesErrorResponse(c, status, code, message)
}
+26 -1
View File
@@ -16,6 +16,29 @@ import (
"github.com/gin-gonic/gin"
)
const gatewayStreamHeartbeatBytesKey = "gateway_stream_heartbeat_bytes"
func recordGatewayStreamHeartbeat(c *gin.Context, written int) {
if c == nil || written <= 0 {
return
}
total, _ := c.Get(gatewayStreamHeartbeatBytesKey)
bytes, _ := total.(int)
c.Set(gatewayStreamHeartbeatBytesKey, bytes+written)
}
func gatewayStreamHasOnlyHeartbeats(c *gin.Context) bool {
if c == nil || c.Writer == nil {
return false
}
value, ok := c.Get(gatewayStreamHeartbeatBytesKey)
if !ok {
return false
}
heartbeatBytes, _ := value.(int)
return heartbeatBytes > 0 && c.Writer.Size() == heartbeatBytes
}
// claudeCodeValidator is a singleton validator for Claude Code client detection
var claudeCodeValidator = service.NewClaudeCodeValidator()
@@ -396,9 +419,11 @@ func (h *ConcurrencyHelper) waitForSlotWithPingTimeout(c *gin.Context, slotType
c.Header("X-Accel-Buffering", "no")
*streamStarted = true
}
if _, err := fmt.Fprint(c.Writer, string(h.pingFormat)); err != nil {
written, err := fmt.Fprint(c.Writer, string(h.pingFormat))
if err != nil {
return nil, err
}
recordGatewayStreamHeartbeat(c, written)
flusher.Flush()
case <-timer.C:
+433 -8
View File
@@ -25,6 +25,22 @@ type gatewayModelsResponseForTest struct {
Data []gatewayModelItemForTest `json:"data"`
}
type codexModelsResponseForTest struct {
Models []struct {
Slug string `json:"slug"`
SupportedReasoningLevels []codexReasoningLevelForTest `json:"supported_reasoning_levels"`
InputModalities []string `json:"input_modalities"`
ModelMessages map[string]json.RawMessage `json:"model_messages"`
TruncationPolicy map[string]json.RawMessage `json:"truncation_policy"`
AvailabilityNUX json.RawMessage `json:"availability_nux"`
Upgrade json.RawMessage `json:"upgrade"`
} `json:"models"`
}
type codexReasoningLevelForTest struct {
Effort string `json:"effort"`
}
type gatewayModelItemForTest struct {
ID string `json:"id"`
Object string `json:"object"`
@@ -52,6 +68,10 @@ func (s *gatewayModelsAccountRepoStub) ListSchedulableByGroupID(ctx context.Cont
return out, nil
}
func (s *gatewayModelsAccountRepoStub) ListByGroup(ctx context.Context, groupID int64) ([]service.Account, error) {
return s.ListSchedulableByGroupID(ctx, groupID)
}
func newGatewayModelsHandlerForTest(repo service.AccountRepository) *GatewayHandler {
return &GatewayHandler{
gatewayService: service.NewGatewayService(
@@ -70,6 +90,247 @@ func TestDefaultModelIDsForCompositeIncludesAntigravityDefaults(t *testing.T) {
require.Contains(t, compositeIDs, antigravityIDs[0])
}
// Scenario: Anthropic defaults contain only Claude while Antigravity keeps its own Gemini models.
func TestDefaultModelIDsForAnthropicExcludeAntigravityGemini(t *testing.T) {
anthropicIDs := defaultModelIDsForPlatform(service.PlatformAnthropic)
require.Contains(t, anthropicIDs, "claude-opus-4-6")
require.NotContains(t, anthropicIDs, "gemini-2.5-flash")
antigravityIDs := defaultModelIDsForPlatform(service.PlatformAntigravity)
require.Contains(t, antigravityIDs, "gemini-2.5-flash")
}
// Scenario: non-OpenAI groups return a Codex manifest instead of a standard model list.
func TestGatewayCodexModels_NonOpenAIGroupsUseMappedModels(t *testing.T) {
tests := []struct {
name string
platform string
model string
efforts []string
modalities []string
}{
{
name: "Grok",
platform: service.PlatformGrok,
model: "grok-4.6",
efforts: []string{"low", "medium", "high", "xhigh"},
modalities: []string{"text", "image"},
},
{
name: "DeepSeek",
platform: service.PlatformDeepseek,
model: "deepseek-v4-pro",
efforts: []string{"low", "high", "max"},
modalities: []string{"text"},
},
{
name: "provider-qualified Claude",
platform: service.PlatformAnthropic,
model: "anthropic/claude-sonnet-4-6",
efforts: []string{"low", "medium", "high", "max"},
modalities: []string{"text"},
},
}
for index, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(100 + index)
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: tt.platform,
Credentials: map[string]any{
"model_mapping": map[string]any{tt.model: tt.model},
},
},
},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: tt.platform},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Len(t, got.Models, 1)
require.Equal(t, tt.model, got.Models[0].Slug)
require.NotEmpty(t, got.Models[0].ModelMessages)
require.NotEmpty(t, got.Models[0].TruncationPolicy)
require.NotNil(t, got.Models[0].AvailabilityNUX)
require.NotNil(t, got.Models[0].Upgrade)
require.Equal(t, tt.efforts, codexReasoningEffortsForTest(got.Models[0].SupportedReasoningLevels))
require.Equal(t, tt.modalities, got.Models[0].InputModalities)
})
}
}
// Scenario: Composite manifests aggregate only administrator-configured models.
func TestGatewayCodexModels_CompositeUsesCompleteEffectiveModelList(t *testing.T) {
gin.SetMode(gin.TestMode)
const groupID int64 = 120
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 3,
Platform: service.PlatformOpenAI,
Status: service.StatusActive,
Schedulable: true,
Credentials: map[string]any{},
},
{
ID: 1,
Platform: service.PlatformOpenAI,
Credentials: map[string]any{
"model_mapping": map[string]any{"gpt-5.5": "gpt-5.5"},
},
},
{
ID: 2,
Platform: service.PlatformGrok,
Credentials: map[string]any{
"model_mapping": map[string]any{"grok-4.6": "grok-4.6"},
},
},
},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: service.PlatformComposite},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Equal(t, []string{"gpt-5.5", "grok-4.6"}, codexModelSlugsForTest(got.Models))
}
func TestGatewayCodexModels_GeneratedManifestUsesFinalBodyETag(t *testing.T) {
gin.SetMode(gin.TestMode)
const groupID int64 = 122
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {{
ID: 1,
Platform: service.PlatformDeepseek,
Credentials: map[string]any{
"model_mapping": map[string]any{"deepseek-v4-pro": "deepseek-v4-pro"},
},
}},
},
})
group := &service.Group{ID: groupID, Platform: service.PlatformDeepseek}
first := httptest.NewRecorder()
firstContext, _ := gin.CreateTestContext(first)
firstContext.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
firstContext.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{Group: group})
h.CodexModels(firstContext)
require.Equal(t, http.StatusOK, first.Code)
etag := first.Header().Get("ETag")
require.NotEmpty(t, etag)
require.Equal(t, service.CodexModelsManifestETag(first.Body.Bytes()), etag)
second := httptest.NewRecorder()
secondContext, _ := gin.CreateTestContext(second)
secondContext.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
secondContext.Request.Header.Set("If-None-Match", "W/"+etag)
secondContext.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{Group: group})
h.CodexModels(secondContext)
require.Equal(t, http.StatusNotModified, second.Code)
require.Empty(t, second.Body.Bytes())
require.Equal(t, etag, second.Header().Get("ETag"))
}
// Scenario: group models_list_config limits the generated Codex manifest.
func TestGatewayCodexModels_CustomModelsListFiltersCompositeManifest(t *testing.T) {
gin.SetMode(gin.TestMode)
const groupID int64 = 121
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformOpenAI,
Credentials: map[string]any{
"model_mapping": map[string]any{"gpt-5.5": "gpt-5.5"},
},
},
{
ID: 2,
Platform: service.PlatformGrok,
Credentials: map[string]any{
"model_mapping": map[string]any{"grok-4.6": "grok-4.6"},
},
},
},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{
ID: groupID,
Platform: service.PlatformComposite,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"grok-4.6"},
},
},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Equal(t, []string{"grok-4.6"}, codexModelSlugsForTest(got.Models))
}
func codexModelSlugsForTest(models []struct {
Slug string `json:"slug"`
SupportedReasoningLevels []codexReasoningLevelForTest `json:"supported_reasoning_levels"`
InputModalities []string `json:"input_modalities"`
ModelMessages map[string]json.RawMessage `json:"model_messages"`
TruncationPolicy map[string]json.RawMessage `json:"truncation_policy"`
AvailabilityNUX json.RawMessage `json:"availability_nux"`
Upgrade json.RawMessage `json:"upgrade"`
}) []string {
slugs := make([]string, 0, len(models))
for _, model := range models {
slugs = append(slugs, model.Slug)
}
return slugs
}
func codexReasoningEffortsForTest(levels []codexReasoningLevelForTest) []string {
efforts := make([]string, 0, len(levels))
for _, level := range levels {
efforts = append(efforts, level.Effort)
}
return efforts
}
func TestGatewayModels_GeminiGroupFallsBackToGeminiModels(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -103,9 +364,38 @@ func TestGatewayModels_GeminiGroupFallsBackToGeminiModels(t *testing.T) {
}
func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T) {
assertGrokGatewayReasoningEfforts(t, 4409, "grok-4.5", []gatewayReasoningEffortOptionForTest{
{Value: "low", Label: "Low"},
{Value: "medium", Label: "Medium"},
{Value: "high", Label: "High", Default: true},
})
}
func TestGatewayModels_Grok46AdvertisesXHighReasoningEffortForGrokBuild(t *testing.T) {
xhighEfforts := []gatewayReasoningEffortOptionForTest{
{Value: "low", Label: "Low"},
{Value: "medium", Label: "Medium"},
{Value: "high", Label: "High", Default: true},
{Value: "xhigh", Label: "xHigh"},
}
tests := []struct {
groupID int64
model string
}{
{groupID: 4410, model: "grok-4.6"},
{groupID: 4411, model: "grok-4.6-latest"},
}
for _, tt := range tests {
t.Run(tt.model, func(t *testing.T) {
assertGrokGatewayReasoningEfforts(t, tt.groupID, tt.model, xhighEfforts)
})
}
}
func assertGrokGatewayReasoningEfforts(t *testing.T, groupID int64, modelID string, want []gatewayReasoningEffortOptionForTest) {
t.Helper()
gin.SetMode(gin.TestMode)
groupID := int64(4409)
h := newGatewayModelsHandlerForTest(
&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
@@ -114,7 +404,7 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
ID: 1,
Platform: service.PlatformGrok,
Credentials: map[string]any{
"model_mapping": map[string]any{"grok-4.5": "grok-4.5"},
"model_mapping": map[string]any{modelID: modelID},
},
},
},
@@ -136,14 +426,10 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Len(t, got.Data, 1)
model := got.Data[0]
require.Equal(t, "grok-4.5", model.ID)
require.Equal(t, modelID, model.ID)
require.True(t, model.SupportsReasoningEffort)
require.Equal(t, "high", model.ReasoningEffort)
require.Equal(t, []gatewayReasoningEffortOptionForTest{
{Value: "low", Label: "Low"},
{Value: "medium", Label: "Medium"},
{Value: "high", Label: "High", Default: true},
}, model.ReasoningEfforts)
require.Equal(t, want, model.ReasoningEfforts)
}
func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) {
@@ -193,6 +479,62 @@ func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) {
require.Equal(t, []string{"gemini-2.5-flash"}, modelIDsForTest(got.Data))
}
// Scenario: a Composite group with only Anthropic accounts must not inherit Antigravity Gemini defaults.
func TestGatewayCodexModels_CompositeAnthropicDoesNotAdvertiseAntigravityDefaults(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(64)
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {{ID: 1, Platform: service.PlatformAnthropic}},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.147.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: service.PlatformComposite},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
slugs := codexModelSlugsForTest(got.Models)
require.Contains(t, slugs, "claude-opus-4-6")
require.NotContains(t, slugs, "gemini-2.5-flash")
}
// Scenario: Antigravity retains its own Claude and Gemini defaults inside Composite groups.
func TestGatewayModels_CompositeAntigravityAdvertisesAntigravityDefaults(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(65)
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {{ID: 1, Platform: service.PlatformAntigravity}},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: service.PlatformComposite},
})
h.Models(c)
require.Equal(t, http.StatusOK, rec.Code)
var got gatewayModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
ids := modelIDsForTest(got.Data)
require.Contains(t, ids, "claude-opus-4-6")
require.Contains(t, ids, "gemini-2.5-flash")
}
func TestGatewayModels_CustomModelsListDisabledKeepsOriginalModels(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -457,6 +799,89 @@ func TestDefaultModelIDsForPlatform_CNProvidersKeepClaudeDefaults(t *testing.T)
}
}
func TestDefaultCodexModelIDsForPlatform_DeepSeekUsesDeepSeekModels(t *testing.T) {
require.Equal(t, []string{"deepseek-v4-pro", "deepseek-v4-flash"}, defaultCodexModelIDsForPlatform(service.PlatformDeepseek))
require.Equal(t, defaultModelIDsForPlatform(service.PlatformAnthropic), defaultCodexModelIDsForPlatform(service.PlatformAnthropic))
}
func TestGatewayCodexModels_DeepSeekWithoutMappingUsesDeepSeekDefaults(t *testing.T) {
gin.SetMode(gin.TestMode)
const groupID int64 = 130
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformDeepseek,
Status: service.StatusActive,
Schedulable: true,
Credentials: map[string]any{},
},
},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.150.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: service.PlatformDeepseek},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
slugs := make([]string, 0, len(got.Models))
for _, model := range got.Models {
slugs = append(slugs, model.Slug)
}
require.Contains(t, slugs, "deepseek-v4-pro")
require.Contains(t, slugs, "deepseek-v4-flash")
require.NotContains(t, slugs, "claude-sonnet-4-6")
require.NotContains(t, slugs, "claude-opus-4-6")
}
func TestGatewayCodexModels_OmitsWildcardMappingKeys(t *testing.T) {
gin.SetMode(gin.TestMode)
const groupID int64 = 131
h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformDeepseek,
Credentials: map[string]any{
"model_mapping": map[string]any{
"foo-*": "deepseek-v4-pro",
"deepseek-v4-pro": "deepseek-v4-pro",
},
},
},
},
},
})
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/models?client_version=0.150.0", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{ID: groupID, Platform: service.PlatformDeepseek},
})
h.CodexModels(c)
require.Equal(t, http.StatusOK, rec.Code)
var got codexModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
slugs := make([]string, 0, len(got.Models))
for _, model := range got.Models {
slugs = append(slugs, model.Slug)
}
require.Equal(t, []string{"deepseek-v4-pro"}, slugs)
}
func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMapping(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -569,6 +569,13 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
// 长上下文规则由计费服务统一持有(模型广场展示同源),入口只负责声明自己适用该规则。
var longContextThreshold int
var longContextMultiplier float64
if rule := h.gatewayService.LegacyLongContextRule(service.PlatformGemini); rule != nil {
longContextThreshold = rule.Threshold
longContextMultiplier = rule.Multiplier
}
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
Result: result,
@@ -583,8 +590,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
LongContextThreshold: 200000, // Gemini 200K 阈值
LongContextMultiplier: 2.0, // 超出部分双倍计费
LongContextThreshold: longContextThreshold,
LongContextMultiplier: longContextMultiplier,
ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
SessionID: sessionID,
+75 -34
View File
@@ -44,39 +44,84 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) {
return
}
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
"",
"grok-4.5",
nil,
service.OpenAIUpstreamTransportHTTPSSE,
// Grok only advertises chat_completions + media capabilities on HEAD.
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
false,
service.PlatformGrok,
)
if err != nil || selection == nil || selection.Account == nil {
h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts")
return
}
var streamStarted bool
reqLog := requestLogger(c, "handler.openai_gateway.grok_realtime")
release, slotStatus := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, true, &streamStarted, reqLog)
if slotStatus != openAISlotAcquireOK {
model := c.Query("model")
if strings.TrimSpace(model) == "" {
model = "grok-voice-latest"
}
// Keep the HTTP response uncommitted while selecting and probing an account.
// Realtime is not an HTTP streaming response; using reqStream=true here would
// let the wait queue flush an SSE ping before the WebSocket handshake succeeds.
failed := map[int64]struct{}{}
var selection *service.AccountSelectionResult
var release func()
var token string
var upstream *service.GrokRealtimeUpstream
var candidateSeen bool
for attempts := 0; attempts < 4; attempts++ {
// Realtime's voice model is not a text-model capability. Passing a
// concrete text model here would reject accounts mapped only to an
// older/default text model before the upstream handshake can decide.
// An empty requested model keeps account selection capability-based;
// the actual voice model remains in the upstream WS query below.
candidate, _, selectErr := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(), apiKey.GroupID, "", "", "", failed,
service.OpenAIUpstreamTransportHTTPSSE,
service.OpenAIEndpointCapabilityChatCompletions,
false, false, false, service.PlatformGrok,
)
if selectErr != nil || candidate == nil || candidate.Account == nil {
break
}
candidateSeen = true
account := candidate.Account
var streamStarted bool
var slotStatus openAISlotAcquireResult
release, slotStatus = h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", candidate, false, &streamStarted, reqLog)
if slotStatus != openAISlotAcquireOK {
if slotStatus == openAISlotAcquireFailed {
return
}
failed[account.ID] = struct{}{}
continue
}
var credErr error
token, _, credErr = h.gatewayService.GetRequestCredential(c.Request.Context(), c, account)
if credErr != nil {
release()
release = nil
failed[account.ID] = struct{}{}
continue
}
probeCtx, cancelProbe := context.WithTimeout(c.Request.Context(), service.DefaultGrokRealtimeDialTimeout)
candidateUpstream, openErr := h.gatewayService.OpenGrokRealtime(probeCtx, account, token, model)
cancelProbe()
if openErr != nil {
reqLog.Warn("grok_realtime.pre_accept_failed", zap.Int64("account_id", account.ID), zap.Error(openErr))
statusCode := http.StatusBadGateway
var dialErr *service.GrokRealtimeDialError
if errors.As(openErr, &dialErr) && dialErr.StatusCode > 0 {
statusCode = dialErr.StatusCode
}
h.gatewayService.HandleGrokRealtimeUpstreamError(c.Request.Context(), account, statusCode, []byte(openErr.Error()))
release()
release = nil
failed[account.ID] = struct{}{}
continue
}
selection, upstream = candidate, candidateUpstream
break
}
if selection == nil || selection.Account == nil || release == nil || upstream == nil {
if !candidateSeen {
h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts")
} else {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Grok realtime upstream unavailable")
}
return
}
defer release()
token, _, err := h.gatewayService.GetRequestCredential(c.Request.Context(), c, selection.Account)
if err != nil {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Grok credential unavailable")
return
}
defer func() { _ = upstream.Close() }()
conn, err := coderws.Accept(c.Writer, c.Request, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
@@ -84,12 +129,8 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) {
}
defer func() { _ = conn.CloseNow() }()
model := c.Query("model")
if strings.TrimSpace(model) == "" {
model = "grok-voice-latest"
}
started := time.Now()
audioObserved, proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model)
audioObserved, proxyErr := h.gatewayService.ProxyGrokRealtimeConn(c.Request.Context(), c, conn, upstream)
elapsed := time.Since(started)
if proxyErr != nil {
reqLog.Info("grok_realtime.proxy_failed", zap.Error(proxyErr))
+8 -6
View File
@@ -346,7 +346,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
return
}
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, nil), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, nil), false, nil)
}
if c.Writer.Size() != writerSizeBeforeForward {
h.handleFailoverExhausted(c, failoverErr, true)
@@ -361,19 +361,21 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("grok_media.pool_mode_same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_limit", retryLimit),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
zap.Duration("retry_delay", retryDelay),
)
select {
case <-requestCtx.Done():
return
case <-time.After(sameAccountRetryDelay):
case <-time.After(retryDelay):
}
continue
}
@@ -398,7 +400,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, nil), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, nil), false, nil)
if !service.IsResponseCommitted(c) && c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
@@ -409,7 +411,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, result), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, grokMediaScheduleModel(account, routingModel, result), true, nil)
if isGrokVideoCreateEndpoint(endpoint) && strings.TrimSpace(result.ResponseID) != "" {
if err := h.gatewayService.BindGrokMediaVideoRequestAccount(
requestCtx, apiKey.GroupID, result.ResponseID, subject.UserID, apiKey.ID, account.ID,
+1
View File
@@ -31,6 +31,7 @@ type AdminHandlers struct {
UserAttribute *admin.UserAttributeHandler
ErrorPassthrough *admin.ErrorPassthroughHandler
TLSFingerprintProfile *admin.TLSFingerprintProfileHandler
Plugin *admin.PluginHandler
APIKey *admin.AdminAPIKeyHandler
ScheduledTest *admin.ScheduledTestHandler
Channel *admin.ChannelHandler
+85 -37
View File
@@ -15,41 +15,63 @@ import (
// 广场路由挂 OptionalJWT 中间件:匿名可访问(除非 require_auth 开启),带 token 则
// 识别用户。可见性规则(橱窗语义,与「可用渠道」的可绑定语义不同):
// - 匿名:仅非专属分组(订阅型照常展示);
// - 登录:非专属分组 + user_allowed_groups 授权的专属分组(不检查订阅有效性)。
// - 登录:非专属分组 + user_allowed_groups 授权的专属分组(不检查订阅有效性);
// 若该用户开启了公开分组限制,则公开分组同样需要落在授权集合内。
type ModelPlazaHandler struct {
channelService *service.ChannelService
plazaService *service.ModelPlazaService
apiKeyService *service.APIKeyService
settingService *service.SettingService
}
// NewModelPlazaHandler 创建模型广场 handler。
func NewModelPlazaHandler(
channelService *service.ChannelService,
plazaService *service.ModelPlazaService,
apiKeyService *service.APIKeyService,
settingService *service.SettingService,
) *ModelPlazaHandler {
return &ModelPlazaHandler{
channelService: channelService,
plazaService: plazaService,
apiKeyService: apiKeyService,
settingService: settingService,
}
}
// modelPlazaOfficialPricing LiteLLM 官方参考价(USD per token)。
// modelPlazaOfficialPricing 官方参考价(USD per token,与计费目录同源)。
type modelPlazaOfficialPricing struct {
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheWrite1hPrice *float64 `json:"cache_write_1h_price,omitempty"`
CacheReadPrice *float64 `json:"cache_read_price"`
// Intervals 官方长上下文阶梯,仅多档模型给出。
Intervals []userPricingIntervalDTO `json:"intervals,omitempty"`
}
// modelPlazaModel 广场模型条目:渠道定价(白名单形态)+ 官方参考价。
// modelPlazaTimePricingPeriod 分时倍率时段(配置时区当天 [start, end))。
type modelPlazaTimePricingPeriod struct {
StartTime string `json:"start_time"`
EndTime string `json:"end_time"`
Multiplier float64 `json:"multiplier"`
}
// modelPlazaTimePricing 计费会生效的分时倍率(仅倍率 ≠ 1 的时段)。
// WeekdaysOnly 为 true 时时段仅周一至周五生效,周末整天按标准价计费。
type modelPlazaTimePricing struct {
Timezone string `json:"timezone"`
WeekdaysOnly bool `json:"weekdays_only,omitempty"`
Periods []modelPlazaTimePricingPeriod `json:"periods"`
}
// modelPlazaModel 广场模型条目:实收口径展示定价(白名单形态)+ 官方参考价。
type modelPlazaModel struct {
Name string `json:"name"`
Platform string `json:"platform"`
Pricing *userSupportedModelPricing `json:"pricing"`
OfficialPricing *modelPlazaOfficialPricing `json:"official_pricing"`
// LongContextBasis 多档时的计价基准:"whole_request"(整单按档)| "marginal"(仅超出部分)。
LongContextBasis string `json:"long_context_basis,omitempty"`
// TimePricing 分时倍率时段,落在时段内的请求整单乘倍率;无分时省略。
TimePricing *modelPlazaTimePricing `json:"time_pricing,omitempty"`
}
// modelPlazaGroup 广场分组条目(白名单字段)。
@@ -68,9 +90,11 @@ type modelPlazaGroup struct {
IsExclusive bool `json:"is_exclusive"`
// 生图独立倍率:为 true 时图片计费模型的实付倍率取 ImageRateMultiplier,
// 不取分组/用户专属倍率。
ImageRateIndependent bool `json:"image_rate_independent"`
ImageRateMultiplier float64 `json:"image_rate_multiplier"`
Models []modelPlazaModel `json:"models"`
ImageRateIndependent bool `json:"image_rate_independent"`
ImageRateMultiplier float64 `json:"image_rate_multiplier"`
// 分组是否启用长上下文阶梯计费;关闭时模型实付列只展示最低档/基础价。
LongContextPricingEnabled bool `json:"long_context_pricing_enabled"`
Models []modelPlazaModel `json:"models"`
}
// modelPlazaResponse 广场页响应。
@@ -98,17 +122,18 @@ func (h *ModelPlazaHandler) Get(c *gin.Context) {
return
}
groups, err := h.channelService.ListPlazaGroups(c.Request.Context())
groups, err := h.plazaService.ListGroups(c.Request.Context())
if err != nil {
response.ErrorFrom(c, err)
return
}
// allowedExclusive == nil 表示匿名;登录用户恒为非 nil(可能为空集合)。
var allowedExclusive map[int64]struct{}
// allowedGroups == nil 表示匿名;登录用户恒为非 nil(可能为空集合)。
var allowedGroups map[int64]struct{}
var restrictPublicGroups bool
var userRates map[int64]float64
if authed {
allowedExclusive, err = h.apiKeyService.GetUserAllowedGroupIDSet(c.Request.Context(), subject.UserID)
allowedGroups, restrictPublicGroups, err = h.apiKeyService.GetUserGroupVisibility(c.Request.Context(), subject.UserID)
if err != nil {
// 可见性数据拿不到时不能静默降级成匿名视图(会错漏专属分组),直接报错。
response.ErrorFrom(c, err)
@@ -122,7 +147,7 @@ func (h *ModelPlazaHandler) Get(c *gin.Context) {
}
}
visible := filterPlazaVisibleGroups(groups, allowedExclusive)
visible := filterPlazaVisibleGroups(groups, allowedGroups, restrictPublicGroups)
out := make([]modelPlazaGroup, 0, len(visible))
for i := range visible {
@@ -135,18 +160,21 @@ func (h *ModelPlazaHandler) Get(c *gin.Context) {
}
// filterPlazaVisibleGroups 按登录态裁剪分组可见性。
// allowedExclusive == nil 表示匿名(仅非专属);非 nil 表示登录(非专属 + 授权专属)。
// allowedGroups == nil 表示匿名(仅非专属);非 nil 表示登录(非专属 + 授权专属)。
// restrictPublicGroups 为 true 时,公开分组也必须落在 allowedGroups 内,否则用户会
// 在广场看到自己实际绑定不了的分组。
func filterPlazaVisibleGroups(
groups []service.PlazaGroup,
allowedExclusive map[int64]struct{},
allowedGroups map[int64]struct{},
restrictPublicGroups bool,
) []service.PlazaGroup {
visible := make([]service.PlazaGroup, 0, len(groups))
for _, g := range groups {
if g.IsExclusive {
if allowedExclusive == nil {
if g.IsExclusive || (restrictPublicGroups && allowedGroups != nil) {
if allowedGroups == nil {
continue
}
if _, ok := allowedExclusive[g.ID]; !ok {
if _, ok := allowedGroups[g.ID]; !ok {
continue
}
}
@@ -161,27 +189,30 @@ func toModelPlazaGroupDTO(g *service.PlazaGroup, userRates map[int64]float64) mo
for i := range g.Models {
m := &g.Models[i]
models = append(models, modelPlazaModel{
Name: m.Name,
Platform: m.Platform,
Pricing: toUserPricing(m.Pricing),
OfficialPricing: toModelPlazaOfficialPricing(m.OfficialPricing),
Name: m.Name,
Platform: m.Platform,
Pricing: toUserPricing(m.Pricing),
OfficialPricing: toModelPlazaOfficialPricing(m.OfficialPricing),
LongContextBasis: string(m.LongContextBasis),
TimePricing: toModelPlazaTimePricing(m.TimePricing),
})
}
dto := modelPlazaGroup{
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
Models: models,
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
LongContextPricingEnabled: g.LongContextPricingEnabled,
Models: models,
}
if rate, ok := userRates[g.ID]; ok {
dto.UserRateMultiplier = &rate
@@ -189,6 +220,22 @@ func toModelPlazaGroupDTO(g *service.PlazaGroup, userRates map[int64]float64) mo
return dto
}
// toModelPlazaTimePricing 转换分时倍率;nil 透传(JSON 省略)。
func toModelPlazaTimePricing(p *service.TimePricingSchedule) *modelPlazaTimePricing {
if p == nil || len(p.Periods) == 0 {
return nil
}
periods := make([]modelPlazaTimePricingPeriod, 0, len(p.Periods))
for _, period := range p.Periods {
periods = append(periods, modelPlazaTimePricingPeriod{
StartTime: period.StartTime,
EndTime: period.EndTime,
Multiplier: period.Multiplier,
})
}
return &modelPlazaTimePricing{Timezone: p.Timezone, WeekdaysOnly: p.WeekdaysOnly, Periods: periods}
}
// toModelPlazaOfficialPricing 转换官方参考价;nil 透传(前端显示 "-")。
func toModelPlazaOfficialPricing(p *service.PlazaOfficialPricing) *modelPlazaOfficialPricing {
if p == nil {
@@ -200,5 +247,6 @@ func toModelPlazaOfficialPricing(p *service.PlazaOfficialPricing) *modelPlazaOff
CacheWritePrice: p.CacheWritePrice,
CacheWrite1hPrice: p.CacheWrite1hPrice,
CacheReadPrice: p.CacheReadPrice,
Intervals: toUserPricingIntervals(p.Intervals),
}
}
@@ -25,7 +25,7 @@ func plazaGroups() []service.PlazaGroup {
func TestFilterPlazaVisibleGroups_AnonymousSeesOnlyNonExclusive(t *testing.T) {
// 匿名(allowedExclusive == nil):仅非专属分组;订阅型公开分组照常可见(橱窗语义)。
visible := filterPlazaVisibleGroups(plazaGroups(), nil)
visible := filterPlazaVisibleGroups(plazaGroups(), nil, false)
require.Len(t, visible, 2)
ids := []int64{visible[0].ID, visible[1].ID}
require.ElementsMatch(t, []int64{1, 3}, ids)
@@ -34,7 +34,7 @@ func TestFilterPlazaVisibleGroups_AnonymousSeesOnlyNonExclusive(t *testing.T) {
func TestFilterPlazaVisibleGroups_AuthedSeesGrantedExclusive(t *testing.T) {
// 登录:非专属 + 授权的专属;未授权的专属仍不可见。
allowed := map[int64]struct{}{2: {}}
visible := filterPlazaVisibleGroups(plazaGroups(), allowed)
visible := filterPlazaVisibleGroups(plazaGroups(), allowed, false)
require.Len(t, visible, 3)
ids := make([]int64, 0, len(visible))
for _, g := range visible {
@@ -46,10 +46,33 @@ func TestFilterPlazaVisibleGroups_AuthedSeesGrantedExclusive(t *testing.T) {
func TestFilterPlazaVisibleGroups_AuthedEmptySetSeesNoExclusive(t *testing.T) {
// 登录但无任何专属授权(空集合,非 nil):与匿名同样只见非专属,
// 但语义区分要保持——空集合不能被当作 nil 匿名分支。
visible := filterPlazaVisibleGroups(plazaGroups(), map[int64]struct{}{})
visible := filterPlazaVisibleGroups(plazaGroups(), map[int64]struct{}{}, false)
require.Len(t, visible, 2)
}
func TestFilterPlazaVisibleGroups_RestrictedUserSeesOnlyGrantedPublic(t *testing.T) {
// 开启公开分组限制后,公开分组也必须落在授权集合内,否则用户会在广场
// 看到自己实际绑定不了的分组。
allowed := map[int64]struct{}{1: {}, 2: {}}
visible := filterPlazaVisibleGroups(plazaGroups(), allowed, true)
ids := make([]int64, 0, len(visible))
for _, g := range visible {
ids = append(ids, g.ID)
}
// 3 是未授权的公开分组,受限后不可见;4 是未授权的专属分组,一贯不可见。
require.ElementsMatch(t, []int64{1, 2}, ids)
}
func TestFilterPlazaVisibleGroups_RestrictionDoesNotAffectAnonymous(t *testing.T) {
// 匿名没有用户记录,限制标志无从谈起,可见性必须与未受限时一致。
visible := filterPlazaVisibleGroups(plazaGroups(), nil, true)
ids := make([]int64, 0, len(visible))
for _, g := range visible {
ids = append(ids, g.ID)
}
require.ElementsMatch(t, []int64{1, 3}, ids)
}
func TestModelPlazaHandler_NilSettingServiceFailsClosed404(t *testing.T) {
gin.SetMode(gin.TestMode)
h := &ModelPlazaHandler{} // settingService == nil → fail-closed
@@ -91,7 +114,7 @@ func TestToModelPlazaGroupDTO_UserRateAndFieldWhitelist(t *testing.T) {
"id", "name", "description", "platform", "subscription_type",
"rate_multiplier", "user_rate_multiplier", "is_exclusive", "models",
"peak_rate_enabled", "peak_start", "peak_end", "peak_rate_multiplier",
"image_rate_independent", "image_rate_multiplier",
"image_rate_independent", "image_rate_multiplier", "long_context_pricing_enabled",
} {
_, exists := decoded[key]
require.Truef(t, exists, "plaza group DTO must expose %q", key)
@@ -109,6 +132,12 @@ func TestToModelPlazaGroupDTO_UserRateAndFieldWhitelist(t *testing.T) {
require.Contains(t, official, "cache_read_price")
_, has1h := official["cache_write_1h_price"]
require.False(t, has1h, "1h 缓存写价为 nil 时应 omitempty")
_, hasOfficialIntervals := official["intervals"]
require.False(t, hasOfficialIntervals, "官方无阶梯时 intervals 应 omitempty")
_, hasBasis := model["long_context_basis"]
require.False(t, hasBasis, "单档模型不输出 long_context_basis")
_, hasTimePricing := model["time_pricing"]
require.False(t, hasTimePricing, "无分时时不输出 time_pricing")
// 无专属倍率:user_rate_multiplier 整个字段省略
dtoNoRate := toModelPlazaGroupDTO(&g, nil)
@@ -124,4 +153,94 @@ func TestToModelPlazaOfficialPricing_NilPassthrough(t *testing.T) {
require.Nil(t, toModelPlazaOfficialPricing(nil))
}
func TestToModelPlazaGroupDTO_LongContextTiersAndBasis(t *testing.T) {
maxTokens := 272000
g := service.PlazaGroup{
ID: 3, Name: "ladder", Platform: "openai", SubscriptionType: "standard", RateMultiplier: 1,
LongContextPricingEnabled: true,
Models: []service.PlazaModel{{
Name: "gpt-5.4",
Platform: "openai",
Pricing: &service.ChannelModelPricing{
BillingMode: service.BillingModeToken,
InputPrice: testPtr(2.5e-6),
Intervals: []service.PricingInterval{
{MinTokens: 0, MaxTokens: &maxTokens, TierLabel: "≤272K", InputPrice: testPtr(2.5e-6)},
{MinTokens: 272000, TierLabel: ">272K", InputPrice: testPtr(5e-6)},
},
},
OfficialPricing: &service.PlazaOfficialPricing{
InputPrice: testPtr(2.5e-6),
Intervals: []service.PricingInterval{
{MinTokens: 0, MaxTokens: &maxTokens, TierLabel: "≤272K", InputPrice: testPtr(2.5e-6)},
{MinTokens: 272000, TierLabel: ">272K", InputPrice: testPtr(5e-6)},
},
},
LongContextBasis: service.ContextPricingBasisWholeRequest,
}},
}
raw, err := json.Marshal(toModelPlazaGroupDTO(&g, nil))
require.NoError(t, err)
var decoded map[string]any
require.NoError(t, json.Unmarshal(raw, &decoded))
require.Equal(t, true, decoded["long_context_pricing_enabled"])
model := decoded["models"].([]any)[0].(map[string]any)
require.Equal(t, "whole_request", model["long_context_basis"])
pricing := model["pricing"].(map[string]any)
paidTiers := pricing["intervals"].([]any)
require.Len(t, paidTiers, 2)
require.Equal(t, ">272K", paidTiers[1].(map[string]any)["tier_label"])
official := model["official_pricing"].(map[string]any)
officialTiers := official["intervals"].([]any)
require.Len(t, officialTiers, 2)
first := officialTiers[0].(map[string]any)
require.Equal(t, "≤272K", first["tier_label"])
require.InDelta(t, 272000, first["max_tokens"].(float64), 0)
require.Contains(t, first, "cache_write_price", "区间 DTO 字段齐全(nil 输出 null)")
}
func testPtr(v float64) *float64 { return &v }
func TestToModelPlazaGroupDTO_TimePricing(t *testing.T) {
g := service.PlazaGroup{
ID: 4, Name: "cn", Platform: "deepseek", SubscriptionType: "standard", RateMultiplier: 1,
Models: []service.PlazaModel{{
Name: "deepseek-chat",
Platform: "deepseek",
Pricing: &service.ChannelModelPricing{BillingMode: service.BillingModeToken, InputPrice: testPtr(0.28e-6)},
TimePricing: &service.TimePricingSchedule{Timezone: "Asia/Shanghai", Periods: []service.TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}},
}, {
Name: "deepseek-reasoner",
Platform: "deepseek",
Pricing: &service.ChannelModelPricing{BillingMode: service.BillingModeToken, InputPrice: testPtr(0.56e-6)},
TimePricing: &service.TimePricingSchedule{Timezone: "Asia/Shanghai", WeekdaysOnly: true, Periods: []service.TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}},
}},
}
raw, err := json.Marshal(toModelPlazaGroupDTO(&g, nil))
require.NoError(t, err)
var decoded map[string]any
require.NoError(t, json.Unmarshal(raw, &decoded))
model := decoded["models"].([]any)[0].(map[string]any)
tp := model["time_pricing"].(map[string]any)
require.Equal(t, "Asia/Shanghai", tp["timezone"])
_, hasWeekdaysOnly := tp["weekdays_only"]
require.False(t, hasWeekdaysOnly, "未开启仅工作日时字段省略")
periods := tp["periods"].([]any)
require.Len(t, periods, 1)
first := periods[0].(map[string]any)
require.Equal(t, "00:30", first["start_time"])
require.Equal(t, "08:30", first["end_time"])
require.InDelta(t, 0.5, first["multiplier"].(float64), 1e-12)
weekdaysModel := decoded["models"].([]any)[1].(map[string]any)
weekdaysTP := weekdaysModel["time_pricing"].(map[string]any)
require.Equal(t, true, weekdaysTP["weekdays_only"])
}
+30 -1
View File
@@ -4,6 +4,8 @@ import (
"context"
"fmt"
"net/http"
"regexp"
"strconv"
"strings"
"github.com/gin-gonic/gin"
@@ -32,6 +34,29 @@ type noAccountErrorClassification struct {
ModelNotFound bool // true when this is a 404 model_not_found classification
}
var selectionModelRateLimitedPattern = regexp.MustCompile(`(?:model_rate_limited|rate_limited)=(\d+)`)
// classifySelectionFailureError preserves the scheduler's compact reason when
// every model-capable account is temporarily rate limited.
func classifySelectionFailureError(err error, fallback noAccountErrorClassification) noAccountErrorClassification {
if err == nil {
return fallback
}
match := selectionModelRateLimitedPattern.FindStringSubmatch(strings.ToLower(err.Error()))
if len(match) != 2 {
return fallback
}
count, parseErr := strconv.Atoi(match[1])
if parseErr != nil || count <= 0 {
return fallback
}
return noAccountErrorClassification{
Status: http.StatusTooManyRequests,
ErrType: "rate_limit_error",
Message: "All available accounts are currently rate-limited. Please retry later.",
}
}
// classifyNoAccountError decides between 404 model_not_found and 503
// api_error for "no available accounts" failures.
//
@@ -106,7 +131,11 @@ func classifyNoAccountErrorFromGin(
if c != nil && c.Request != nil {
ctx = c.Request.Context()
}
return classifyNoAccountError(ctx, diag, apiKey, routingModel, displayModel, platform)
classification := classifyNoAccountError(ctx, diag, apiKey, routingModel, displayModel, platform)
if classification.ModelNotFound {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalModelConfiguration)
}
return classification
}
func classifyOpenAICompatibleNoAccountErrorFromGin(
@@ -61,6 +61,21 @@ func TestClassifyNoAccountError_NilDiagnoser_Falls503(t *testing.T) {
require.False(t, cls.ModelNotFound)
}
func TestClassifySelectionFailureError_RateLimitedPool(t *testing.T) {
fallback := noAccountErrorClassification{Status: http.StatusServiceUnavailable, ErrType: "api_error", Message: "Service temporarily unavailable"}
got := classifySelectionFailureError(
fmt.Errorf("no available accounts supporting model: gpt-5.6-sol (total=3 eligible=0 model_rate_limited=3)"),
fallback,
)
require.Equal(t, http.StatusTooManyRequests, got.Status)
require.Equal(t, "rate_limit_error", got.ErrType)
require.Contains(t, got.Message, "rate-limited")
require.Equal(t, fallback, classifySelectionFailureError(fmt.Errorf("model_rate_limited=0"), fallback))
require.Equal(t, fallback, classifySelectionFailureError(fmt.Errorf("no available accounts"), fallback))
}
func TestClassifyNoAccountError_NilAPIKey_Falls503(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: false}}
@@ -113,6 +128,8 @@ func TestClassifyNoAccountError_ModelNotSupported_Returns404(t *testing.T) {
require.Equal(t, service.PlatformOpenAI, fd.calls[0].Platform)
require.NotNil(t, fd.calls[0].GroupID)
require.Equal(t, int64(42), *fd.calls[0].GroupID)
require.True(t, service.HasOpsClientBusinessLimited(c))
require.Equal(t, service.OpsClientBusinessLimitedReasonLocalModelConfiguration, service.OpsClientBusinessLimitedReason(c))
}
func TestClassifyOpenAICompatibleNoAccountError_GrokUsesGrokPlatform(t *testing.T) {
@@ -134,6 +151,8 @@ func TestClassifyOpenAICompatibleNoAccountError_GrokUsesGrokPlatform(t *testing.
require.True(t, cls.ModelNotFound)
require.Len(t, fd.calls, 1)
require.Equal(t, service.PlatformGrok, fd.calls[0].Platform)
require.True(t, service.HasOpsClientBusinessLimited(c))
require.Equal(t, service.OpsClientBusinessLimitedReasonLocalModelConfiguration, service.OpsClientBusinessLimitedReason(c))
logErr := openAICompatibleSelectionErrorForLog(
fmt.Errorf("no available OpenAI accounts supporting model: grok-4.5"),
@@ -142,6 +161,18 @@ func TestClassifyOpenAICompatibleNoAccountError_GrokUsesGrokPlatform(t *testing.
require.EqualError(t, logErr, "no available Grok accounts supporting model: grok-4.5")
}
func TestClassifyNoAccountError_PureClassifierDoesNotMarkGinContext(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: false}}
apiKey := &service.APIKey{GroupID: ptrInt64(7)}
cls := classifyNoAccountError(c.Request.Context(), fd, apiKey, "gpt-5", "gpt-5", service.PlatformOpenAI)
require.True(t, cls.ModelNotFound)
require.False(t, service.HasOpsClientBusinessLimited(c))
require.Empty(t, service.OpsClientBusinessLimitedReason(c))
}
func TestClassifyNoAccountError_HasModelSupport_KeepsRoutingMessageGenerationToCaller(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: true}}
@@ -200,4 +231,6 @@ func TestClassifyNoAccountError_FromGin_NilContextStillSafe(t *testing.T) {
require.Equal(t, http.StatusNotFound, cls.Status, "even with a nil gin context the classifier must still run and yield a coherent response")
require.True(t, cls.ModelNotFound)
require.False(t, service.HasOpsClientBusinessLimited(nil))
require.Empty(t, service.OpsClientBusinessLimitedReason(nil))
}
@@ -113,6 +113,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID)
profitVetoCount := 0
failedAccountIDs := make(map[int64]struct{})
sameAccountRetryCount := make(map[int64]int)
var lastFailoverErr *service.UpstreamFailoverError
switchCount := 0
var oauth429FailoverState service.OpenAIOAuth429FailoverState
@@ -186,7 +187,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds())
if err == nil {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), true, nil)
if result != nil {
h.recordAlphaSearchUsage(c, apiKey, account, subscription, channelMapping, requestedModel, body, result, subject.UserID)
}
@@ -195,7 +196,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
var failoverErr *service.UpstreamFailoverError
if !errors.As(err, &failoverErr) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), false, nil, err)
if c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
@@ -203,7 +204,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestedModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestedModel, false, result), false, nil, err)
if c.Writer.Size() != writerSizeBeforeForward {
h.handleFailoverExhausted(c, failoverErr, true)
return
@@ -215,6 +216,26 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
)
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai_alpha_search.same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_limit", retryLimit),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
zap.Duration("retry_delay", retryDelay),
)
select {
case <-c.Request.Context().Done():
return
case <-time.After(retryDelay):
}
continue
}
}
h.gatewayService.RecordOpenAIAccountSwitch()
failedAccountIDs[account.ID] = struct{}{}
lastFailoverErr = failoverErr
@@ -55,6 +55,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
logRequestBodyReadFailure(reqLog, c.Request, err)
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
@@ -88,6 +89,10 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
return
}
if _, err := service.ValidateOpenAIServiceTierField(body); err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if service.IsGPTImageGenerationModel(reqModel) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "This model is not supported on the Chat Completions endpoint")
return
@@ -182,6 +187,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
)
if len(failedAccountIDs) == 0 {
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
cls = classifySelectionFailureError(err, cls)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -239,11 +245,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}()
return h.gatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, promptCacheKey, "")
}()
cyberBlockKeyChat := ""
var cyberBlockBodyChat []byte
if service.GetOpsCyberPolicy(c) != nil {
cyberBlockKeyChat = service.CyberSessionBlockKey(apiKey.ID, c, body)
cyberBlockBodyChat = body
}
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyChat, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyChat, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
forwardDurationMs := time.Since(forwardStart).Milliseconds()
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
@@ -261,6 +267,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
if res == nil {
return
}
stampOpenAIRequestedReasoningEffort(res, c)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
inboundEndpoint := GetInboundEndpoint(c)
@@ -315,11 +322,12 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
return
}
if c.Writer.Size() != writerSizeBeforeForward {
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
h.handleFailoverExhausted(c, failoverErr, true)
return
}
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, nil), false, nil, err)
}
if !failoverErr.ShouldRetryNextAccount() {
h.handleFailoverExhausted(c, failoverErr, streamStarted)
@@ -327,8 +335,8 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
// Pool mode: retry on the same account
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai_chat_completions.pool_mode_same_account_retry",
@@ -366,7 +374,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, nil), false, nil, err)
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
wroteFallback := false
if !upstreamErrorAlreadyCommunicated {
@@ -386,9 +394,9 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
}
if result != nil {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, result.FirstTokenMs)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, result.FirstTokenMs)
} else {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, nil)
}
submitChatUsage(result)
@@ -15,9 +15,9 @@ import (
// Codex CLI and the Codex desktop app refresh their model picker from
// GET {base_url}/models?client_version=... (custom provider mode) or
// GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land
// here. ChatGPT manifests are proxied verbatim; custom API key manifests receive
// provider-compatibility normalization and use a short-lived, asynchronously
// revalidated cache to tolerate canceled client requests.
// here. Groups with explicit account model mappings are generated locally;
// otherwise ChatGPT manifests are proxied verbatim and custom API key manifests
// receive provider-compatibility normalization plus short-lived caching.
func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
if c.Request.Context().Err() != nil {
return
@@ -32,6 +32,24 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
return
}
ifNoneMatch := c.GetHeader("If-None-Match")
configuredManifest, configured, err := h.gatewayService.BuildGroupConfiguredCodexModelsManifest(
c.Request.Context(),
apiKey.Group,
ifNoneMatch,
)
if err != nil {
if c.Request.Context().Err() != nil {
return
}
h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest")
return
}
if configured {
writeCodexModelsManifestResponse(c, configuredManifest)
return
}
maxAccountSwitches := h.maxAccountSwitches
if maxAccountSwitches <= 0 {
maxAccountSwitches = 3
@@ -56,7 +74,9 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
// 让 ops 错误日志携带实际选中的上游账号,便于定位失效账号(#4544)。
setOpsSelectedAccount(c, account.ID, account.Platform)
manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match"))
// The client ETag represents the final group-specific body, so fetch the
// source manifest before applying local filtering and alias metadata.
manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), "")
if err != nil {
if c.Request.Context().Err() != nil {
return
@@ -70,18 +90,31 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err))
return
}
if err := h.gatewayService.CompleteAPIKeyCodexModelsManifestForClient(manifest, account); err != nil {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to complete Codex models manifest")
return
}
if err := h.gatewayService.MergeGroupConfiguredCodexModels(c.Request.Context(), apiKey.Group, manifest, ifNoneMatch); err != nil {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to build Codex models manifest")
return
}
if c.Request.Context().Err() != nil {
return
}
if manifest.ETag != "" {
c.Header("ETag", manifest.ETag)
}
if manifest.NotModified {
c.Status(http.StatusNotModified)
return
}
c.Data(http.StatusOK, "application/json", manifest.Body)
writeCodexModelsManifestResponse(c, manifest)
return
}
}
func writeCodexModelsManifestResponse(c *gin.Context, manifest *service.CodexModelsManifest) {
if manifest.ETag != "" {
c.Header("ETag", manifest.ETag)
}
if manifest.NotModified {
c.Status(http.StatusNotModified)
c.Writer.WriteHeaderNow()
return
}
c.Data(http.StatusOK, "application/json", manifest.Body)
}
@@ -2,6 +2,7 @@ package handler
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
@@ -16,6 +17,7 @@ import (
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type codexModelsFailoverAccountRepo struct {
@@ -43,6 +45,14 @@ func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Cont
return accounts, nil
}
func (r codexModelsFailoverAccountRepo) ListSchedulableByGroupID(_ context.Context, _ int64) ([]service.Account, error) {
return append([]service.Account(nil), r.accounts...), nil
}
func (r codexModelsFailoverAccountRepo) ListByGroup(_ context.Context, _ int64) ([]service.Account, error) {
return append([]service.Account(nil), r.accounts...), nil
}
type codexModelsFailoverHTTPUpstream struct {
service.HTTPUpstream
mu sync.Mutex
@@ -116,6 +126,236 @@ func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) {
}
}
func TestCodexModelsAppliesLocalFiltersBeforeClientETag(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(43)
repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{
{
ID: 1,
Name: "custom-openai",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://upstream.example/v1",
},
},
}}
upstream := &codexModelsFailoverHTTPUpstream{
firstBody: `{"object":"list","data":[{"id":"codex-auto-review"},{"id":"gpt-5.6"}]}`,
}
gatewayService := service.NewOpenAIGatewayService(
repo,
nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil,
upstream,
nil, nil, nil, nil, nil, nil, nil, nil,
)
handler := &OpenAIGatewayHandler{gatewayService: gatewayService}
group := &service.Group{
ID: groupID,
Platform: service.PlatformOpenAI,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"codex-auto-review", "gpt-5.6"},
},
}
first := performCodexModelsRequestForGroup(t, handler, group, "")
if first.Code != http.StatusOK {
t.Fatalf("first status: got %d, want %d; body=%s", first.Code, http.StatusOK, first.Body.String())
}
if body := first.Body.String(); !strings.Contains(body, "codex-auto-review") || !strings.Contains(body, "gpt-5.6") {
t.Fatalf("first body did not include the explicitly selected models: %s", body)
}
oldETag := first.Header().Get("ETag")
if oldETag == "" {
t.Fatal("first response did not include an ETag")
}
group.ModelsListConfig.Enabled = false
second := performCodexModelsRequestForGroup(t, handler, group, oldETag)
if second.Code != http.StatusOK {
t.Fatalf("second status: got %d, want %d; body=%s", second.Code, http.StatusOK, second.Body.String())
}
if body := second.Body.String(); strings.Contains(body, "codex-auto-review") || !strings.Contains(body, "gpt-5.6") {
t.Fatalf("second body was not the filtered manifest: %s", body)
}
if newETag := second.Header().Get("ETag"); newETag == "" || newETag == oldETag {
t.Fatalf("second ETag: got %q, want a new final-body ETag", newETag)
}
third := performCodexModelsRequestForGroup(t, handler, group, second.Header().Get("ETag"))
if third.Code != http.StatusNotModified {
t.Fatalf("third status: got %d, want %d; body=%s", third.Code, http.StatusNotModified, third.Body.String())
}
if third.Body.Len() != 0 {
t.Fatalf("third body: got %q, want empty", third.Body.String())
}
}
func TestCodexModelsAPIKeyCacheDoesNotLeakGroupFilters(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{
{
ID: 1,
Name: "shared-api-key",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-shared",
"base_url": "https://upstream.example/v1",
},
},
}}
upstream := &codexModelsFailoverHTTPUpstream{
firstBody: `{"object":"list","data":[{"id":"model-a"},{"id":"model-b"}]}`,
}
gatewayService := service.NewOpenAIGatewayService(
repo,
nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil,
upstream,
nil, nil, nil, nil, nil, nil, nil, nil,
)
handler := &OpenAIGatewayHandler{gatewayService: gatewayService}
groupA := &service.Group{
ID: 91,
Platform: service.PlatformOpenAI,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"model-a"},
},
}
groupB := &service.Group{
ID: 92,
Platform: service.PlatformOpenAI,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"model-b"},
},
}
firstA := performCodexModelsRequestForGroup(t, handler, groupA, "")
require.Equal(t, http.StatusOK, firstA.Code, firstA.Body.String())
require.Equal(t, []string{"model-a"}, codexHandlerManifestSlugs(t, firstA))
firstB := performCodexModelsRequestForGroup(t, handler, groupB, "")
require.Equal(t, http.StatusOK, firstB.Code, firstB.Body.String())
require.Equal(t, []string{"model-b"}, codexHandlerManifestSlugs(t, firstB))
etagA := firstA.Header().Get("ETag")
require.NotEmpty(t, etagA)
var wg sync.WaitGroup
results := make([]*httptest.ResponseRecorder, 8)
for i := range results {
wg.Add(1)
go func(index int) {
defer wg.Done()
if index%2 == 0 {
results[index] = performCodexModelsRequestForGroup(t, handler, groupA, etagA)
return
}
results[index] = performCodexModelsRequestForGroup(t, handler, groupB, "")
}(i)
}
wg.Wait()
sawGroupB := false
for _, recorder := range results {
require.NotNil(t, recorder)
switch recorder.Code {
case http.StatusNotModified:
require.Empty(t, recorder.Body.Bytes())
case http.StatusOK:
slugs := codexHandlerManifestSlugs(t, recorder)
if len(slugs) == 1 && slugs[0] == "model-b" {
sawGroupB = true
continue
}
require.Equal(t, []string{"model-a"}, slugs)
default:
t.Fatalf("unexpected status %d body=%s", recorder.Code, recorder.Body.String())
}
}
require.True(t, sawGroupB)
}
// Scenario: OpenAI 分组内混用 OAuth 和第三方 API Key 时,管理员模型配置优先。
func TestCodexModelsUsesConfiguredModelsBeforeUpstreamDiscovery(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(44)
repo := &codexModelsFailoverAccountRepo{accounts: []service.Account{
{
ID: 1,
Name: "ark-compatible",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Priority: 0,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-ark",
"base_url": "https://ark.example/v1",
"model_mapping": map[string]any{
"glm-5.3": "glm-5.3",
},
},
},
{
ID: 2,
Name: "chatgpt-oauth",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
Schedulable: true,
Priority: 1,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "oauth-test",
},
},
}}
upstream := &codexModelsFailoverHTTPUpstream{firstStatus: http.StatusNotFound}
gatewayService := service.NewOpenAIGatewayService(
repo,
nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil, nil, nil, nil, nil,
upstream,
nil, nil, nil, nil, nil, nil, nil, nil,
)
handler := &OpenAIGatewayHandler{gatewayService: gatewayService}
recorder := performCodexModelsRequestForGroup(t, handler, &service.Group{
ID: groupID,
Platform: service.PlatformOpenAI,
}, "")
if got := upstream.calls(); len(got) != 0 {
t.Fatalf("upstream account calls: got %v, want none", got)
}
if recorder.Code != http.StatusOK {
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
}
var envelope struct {
Models []map[string]any `json:"models"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String())
}
if len(envelope.Models) != 1 || envelope.Models[0]["slug"] != "glm-5.3" {
t.Fatalf("models: got %v, want only glm-5.3", envelope.Models)
}
if _, ok := envelope.Models[0]["supported_reasoning_levels"]; !ok {
t.Fatalf("configured model is missing the Codex descriptor contract: %v", envelope.Models[0])
}
}
func TestCompositeCodexModelsReusesExistingManifestSelection(t *testing.T) {
handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
@@ -148,8 +388,23 @@ func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) {
if recorder.Code != http.StatusOK {
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
}
if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want {
t.Fatalf("body: got %q, want %q", got, want)
requireCompleteCodexModelsHandlerResponse(t, recorder, "gpt-5.6-sol")
})
}
}
// Scenario: an API-key upstream without /models is excluded only for this discovery request.
func TestCodexModelsFailsOverWhenAPIKeyModelsEndpointIsUnavailable(t *testing.T) {
for _, status := range []int{http.StatusNotFound, http.StatusMethodNotAllowed} {
t.Run(http.StatusText(status), func(t *testing.T) {
handler, upstream, groupID := newCodexModelsFailoverTestHandler(status)
recorder := performCodexModelsRequest(t, handler, groupID)
if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
t.Fatalf("upstream account calls: got %v, want %v", got, want)
}
if recorder.Code != http.StatusOK {
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
}
})
}
@@ -183,9 +438,7 @@ func TestCodexModelsFailsOverFromInvalidManifestEnvelope(t *testing.T) {
if recorder.Code != http.StatusOK {
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
}
if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want {
t.Fatalf("body: got %q, want %q", got, want)
}
requireCompleteCodexModelsHandlerResponse(t, recorder, "gpt-5.6-sol")
}
func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) {
@@ -193,7 +446,6 @@ func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) {
http.StatusBadRequest,
http.StatusUnauthorized,
http.StatusForbidden,
http.StatusNotFound,
600,
}
for _, status := range statuses {
@@ -300,23 +552,79 @@ func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount
}
func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder {
return performCodexModelsRequestForPlatform(t, handler, groupID, service.PlatformOpenAI)
return performCodexModelsRequestForGroup(t, handler, &service.Group{ID: groupID, Platform: service.PlatformOpenAI}, "")
}
func performCodexModelsRequestForPlatform(t *testing.T, handler *OpenAIGatewayHandler, groupID int64, platform string) *httptest.ResponseRecorder {
return performCodexModelsRequestForGroup(t, handler, &service.Group{ID: groupID, Platform: platform}, "")
}
func performCodexModelsRequestForGroup(t *testing.T, handler *OpenAIGatewayHandler, group *service.Group, etag string) *httptest.ResponseRecorder {
t.Helper()
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil)
if etag != "" {
c.Request.Header.Set("If-None-Match", etag)
}
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
GroupID: &groupID,
Group: &service.Group{ID: groupID, Platform: platform},
GroupID: &group.ID,
Group: group,
})
handler.CodexModels(c)
return recorder
}
func codexHandlerManifestSlugs(t *testing.T, recorder *httptest.ResponseRecorder) []string {
t.Helper()
var envelope struct {
Models []struct {
Slug string `json:"slug"`
} `json:"models"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String())
}
slugs := make([]string, 0, len(envelope.Models))
for _, model := range envelope.Models {
slugs = append(slugs, model.Slug)
}
return slugs
}
func requireCompleteCodexModelsHandlerResponse(t *testing.T, recorder *httptest.ResponseRecorder, slug string) {
t.Helper()
var envelope struct {
Models []map[string]any `json:"models"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatalf("decode body: %v; body=%s", err, recorder.Body.String())
}
if len(envelope.Models) != 1 {
t.Fatalf("models count: got %d, want 1; body=%s", len(envelope.Models), recorder.Body.String())
}
model := envelope.Models[0]
if got := model["slug"]; got != slug {
t.Fatalf("slug: got %v, want %q", got, slug)
}
if levels, ok := model["supported_reasoning_levels"].([]any); !ok || len(levels) == 0 {
t.Fatalf("supported_reasoning_levels must be populated: %v", model["supported_reasoning_levels"])
}
if messages, ok := model["model_messages"].(map[string]any); !ok || messages["instructions_template"] == "" {
t.Fatalf("model_messages.instructions_template must be populated: %v", model["model_messages"])
}
if policy, ok := model["truncation_policy"].(map[string]any); !ok || len(policy) == 0 {
t.Fatalf("truncation_policy must be populated: %v", model["truncation_policy"])
}
modalities, ok := model["input_modalities"].([]any)
if !ok || len(modalities) != 1 || modalities[0] != "text" {
t.Fatalf("custom OpenAI-compatible endpoint modalities: got %v, want [text]", model["input_modalities"])
}
}
func equalInt64Slices(got, want []int64) bool {
if len(got) != len(want) {
return false
@@ -215,7 +215,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
h.handleFailoverExhausted(c, failoverErr, true)
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), false, nil, err)
if failoverClientGone(c) {
reqLog.Info("openai_embeddings.failover_aborted_client_disconnected",
zap.Int64("account_id", account.ID),
@@ -239,7 +239,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), false, nil, err)
if c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
@@ -250,7 +250,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, reqModel, false, result), true, nil)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
inboundEndpoint := GetInboundEndpoint(c)
@@ -1,11 +1,10 @@
package handler
// CN 分组 /v1/messages 调度闸门回归(修复:正常途径创建的 CN 分组曾恒 403):
// sanitizeGroupMessagesDispatchFields 对非 openai 平台强制 AllowMessagesDispatch
// sanitizeGroupMessagesDispatchFields 对非 openai/composite 平台强制 AllowMessagesDispatch
// =false,故 CN 分组必须与 grok 一样在闸门处豁免,否则原生 Anthropic 直通
//(Claude Code 主用例)永远不可达。composite 分组同理:sanitize 对 composite
// 恒置 false,解析到 grok/CN 目标时必须按目标平台豁免,解析到 openai 目标
// 仍受开关控制。
//(Claude Code 主用例)永远不可达。composite 分组解析到 grok/CN 目标时按
// 目标平台豁免,解析到 openai 目标仍受其可配置开关控制。
import (
"net/http/httptest"
@@ -35,23 +34,25 @@ func TestAllowOpenAICompatibleMessagesDispatch_CNProvidersExempt(t *testing.T) {
func TestAllowOpenAICompatibleMessagesDispatch_CompositeResolvedTargets(t *testing.T) {
gin.SetMode(gin.TestMode)
newCompositeCtx := func(model string) (*gin.Context, *service.APIKey) {
newCompositeCtx := func(model string, allow bool) (*gin.Context, *service.APIKey) {
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest("POST", "/v1/messages", nil)
apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite, AllowMessagesDispatch: false}}
apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite, AllowMessagesDispatch: allow}}
ensureCompositeTargetPlatform(c, apiKey, model)
return c, apiKey
}
// 解析到 grok/CN 目标:与对应独立分组同语义豁免。
for _, model := range []string{"grok-4.3", "kimi-k2-thinking", "glm-5.2", "deepseek-v3.2"} {
c, apiKey := newCompositeCtx(model)
c, apiKey := newCompositeCtx(model, false)
require.True(t, allowOpenAICompatibleMessagesDispatch(c, apiKey), "model=%s", model)
}
// 解析到 openai 目标:仍受开关控制(composite 被 sanitize 恒置 false ⇒ 拒绝)。
c, apiKey := newCompositeCtx("gpt-5.5")
// 解析到 openai 目标:受 composite 分组自身开关控制。
c, apiKey := newCompositeCtx("gpt-5.5", false)
require.False(t, allowOpenAICompatibleMessagesDispatch(c, apiKey))
c, apiKey = newCompositeCtx("gpt-5.5", true)
require.True(t, allowOpenAICompatibleMessagesDispatch(c, apiKey))
// 未解析出目标平台:保持拒绝,不放宽。
cNone, _ := gin.CreateTestContext(httptest.NewRecorder())
@@ -96,18 +96,12 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
requestStart := time.Now()
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
account, err := h.gatewayService.SelectAccountForTokenCount(
c.Request.Context(),
apiKey.GroupID,
"",
sessionHash,
routingModel,
nil,
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
false,
requestPlatform,
)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
@@ -120,7 +114,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if selection == nil || selection.Account == nil {
if account == nil {
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, routingModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
@@ -129,11 +123,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
return
}
account := selection.Account
setOpsSelectedAccount(c, account.ID, account.Platform)
if selection.Acquired && selection.ReleaseFunc != nil {
defer selection.ReleaseFunc()
}
if err := h.gatewayService.ForwardResponsesInputTokens(c.Request.Context(), c, account, forwardBody); err != nil {
reqLog.Error("openai_input_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
}
@@ -182,8 +172,7 @@ func (h *OpenAIGatewayHandler) GrokCountTokens(c *gin.Context) {
}
// CountTokens handles Anthropic-compatible POST /v1/messages/count_tokens for OpenAI groups.
// It validates billing and routes to an OpenAI token-count bridge without taking concurrency slots
// or recording usage.
// It validates billing and routes to an OpenAI token-count bridge without recording usage.
func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok {
@@ -279,18 +268,12 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
if preferredMappedModel != "" {
currentRoutingModel = preferredMappedModel
}
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
account, err := h.gatewayService.SelectAccountForTokenCount(
c.Request.Context(),
apiKey.GroupID,
"",
sessionHash,
currentRoutingModel,
nil,
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
false,
openAICompatibleRequestPlatform(c.Request.Context(), apiKey),
)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
@@ -304,7 +287,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
h.anthropicErrorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if selection == nil || selection.Account == nil {
if account == nil {
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
@@ -313,11 +296,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
return
}
account := selection.Account
setOpsSelectedAccount(c, account.ID, account.Platform)
if selection.Acquired && selection.ReleaseFunc != nil {
defer selection.ReleaseFunc()
}
forwardBody := mappedBodyForMessages(channelMapping.Mapped, channelMapping.MappedModel)
defaultMappedModel := preferredMappedModel
@@ -13,6 +13,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestGatewayChatCredentialStopDoesNotSelectAnotherAccountAndReturnsSafe503(t *testing.T) {
@@ -64,6 +65,113 @@ func TestGatewayChatAntigravityCredentialFailureReturnsActionableMessage(t *test
require.NotContains(t, strings.ToLower(recorder.Body.String()), "refresh_token")
}
func TestOpenAIAccessStateCredentialFailureUsesTypedSafeResponse(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
(&OpenAIGatewayHandler{}).handleFailoverExhausted(c, &service.UpstreamFailoverError{
StatusCode: http.StatusForbidden,
Stage: service.GatewayFailureStageAccountAuth,
Scope: service.GatewayFailureScopeAccount,
Reason: service.OpenAIUpstreamAccessStateReason,
NextAccountAction: service.NextAccountRetry,
ClientStatusCode: http.StatusBadGateway,
ClientMessage: "Upstream access is temporarily unavailable, please retry later",
ResponseBody: []byte(`{"error":{"message":"Your workspace is deactivated","token":"must-not-leak"}}`),
}, false)
require.Equal(t, http.StatusBadGateway, recorder.Code)
require.Contains(t, recorder.Body.String(), "Upstream access is temporarily unavailable")
require.NotContains(t, strings.ToLower(recorder.Body.String()), "deactivated")
require.NotContains(t, recorder.Body.String(), "must-not-leak")
}
func TestOpenAICapacityFailoverExhaustionPreservesMessageAsServerError(t *testing.T) {
gin.SetMode(gin.TestMode)
message := "Our servers are currently overloaded. Please try again later."
failoverErr := &service.UpstreamFailoverError{
StatusCode: http.StatusBadRequest,
ResponseBody: []byte(`{"error":{"code":"server_is_overloaded","message":"` + message + `"}}`),
RetryableOnSameAccount: true,
RequestScopedTransient: true,
ClientStatusCode: http.StatusServiceUnavailable,
ClientMessage: message,
}
t.Run("native_openai", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
(&OpenAIGatewayHandler{}).handleFailoverExhausted(c, failoverErr, false)
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
require.Equal(t, "server_error", gjson.Get(recorder.Body.String(), "error.type").String())
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
require.NotContains(t, recorder.Body.String(), "server_is_overloaded")
})
t.Run("responses_compat", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
(&GatewayHandler{}).handleResponsesFailoverExhausted(c, failoverErr, false)
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
require.Equal(t, "server_error", gjson.Get(recorder.Body.String(), "error.code").String())
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
})
t.Run("anthropic_compat", func(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
(&OpenAIGatewayHandler{}).handleAnthropicFailoverExhausted(c, failoverErr, false)
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
require.Equal(t, "api_error", gjson.Get(recorder.Body.String(), "error.type").String())
require.Equal(t, message, gjson.Get(recorder.Body.String(), "error.message").String())
})
}
func TestResponsesFailoverExhaustedAfterForwardedTerminalMarksOpsWithoutDuplicateFrame(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
official := "event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"code\":\"server_error\",\"message\":\"official failure\"}}}\n\n"
_, err := c.Writer.Write([]byte(official))
require.NoError(t, err)
service.MarkOpsStreamError(c, "server_error", "official failure", http.StatusBadGateway)
(&GatewayHandler{}).handleResponsesFailoverExhausted(c, &service.UpstreamFailoverError{
StatusCode: http.StatusBadGateway,
ResponseBody: []byte(`{"error":{"message":"fallback failure"}}`),
}, true)
require.Equal(t, official, recorder.Body.String())
streamErr, ok := service.GetOpsStreamError(c)
require.True(t, ok)
require.Equal(t, "official failure", streamErr.Message)
markerRecorder := httptest.NewRecorder()
markerContext, _ := gin.CreateTestContext(markerRecorder)
(&GatewayHandler{}).handleResponsesFailoverExhausted(markerContext, &service.UpstreamFailoverError{
StatusCode: http.StatusTooManyRequests,
}, true)
require.Contains(t, markerRecorder.Body.String(), "event: response.failed")
require.Equal(t, 1, strings.Count(markerRecorder.Body.String(), "event: response.failed"))
streamErr, ok = service.GetOpsStreamError(markerContext)
require.True(t, ok)
require.Equal(t, http.StatusTooManyRequests, streamErr.IntendedStatus)
require.Equal(t, "rate_limit_error", streamErr.ErrType)
heartbeatRecorder := httptest.NewRecorder()
heartbeatContext, _ := gin.CreateTestContext(heartbeatRecorder)
heartbeat := ": keepalive\n\n"
written, err := heartbeatRecorder.Write([]byte(heartbeat))
require.NoError(t, err)
recordGatewayStreamHeartbeat(heartbeatContext, written)
(&GatewayHandler{}).handleResponsesFailoverExhausted(heartbeatContext, &service.UpstreamFailoverError{
StatusCode: http.StatusBadGateway,
}, true)
require.True(t, strings.HasPrefix(heartbeatRecorder.Body.String(), heartbeat))
require.Equal(t, 1, strings.Count(heartbeatRecorder.Body.String(), "event: response.failed"))
}
func TestGatewayChatInferenceExhaustionRestoresRetryAfter(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
@@ -197,7 +305,7 @@ func TestOpsClassificationTreatsCredentialFailureAsAuthNotInference(t *testing.T
require.Equal(t, http.StatusForbidden, entry.UpstreamErrors[0].UpstreamStatusCode)
}
func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T) {
func TestOpsRecoveredCredentialFailoverDoesNotCreateRequestError(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
@@ -220,18 +328,75 @@ func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "account_auth", job.entry.ErrorPhase)
require.Equal(t, "provider", job.entry.ErrorOwner)
require.Equal(t, "gateway", job.entry.ErrorSource)
require.Contains(t, job.entry.ErrorMessage, "Recovered account authentication failure")
require.NotContains(t, job.entry.ErrorMessage, "403")
require.NotContains(t, job.entry.ErrorMessage, "earlier inference failure")
require.NotNil(t, job.entry.UpstreamStatusCode)
require.Zero(t, *job.entry.UpstreamStatusCode)
require.Nil(t, job.entry.UpstreamErrors)
require.Equal(t, http.StatusOK, job.entry.StatusCode)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 2)
require.Equal(t, http.StatusForbidden, events[0].UpstreamStatusCode)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), events[1].Stage)
}
func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.GET("/openai/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
Stage: string(service.GatewayFailureStageAccountAuth), Scope: string(service.GatewayFailureScopeAccount),
Reason: string(service.GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action",
}})
})
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil)
request.Header.Set("Connection", "Upgrade")
request.Header.Set("Upgrade", "websocket")
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusOK, job.entry.StatusCode)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 1)
require.Equal(t, string(service.GatewayFailureStageAccountAuth), events[0].Stage)
}
func TestOpsWebSocketCredentialFailoverExhaustedIsRecorded(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.GET("/openai/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
Stage: string(service.GatewayFailureStageAccountAuth), Scope: string(service.GatewayFailureScopeAccount),
Reason: string(service.GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action",
}})
closeOpenAIWSFailoverExhausted(c, nil, &service.UpstreamFailoverError{
Stage: service.GatewayFailureStageAccountAuth,
Scope: service.GatewayFailureScopeAccount,
Reason: service.GrokCredentialReasonRevoked,
NextAccountAction: service.NextAccountStop,
})
})
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil)
request.Header.Set("Connection", "Upgrade")
request.Header.Set("Upgrade", "websocket")
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "account_auth", job.entry.ErrorPhase)
require.Equal(t, http.StatusServiceUnavailable, job.entry.StatusCode)
require.Equal(t, service.GrokCredentialUnavailableClientMessage, job.entry.ErrorMessage)
}
@@ -24,7 +24,7 @@ func TestRecordCyberPolicyIfMarked_NoMark(t *testing.T) {
c := newTestGinContext()
h := &OpenAIGatewayHandler{}
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, nil, service.ChannelUsageFields{}, "")
// Flag must NOT be set when there was no mark.
require.False(t, c.GetBool(cyberPolicyRecordedKey),
@@ -47,14 +47,14 @@ func TestRecordCyberPolicyIfMarked_WithMark(t *testing.T) {
// First call: should set the flag.
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, nil, service.ChannelUsageFields{}, "")
})
require.True(t, c.GetBool(cyberPolicyRecordedKey),
"cyberPolicyRecordedKey must be true after first call with a mark")
// Second call: flag already set — must be a no-op (idempotent).
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
})
// Flag should still be true (not toggled or cleared).
require.True(t, c.GetBool(cyberPolicyRecordedKey),
@@ -75,7 +75,7 @@ func TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog(t *testing.T) {
h := &OpenAIGatewayHandler{}
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, nil, service.ChannelUsageFields{}, "")
})
require.True(t, c.GetBool(cyberPolicyRecordedKey))
}
@@ -88,7 +88,7 @@ func TestClearCyberPolicyTurnState(t *testing.T) {
h := &OpenAIGatewayHandler{}
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn1", UpstreamStatus: 200})
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
require.True(t, c.GetBool(cyberPolicyRecordedKey))
clearCyberPolicyTurnState(c)
@@ -97,7 +97,7 @@ func TestClearCyberPolicyTurnState(t *testing.T) {
// turn2: a fresh cyber hit must be recordable again.
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn2", UpstreamStatus: 200})
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, nil, service.ChannelUsageFields{}, "")
require.True(t, c.GetBool(cyberPolicyRecordedKey))
require.Equal(t, "turn2", service.GetOpsCyberPolicy(c).Message)
}
@@ -139,6 +139,23 @@ func TestRejectIfCyberSessionBlocked_FailOpen(t *testing.T) {
require.False(t, h2.rejectIfCyberSessionBlocked(c, key, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil gateway service → pass")
}
func TestBuildCyberSessionBlockWritePlanCombinesExplicitAndTranscriptKeys(t *testing.T) {
body := []byte(`{"messages":[{"role":"user","content":"setup"},{"role":"assistant","content":"ready"},{"role":"user","content":"trigger"}]}`)
c := newTestGinContext()
c.Request = httptest.NewRequest("POST", "/openai/v1/responses", strings.NewReader(string(body)))
c.Request.RemoteAddr = "203.0.113.44:12345"
c.Request.Header.Set("User-Agent", "client/1.2.3")
plan := buildCyberSessionBlockWritePlan(7, c, body)
require.Len(t, plan.keys, 2)
require.NotEmpty(t, plan.scopeKey)
c.Request.Header.Set("session_id", "sess-explicit")
plan = buildCyberSessionBlockWritePlan(7, c, body)
require.Len(t, plan.keys, 3)
require.NotEmpty(t, plan.scopeKey)
}
// TestRecordCyberPolicyIfMarked_BlockKeyPlumbed verifies the 6th param is
// accepted and a non-empty key with nil gateway service does not panic
// (write-side guards live in the service layer).
@@ -147,7 +164,7 @@ func TestRecordCyberPolicyIfMarked_BlockKeyPlumbed(t *testing.T) {
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "x", UpstreamStatus: 400})
h := &OpenAIGatewayHandler{}
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "deadbeef", service.ChannelUsageFields{}, "")
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, []byte(`{"input":"deadbeef"}`), service.ChannelUsageFields{}, "")
})
}
@@ -58,7 +58,56 @@ func newOpenAIWSUnsupportedModelSwitchError(model string) error {
}
func shouldReportOpenAIWSProxyAccountFailure(err error) bool {
return err != nil && !errors.Is(err, errOpenAIWSUnsupportedModelSwitch)
return err != nil && !errors.Is(err, errOpenAIWSUnsupportedModelSwitch) && !service.IsOpenAIWSSessionPreemptedError(err)
}
// openAIWSIngressEndedByClient reports whether a finished ingress WebSocket turn
// ended the way a healthy client ends one, rather than through an upstream or
// account fault.
//
// Three error shapes describe that same benign outcome and only the first was
// recognised:
//
// - *service.OpenAIWSClientCloseError carrying 1000 — the gateway closing the
// socket on its own terms, e.g. the inter-turn idle timeout.
// - a bare coderws.CloseError{Code: 1000} — what coder/websocket returns when
// the client closes cleanly. ReadOpenAIWSClientMessage hands conn.Read's
// error back verbatim, so nothing ever wraps it into the type above and an
// errors.As against that type cannot see it.
// - context.Canceled — the client went away mid-turn. That path closes with
// StatusGoingAway (1001) and carries the cancellation as its cause, so a
// check for 1000 alone never matched it either.
//
// The last two fell through to shouldReportOpenAIWSProxyAccountFailure, which
// filters only model-switch and session-preemption errors. Everything else
// reaches ObserveOpenAIAPIKeyHealthFailure and scheduler.ReportResult(false), so
// a client that merely disconnected counted against the upstream account's
// health and could trip it out of scheduling.
//
// failoverClientGone already states the rule this restores for the HTTP failover
// path — a cancelled client context "被误报成账号耗尽" is a bug, not a signal —
// and summarizeWSCloseErrorForLog already reads the close code the correct way,
// which is why the resulting WARN printed close_status=1000(StatusNormalClosure)
// for an error that was, in the same breath, being charged to the account.
//
// Deliberately narrow. StatusGoingAway is not matched on its own: the gateway
// emits 1001 when it tears a session down for its own reasons too, and the
// client-cancellation case is already covered by context.Canceled.
// context.DeadlineExceeded is left out as well — the idle-timeout path wraps it
// in a 1000 close error and stays benign through the first check, while any
// other deadline is a genuine stall worth reporting.
func openAIWSIngressEndedByClient(err error) bool {
if err == nil {
return true
}
var closeErr *service.OpenAIWSClientCloseError
if errors.As(err, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure {
return true
}
if coderws.CloseStatus(err) == coderws.StatusNormalClosure {
return true
}
return errors.Is(err, context.Canceled)
}
func openAIWSTurnBillingModel(result *service.OpenAIForwardResult, mapping service.ChannelMappingResult, requestedModel, upstreamModel string) string {
@@ -98,6 +147,22 @@ func openAIForwardSucceededForScheduling(result *service.OpenAIForwardResult) bo
return result.SucceededForScheduling()
}
func openAIAccountScheduleModel(c *gin.Context, account *service.Account, forwardModel string, requireCompact bool, result *service.OpenAIForwardResult) string {
if result != nil {
if actual := strings.TrimSpace(result.UpstreamModel); actual != "" {
return actual
}
}
if c != nil {
if value, ok := c.Get(service.OpsUpstreamModelKey); ok {
if actual, ok := value.(string); ok && strings.TrimSpace(actual) != "" {
return strings.TrimSpace(actual)
}
}
}
return service.ResolveOpenAIAccountUpstreamModelForRequest(account, forwardModel, requireCompact)
}
func resolveOpenAIMessagesDispatchMappedModel(c *gin.Context, apiKey *service.APIKey, requestedModel string) string {
if apiKey == nil || apiKey.Group == nil {
return ""
@@ -207,14 +272,13 @@ func allowOpenAICompatibleMessagesDispatch(c *gin.Context, apiKey *service.APIKe
}
// 国产供应商分组与 grok 同语义:/v1/messages 就是其主要服务形态(anthropic
// 协议账号原生直通 Claude Code),无需 allow_messages_dispatch 开关授权——
// 该开关对非 openai 平台恒被 sanitizeGroupMessagesDispatchFields 置 false,
// 该开关对非 openai/composite 平台恒被 sanitizeGroupMessagesDispatchFields 置 false,
// 若不豁免,CN 分组将永远 403。
if service.IsCNProvider(apiKey.Group.Platform) {
return true
}
// composite 分组解析到 grok/CN 目标时与对应独立分组同语义豁免:sanitize
// 对 composite 同样恒置 false,不豁免则这些目标的 /v1/messages 永远 403;
// 解析到 openai 目标仍受开关控制,维持现状。
// composite 分组解析到 grok/CN 目标时与对应独立分组同语义豁免;
// 解析到 openai 目标则受 composite 分组自身的可配置开关控制。
if apiKey.Group.Platform == service.PlatformComposite && c != nil && c.Request != nil {
if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok &&
(platform == service.PlatformGrok || service.IsCNProvider(platform)) {
@@ -320,6 +384,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
logRequestBodyReadFailure(reqLog, c.Request, err)
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
@@ -376,6 +441,10 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
return
}
if _, err := service.ValidateOpenAIServiceTierField(body); err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
previousResponseID := strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String())
if previousResponseID != "" {
@@ -392,12 +461,27 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id must be a response.id (resp_*), not a message id")
return
}
reqLog.Warn("openai.request_validation_failed",
zap.String("reason", "previous_response_id_requires_wsv2"),
groupID := int64(0)
if apiKey.GroupID != nil {
groupID = *apiKey.GroupID
}
owned, ownershipErr := h.gatewayService.ValidateOpenAIHTTPResponseOwner(
c.Request.Context(),
groupID,
previousResponseID,
subject.UserID,
apiKey.ID,
)
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id is only supported on Responses WebSocket v2")
return
if ownershipErr != nil {
reqLog.Warn("openai.previous_response_owner_lookup_failed", zap.Error(ownershipErr))
}
if !owned {
reqLog.Warn("openai.request_validation_failed", zap.String("reason", "previous_response_owner_mismatch"))
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id is not available for this user")
return
}
}
service.SetOpenAIHTTPResponseOwner(c, subject.UserID, apiKey.ID)
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
@@ -483,6 +567,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
if h.rejectIfCyberSessionBlocked(c, apiKey, sessionHashBody, reqModel, cyberBlockFormatResponses) {
return
}
c.Request = c.Request.WithContext(service.WithOpenAIGuardianParentAffinity(
c.Request.Context(), c, sessionHashBody, reqModel,
))
requireCompact := legacyCompact
maxAccountSwitches := h.maxAccountSwitches
@@ -549,6 +636,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
return
}
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, requestPlatform)
cls = classifySelectionFailureError(err, cls)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -583,6 +671,29 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
zap.Float64("load_skew", scheduleDecision.LoadSkew),
)
account := selection.Account
if previousResponseID != "" && requestPlatform == service.PlatformOpenAI && !account.IsOpenAIApiKey() {
// The public Responses HTTP API supports previous_response_id on API-key
// accounts. OAuth/SetupToken upstreams do not, so keep searching instead
// of silently deleting continuation state from a mixed account pool.
failedAccountIDs[account.ID] = struct{}{}
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
selection.ReleaseFunc = nil
}
lastFailoverErr = &service.UpstreamFailoverError{
StatusCode: http.StatusBadRequest,
Stage: service.GatewayFailureStageInference,
Scope: service.GatewayFailureScopeRequest,
Reason: service.OpenAIHTTPContinuationUnsupportedReason,
ClientStatusCode: http.StatusBadRequest,
ClientMessage: "previous_response_id requires an OpenAI API-key account for HTTP requests",
}
reqLog.Debug("openai.account_skipped_http_continuation_unsupported",
zap.Int64("account_id", account.ID),
zap.String("account_type", account.Type),
)
continue
}
sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account)
reqLog.Debug("openai.account_selected", zap.Int64("account_id", account.ID), zap.String("account_name", account.Name))
setOpsSelectedAccount(c, account.ID, account.Platform)
@@ -619,11 +730,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
}()
return h.gatewayService.Forward(c.Request.Context(), c, account, attemptBody)
}()
cyberBlockKeyHTTP := ""
var cyberBlockBodyHTTP []byte
if service.GetOpsCyberPolicy(c) != nil {
cyberBlockKeyHTTP = service.CyberSessionBlockKey(apiKey.ID, c, sessionHashBody)
cyberBlockBodyHTTP = sessionHashBody
}
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyHTTP, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyHTTP, clientRequestedUsageFields(c, channelMapping, reqModel, ""), service.HashUsageRequestPayload(body))
forwardDurationMs := time.Since(forwardStart).Milliseconds()
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
responseLatencyMs := forwardDurationMs
@@ -640,6 +751,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
if res == nil {
return
}
stampOpenAIRequestedReasoningEffort(res, c)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
@@ -696,6 +808,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
return
}
if !openAIForwardMayFailover(c, writerSizeBeforeForward, failoverErr) {
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
h.handleFailoverExhausted(c, failoverErr, true)
return
}
@@ -705,7 +818,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
streamStarted = true
}
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, nil), false, nil, err)
}
if !failoverErr.ShouldRetryNextAccount() {
h.handleFailoverExhausted(c, failoverErr, streamStarted)
@@ -717,8 +830,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
}
// 池模式:同账号重试
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai.pool_mode_same_account_retry",
@@ -767,7 +880,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
reqLog.Warn("openai.upstream_failover_switching", failoverSwitchFields...)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), false, nil, err)
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
wroteFallback := false
if !upstreamErrorAlreadyCommunicated {
@@ -793,9 +906,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
if account.Type == service.AccountTypeOAuth && !account.IsShadow() {
h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(c.Request.Context(), account.ID, result.ResponseHeaders)
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
} else {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), openAIForwardSucceededForScheduling(result), nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, forwardModel, requireCompact, result), openAIForwardSucceededForScheduling(result), nil)
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
@@ -839,6 +952,11 @@ func isOpenAIRemoteCompactionV2Request(body []byte) bool {
func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) {
isCompactRequest := isOpenAILegacyCompactPath(c)
if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) {
if normalized, changed, err := service.NormalizeCompactionTriggerInputOrder(body); err != nil {
reqLog.Warn("codex.remote_compact.trigger_order_normalization_failed", zap.Error(err))
} else if changed {
body = normalized
}
if isOpenAIRemoteCompactionV2Request(body) {
return body, true
}
@@ -1176,11 +1294,11 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}()
return h.gatewayService.ForwardAsAnthropic(c.Request.Context(), c, account, forwardBody, promptCacheKey, defaultMappedModel)
}()
cyberBlockKeyMsg := ""
var cyberBlockBodyMsg []byte
if service.GetOpsCyberPolicy(c) != nil {
cyberBlockKeyMsg = service.CyberSessionBlockKey(apiKey.ID, c, body)
cyberBlockBodyMsg = body
}
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyMsg, clientRequestedUsageFields(c, channelMappingMsg, reqModel, ""), service.HashUsageRequestPayload(body))
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockBodyMsg, clientRequestedUsageFields(c, channelMappingMsg, reqModel, ""), service.HashUsageRequestPayload(body))
forwardDurationMs := time.Since(forwardStart).Milliseconds()
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
responseLatencyMs := forwardDurationMs
@@ -1198,6 +1316,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
if res == nil {
return
}
stampOpenAIRequestedReasoningEffort(res, c)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
@@ -1254,11 +1373,12 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
return
}
if c.Writer.Size() != writerSizeBeforeForward {
h.gatewayService.ObserveOpenAIAccountHealthFailure(c.Request.Context(), account, err)
h.handleAnthropicFailoverExhausted(c, failoverErr, true)
return
}
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, nil), false, nil, err)
}
if !failoverErr.ShouldRetryNextAccount() {
h.handleAnthropicFailoverExhausted(c, failoverErr, streamStarted)
@@ -1266,8 +1386,8 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
// 池模式:同账号重试
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai_messages.pool_mode_same_account_retry",
@@ -1315,7 +1435,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
submitMessagesUsage(result)
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), false, nil, err)
wroteFallback := h.ensureAnthropicErrorResponse(c, streamStarted)
reqLog.Warn("openai_messages.forward_failed",
zap.Int64("account_id", account.ID),
@@ -1327,9 +1447,9 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
}
if result != nil {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), true, result.FirstTokenMs)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), true, result.FirstTokenMs)
} else {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(currentRoutingModel), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, currentRoutingModel, false, result), true, nil)
}
submitMessagesUsage(result)
@@ -1397,6 +1517,14 @@ func (h *OpenAIGatewayHandler) handleAnthropicFailoverExhausted(c *gin.Context,
h.anthropicStreamingAwareError(c, status, "api_error", message, streamStarted)
return
}
if failoverErr != nil && failoverErr.IsOpenAICapacityShed() && strings.TrimSpace(failoverErr.ClientMessage) != "" {
status := failoverErr.ClientStatusCode
if status <= 0 {
status = http.StatusServiceUnavailable
}
h.anthropicStreamingAwareError(c, status, "api_error", failoverErr.ClientMessage, streamStarted)
return
}
status, errType, errMsg := h.mapUpstreamError(failoverErr.StatusCode)
h.anthropicStreamingAwareError(c, status, errType, errMsg, streamStarted)
}
@@ -1539,9 +1667,32 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
streamStarted *bool,
reqLog *zap.Logger,
) (func(), openAISlotAcquireResult) {
return h.acquireOpenAIAccountSlot(c, groupID, sessionHash, selection, reqStream, streamStarted, reqLog, nil)
}
type openAISlotErrorWriter func(status int, errType, message string)
// acquireOpenAIAccountSlot centralizes scheduler selection admission. The
// optional error writer lets non-Responses endpoints retain their wire format
// while sharing the same WaitPlan, cancellation, and release semantics.
func (h *OpenAIGatewayHandler) acquireOpenAIAccountSlot(
c *gin.Context,
groupID *int64,
sessionHash string,
selection *service.AccountSelectionResult,
reqStream bool,
streamStarted *bool,
reqLog *zap.Logger,
writeError openAISlotErrorWriter,
) (func(), openAISlotAcquireResult) {
if writeError == nil {
writeError = func(status int, errType, message string) {
h.handleStreamingAwareError(c, status, errType, message, *streamStarted)
}
}
if selection == nil || selection.Account == nil {
markOpsRoutingCapacityLimited(c)
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted)
writeError(http.StatusServiceUnavailable, "api_error", "No available accounts")
return nil, openAISlotAcquireFailed
}
@@ -1571,7 +1722,7 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
}
if selection.WaitPlan == nil {
markOpsRoutingCapacityLimited(c)
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted)
writeError(http.StatusServiceUnavailable, "api_error", "No available accounts")
return nil, openAISlotAcquireFailed
}
@@ -1582,7 +1733,8 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
)
if err != nil {
reqLog.Warn("openai.account_slot_quick_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
h.handleConcurrencyError(c, err, "account", *streamStarted)
status, errType, message := concurrencyErrorResponse(err, "account")
writeError(status, errType, message)
return nil, openAISlotAcquireFailed
}
if fastAcquired {
@@ -1612,7 +1764,7 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
zap.Int64("account_id", account.ID),
zap.Int("max_waiting", selection.WaitPlan.MaxWaiting),
)
h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later", *streamStarted)
writeError(http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later")
return nil, openAISlotAcquireFailed
}
@@ -1635,7 +1787,8 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot(
)
if err != nil {
reqLog.Warn("openai.account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
h.handleConcurrencyError(c, err, "account", *streamStarted)
status, errType, message := concurrencyErrorResponse(err, "account")
writeError(status, errType, message)
return nil, openAISlotAcquireFailed
}
@@ -1814,19 +1967,36 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
return
}
// F5a: 握手层会话屏蔽检查。WS 握手无 body,显式标识仅来自握手 header
// (session_id / conversation_id);无标识则放行,连接内仍有本地 flag 兜底。
cyberBlockKey := service.CyberSessionBlockKey(apiKey.ID, c, nil)
if cyberBlockKey != "" && h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), cyberBlockKey) {
// The first response.create frame is available here, so explicit IDs are
// checked directly and body-derived sessions use the coarse scope gate.
if cyberBlockKey := findBlockedCyberSessionKey(c.Request.Context(), h.gatewayService, apiKey.ID, c, firstMessage); cyberBlockKey != "" {
writeCyberSessionBlockedWSError(c.Request.Context(), wsConn)
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "session blocked by cyber-security policy")
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, reqModel, cyberBlockKey)
return
}
cyberBlockedThisConn := false
var cyberTurnBodiesMu sync.Mutex
cyberTurnBodies := map[int][]byte{1: append([]byte(nil), firstMessage...)}
setCyberTurnBody := func(turn int, payload []byte) {
cyberTurnBodiesMu.Lock()
cyberTurnBodies[turn] = append([]byte(nil), payload...)
cyberTurnBodiesMu.Unlock()
}
takeCyberTurnBody := func(turn int) []byte {
cyberTurnBodiesMu.Lock()
body := cyberTurnBodies[turn]
delete(cyberTurnBodies, turn)
cyberTurnBodiesMu.Unlock()
return body
}
// 解析渠道级模型映射
channelMappingWS, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, reqModel)
wsForwardModel := reqModel
if channelMappingWS.Mapped && strings.TrimSpace(channelMappingWS.MappedModel) != "" {
wsForwardModel = strings.TrimSpace(channelMappingWS.MappedModel)
}
var currentUserRelease func()
var currentAccountRelease func()
@@ -1892,23 +2062,48 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
firstMessage,
openAIWSIngressFallbackSessionSeed(subject.UserID, apiKey.ID, apiKey.GroupID),
)
ctx = service.WithOpenAIGuardianParentAffinity(ctx, c, firstMessage, reqModel)
maxAccountSwitches := h.maxAccountSwitches
switchCount := 0
profitVetoCount := 0
failedAccountIDs := make(map[int64]struct{})
sameAccountRetryCount := make(map[int64]int)
var lastFailoverErr *service.UpstreamFailoverError
var oauth429FailoverState service.OpenAIOAuth429FailoverState
wsAttemptMessage := append([]byte(nil), firstMessage...)
waitForWSSameAccountRetry := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool {
if account == nil || failoverErr == nil || failoverErr.StatusCode != http.StatusTooManyRequests || failoverErr.SameAccountRetryDeadline.IsZero() {
return false
}
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if !sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
return false
}
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai.websocket.same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
zap.Duration("retry_delay", retryDelay),
)
select {
case <-ctx.Done():
return false
case <-time.After(retryDelay):
return true
}
}
handleWSFailover := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool {
if ctx.Err() != nil {
return false
}
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, failoverErr)
}
releaseAccountSlot()
if !failoverErr.ShouldRetryNextAccount() {
closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr)
return false
}
if ctx.Err() != nil {
@@ -1918,12 +2113,12 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
failedAccountIDs[account.ID] = struct{}{}
lastFailoverErr = failoverErr
if switchCount >= maxAccountSwitches {
closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr)
return false
}
switchCount++
if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) {
closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr)
return false
}
reqLog.Warn("openai.websocket_upstream_failover_switching",
@@ -1979,7 +2174,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if lastFailoverErr != nil {
closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, lastFailoverErr)
} else {
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
}
@@ -1987,7 +2182,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
if selection == nil || selection.Account == nil {
if lastFailoverErr != nil {
closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, lastFailoverErr)
} else {
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
}
@@ -2061,6 +2256,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
// 准入完成:门并入连接 ctx,turn 级复核与 failover 重选共用。
ctx = admissionCtx
// Account selection starts a fresh upstream attempt. Clear any model
// captured by the previous failover account before credential lookup.
setOpsSelectedAccount(c, account.ID, account.Platform)
currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
if err := h.gatewayService.BindStickySessionAfterProfitAdmission(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil {
reqLog.Warn("openai.websocket_bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err))
@@ -2125,6 +2323,15 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
TurnStarted: recordTurnStart,
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
c.Set(securityAuditWSTurnContextKey, turn)
service.BeginOpsStreamTurn(c, turn)
setCyberTurnBody(turn, payload)
// Passthrough ingress intentionally skips BeforeTurn, so enforce only
// the connection-level cyber session gate here as well. Native ingress
// visits this hook first and gets the same side-effect-free close error;
// its original BeforeTurn guard remains as defense in depth.
if cyberBlockedThisConn {
return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, cyberSessionBlockedClientMsg, nil)
}
if turn == 1 {
return nil
}
@@ -2149,6 +2356,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
if model == "" {
model = reqModel
}
setOpsRequestContext(c, model, true)
mapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, model)
mappedModelUnchanged := false
if previous := turnChannelMapping.Load(); previous != nil && previous.turn < turn {
@@ -2209,6 +2417,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
},
AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
turnStart := getTurnStart(turn)
cyberBlockBody := takeCyberTurnBody(turn)
// F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行;
// CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行,
// 届时 defer 已清除标记)。
@@ -2234,7 +2443,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
turnUpstreamModel = turnRequestedModel
}
turnUsageFields := turnMapping.ToUsageFields(turnRequestedModel, turnUpstreamModel)
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, turnRequestedModel, turnErr != nil, cyberBlockKey, turnUsageFields, requestPayloadHash)
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, turnRequestedModel, turnErr != nil, cyberBlockBody, turnUsageFields, requestPayloadHash)
if service.GetOpsCyberPolicy(c) != nil {
cyberBlockedThisConn = true
}
@@ -2271,7 +2480,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
if scheduleModel == "" {
scheduleModel = turnRequestedModel
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, scheduleModel, openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, scheduleModel, openAIForwardSucceededForScheduling(result), result.FirstTokenMs)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
@@ -2322,14 +2531,26 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
// WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
if preemptCtx, cleanupPreempt, armed := h.gatewayService.BeginOpenAIWSIngressSessionPreemption(ctx, c, account, wsFirstMessage); armed {
ctx = preemptCtx
defer cleanupPreempt()
}
if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
for {
err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks)
if err == nil {
reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
return
}
if service.IsOpenAIWSSessionPreemptedError(err) {
return
}
var failoverErr *service.UpstreamFailoverError
if errors.As(err, &failoverErr) {
retryPayload, retryCurrentTurn := service.OpenAIWSCurrentTurnRetryPayload(err)
nextAttemptMessage, retrySafe := openAIWSNextAttemptMessage(wsAttemptMessage, retryPayload, retryCurrentTurn)
if !retrySafe {
closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr)
return
}
wsAttemptMessage = nextAttemptMessage
@@ -2341,9 +2562,31 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
zap.Int("retry_payload_bytes", len(retryPayload)),
)
}
if handleWSFailover(account, failoverErr) {
if waitForWSSameAccountRetry(account, failoverErr) {
if failoverErr.ShouldReportAccountScheduleFailure() {
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, err)
}
if !ensureUserSlotHeld() {
return
}
if currentAccountRelease == nil {
accountRelease, acquired, acquireErr := h.concurrencyHelper.TryAcquireAccountSlot(ctx, account.ID, accountMaxConcurrency)
if acquireErr != nil || !acquired {
reqLog.Warn("openai.websocket_same_account_retry_slot_unavailable",
zap.Int64("account_id", account.ID),
zap.Error(acquireErr),
)
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
return
}
currentAccountRelease = wrapReleaseOnDone(ctx, accountRelease)
}
wsFirstMessage = wsAttemptMessage
continue
}
if handleWSFailover(account, failoverErr) {
break
}
return
}
@@ -2357,17 +2600,28 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
var closeErr *service.OpenAIWSClientCloseError
if errors.As(err, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure {
reqLog.Info("openai.websocket_ingress_closed_normally",
zap.Int64("account_id", account.ID),
zap.String("reason", closeErr.Reason()),
)
closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
hasClientCloseErr := errors.As(err, &closeErr)
if openAIWSIngressEndedByClient(err) {
closedFields := []zap.Field{zap.Int64("account_id", account.ID)}
if hasClientCloseErr {
closedFields = append(closedFields, zap.String("reason", closeErr.Reason()))
} else {
closedFields = append(closedFields, zap.Error(err))
}
reqLog.Info("openai.websocket_ingress_closed_normally", closedFields...)
// A bare coderws.CloseError or a plain cancellation carries no
// gateway-chosen close frame; mirror the client's clean 1000
// rather than the 1011 the proxy-failure tail would have sent.
if hasClientCloseErr {
closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
} else {
closeOpenAIClientWS(wsConn, coderws.StatusNormalClosure, "")
}
return
}
if shouldReportOpenAIWSProxyAccountFailure(err) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, wsForwardModel, false, nil), false, nil, err)
}
closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
proxyFailedFields := []zap.Field{
@@ -2387,15 +2641,13 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
proxyFailedFields = append(proxyFailedFields, zap.Int64p("proxy_id", account.ProxyID))
}
reqLog.Warn("openai.websocket_proxy_failed", proxyFailedFields...)
if errors.As(err, &closeErr) {
if hasClientCloseErr {
closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
return
}
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "upstream websocket proxy failed")
return
}
reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
return
}
}
@@ -2611,12 +2863,28 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE
)
return
}
if failoverErr.Reason == service.OpenAIHTTPContinuationUnsupportedReason {
message := strings.TrimSpace(failoverErr.ClientMessage)
if message == "" {
message = "previous_response_id requires an OpenAI API-key account for HTTP requests"
}
h.handleStreamingAwareError(c, http.StatusBadRequest, "invalid_request_error", message, streamStarted)
return
}
copyFailoverRetryAfter(c, failoverErr.ResponseHeaders)
if failoverErr.IsCredentialFailure() {
status, message := credentialFailoverClientResponse(failoverErr)
h.handleStreamingAwareError(c, status, "upstream_error", message, streamStarted)
return
}
if failoverErr.IsOpenAICapacityShed() && strings.TrimSpace(failoverErr.ClientMessage) != "" {
status := failoverErr.ClientStatusCode
if status <= 0 {
status = http.StatusServiceUnavailable
}
h.handleStreamingAwareError(c, status, "server_error", failoverErr.ClientMessage, streamStarted)
return
}
statusCode := failoverErr.StatusCode
responseBody := failoverErr.ResponseBody
if service.IsOpenAISilentRefusalErrorBody(responseBody) {
@@ -2659,6 +2927,13 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE
}
func credentialFailoverClientResponse(failoverErr *service.UpstreamFailoverError) (int, string) {
if failoverErr != nil && failoverErr.Reason == service.OpenAIUpstreamAccessStateReason && strings.TrimSpace(failoverErr.ClientMessage) != "" {
status := failoverErr.ClientStatusCode
if status <= 0 {
status = http.StatusServiceUnavailable
}
return status, failoverErr.ClientMessage
}
if failoverErr != nil && failoverErr.Reason == service.AntigravityCredentialRejectedReason {
return http.StatusBadGateway, service.AntigravityCredentialRejectedClientMessage
}
@@ -2993,25 +3268,44 @@ func openAIWSNextAttemptMessage(current, retryPayload []byte, retryCurrentTurn b
return append([]byte(nil), retryPayload...), true
}
func closeOpenAIWSFailoverExhausted(conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) {
if failoverErr == nil {
closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed")
return
}
if failoverErr.Stage == service.GatewayFailureStageAccountAuth {
closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, service.GrokCredentialUnavailableClientMessage)
return
}
switch failoverErr.StatusCode {
case http.StatusTooManyRequests:
closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream rate limit exceeded, please retry later")
case 529, http.StatusInternalServerError, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream service temporarily unavailable")
case http.StatusUnauthorized, http.StatusForbidden:
closeOpenAIClientWS(conn, coderws.StatusPolicyViolation, "upstream websocket authentication failed")
default:
closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed")
func closeOpenAIWSFailoverExhausted(c *gin.Context, conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) {
intendedStatus := http.StatusBadGateway
errorType := "upstream_error"
errorCode := "upstream_ws_failover_exhausted"
message := "upstream websocket proxy failed"
closeStatus := coderws.StatusInternalError
if failoverErr != nil {
if reason := strings.TrimSpace(string(failoverErr.Reason)); reason != "" {
errorCode = reason
}
if failoverErr.Stage == service.GatewayFailureStageAccountAuth {
intendedStatus = http.StatusServiceUnavailable
errorType = "api_error"
message = service.GrokCredentialUnavailableClientMessage
closeStatus = coderws.StatusTryAgainLater
} else {
switch failoverErr.StatusCode {
case http.StatusTooManyRequests:
intendedStatus = http.StatusTooManyRequests
errorType = "rate_limit_error"
message = "upstream rate limit exceeded, please retry later"
closeStatus = coderws.StatusTryAgainLater
case 529, http.StatusInternalServerError, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
intendedStatus = failoverErr.StatusCode
message = "upstream service temporarily unavailable"
closeStatus = coderws.StatusTryAgainLater
case http.StatusUnauthorized, http.StatusForbidden:
intendedStatus = failoverErr.StatusCode
errorType = "authentication_error"
message = "upstream websocket authentication failed"
closeStatus = coderws.StatusPolicyViolation
}
}
}
service.MarkOpsStreamFailure(c, errorType, errorCode, message, intendedStatus)
closeOpenAIClientWS(conn, closeStatus, message)
}
func writeContentModerationWSError(ctx context.Context, conn *coderws.Conn, decision *service.ContentModerationDecision) {
@@ -3206,13 +3500,10 @@ func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKe
if enabled, _ := h.gatewayService.CyberSessionBlockRuntime(c.Request.Context()); !enabled {
return false
}
key := service.CyberSessionBlockKey(apiKey.ID, c, body)
key := findBlockedCyberSessionKey(c.Request.Context(), h.gatewayService, apiKey.ID, c, body)
if key == "" {
return false
}
if !h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), key) {
return false
}
// body-signal compact 心跳可能已把响应头提交为 200(cyber 检查在用户槽位
// 长等待之后执行):以 response.failed 终止事件回传;未提交时停拍后照常
// 写 JSON(#3887)。
@@ -3240,12 +3531,56 @@ func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKe
return true
}
type cyberSessionBlockWritePlan struct {
scopeKey string
keys []string
}
func buildCyberSessionBlockWritePlan(apiKeyID int64, c *gin.Context, body []byte) cyberSessionBlockWritePlan {
plan := cyberSessionBlockWritePlan{}
if key := service.CyberSessionExplicitBlockKey(apiKeyID, c, body); key != "" {
plan.keys = append(plan.keys, key)
}
transcriptKeys := service.CyberSessionTranscriptBlockKeys(apiKeyID, body)
for _, key := range transcriptKeys {
if len(plan.keys) == 0 || key != plan.keys[0] {
plan.keys = append(plan.keys, key)
}
}
if len(transcriptKeys) > 0 {
plan.scopeKey = cyberSessionScopeKey(apiKeyID, c)
}
return plan
}
func findBlockedCyberSessionKey(ctx context.Context, gatewayService *service.OpenAIGatewayService, apiKeyID int64, c *gin.Context, body []byte) string {
if gatewayService == nil {
return ""
}
clientIP, userAgent := "", ""
if c != nil {
clientIP = strings.TrimSpace(ip.GetClientIP(c))
userAgent = c.GetHeader("User-Agent")
}
return gatewayService.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent)
}
func cyberSessionScopeKey(apiKeyID int64, c *gin.Context) string {
if c == nil {
return ""
}
return service.CyberSessionScopeKey(apiKeyID, strings.TrimSpace(ip.GetClientIP(c)), c.GetHeader("User-Agent"))
}
// enqueueCyberSessionBlockedOpsEntry captures request meta and enqueues the
// ops_error_logs entry for a locally blocked request.
func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context, apiKey *service.APIKey, model string, sessionBlockKey string) {
if h.opsService == nil {
return
}
// The dedicated cyber_session_blocked entry owns Ops semantics for this
// request; suppress the generic middleware record of the same 403 response.
c.Set(opsDedicatedErrorRecordedKey, true)
meta := cyberPolicyOpsErrorMeta{Model: model, InboundEndpoint: GetInboundEndpoint(c), CreatedAt: time.Now(), SessionBlockKey: sessionBlockKey}
meta.RequestID = c.Writer.Header().Get("X-Request-Id")
if c.Request != nil && c.Request.URL != nil {
@@ -3279,7 +3614,7 @@ func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context
// 并在 forward 返回错误时写一条 tokens=0 用量行。标记由 gateway 服务层在透传 cyber 后设置;
// 当前请求已发给用户,本方法只做事后记录,不影响响应。forwardErrored 为 true 时才写用量行,
// 避免与正常 RecordUsage(forward 成功路径)重复。每请求至多记录一次。
func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, model string, forwardErrored bool, cyberBlockKey string, channelFields service.ChannelUsageFields, requestPayloadHash string) {
func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, model string, forwardErrored bool, cyberBlockBody []byte, channelFields service.ChannelUsageFields, requestPayloadHash string) {
mark := service.GetOpsCyberPolicy(c)
if mark == nil {
return
@@ -3361,6 +3696,14 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
ClientIP: clientIPStr,
CreatedAt: time.Now(),
}
if gwSvc != nil && apiKey != nil {
plan := buildCyberSessionBlockWritePlan(apiKey.ID, c, cyberBlockBody)
if len(plan.keys) > 0 {
blockCtx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
gwSvc.MarkCyberSessionBlocked(blockCtx, plan.scopeKey, plan.keys)
cancel()
}
}
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
@@ -3402,9 +3745,6 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
ChannelUsageFields: channelFields,
})
}
if gwSvc != nil && cyberBlockKey != "" {
gwSvc.MarkCyberSessionBlocked(ctx, cyberBlockKey)
}
if opsSvc != nil {
enqueueOpsErrorLog(opsSvc, buildCyberPolicyOpsErrorEntry(opsMeta, mark))
}
@@ -672,7 +672,7 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
Platform: service.PlatformGrok,
},
}
require.Equal(t, "grok-4.5", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5"))
require.Equal(t, "grok-4.6", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5"))
require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "grok"))
})
@@ -855,7 +855,7 @@ func TestOpenAIResponses_RejectsMessageIDAsPreviousResponseID(t *testing.T) {
require.Contains(t, w.Body.String(), "previous_response_id must be a response.id")
}
func TestOpenAIResponses_RejectsHTTPContinuationPreviousResponseID(t *testing.T) {
func TestOpenAIResponses_AcceptsHTTPContinuationPreviousResponseIDBeforeRouting(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
@@ -877,11 +877,59 @@ func TestOpenAIResponses_RejectsHTTPContinuationPreviousResponseID(t *testing.T)
})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_123456", 1, 101))
h.Responses(c)
require.NotEqual(t, http.StatusBadRequest, w.Code)
require.NotContains(t, w.Body.String(), "Responses WebSocket v2")
}
func TestOpenAIResponses_RejectsHTTPContinuationOwnedByAnotherUser(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(
`{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_other_tenant","input":"hello"}`,
))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(2)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{
ID: 202,
UserID: 2,
GroupID: &groupID,
User: &service.User{ID: 2},
})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 2, Concurrency: 1})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_other_tenant", 1, 101))
h.Responses(c)
require.Equal(t, http.StatusBadRequest, w.Code)
require.Contains(t, w.Body.String(), "Responses WebSocket v2")
require.Contains(t, w.Body.String(), "previous_response_id")
require.Contains(t, w.Body.String(), "previous_response_id is not available for this user")
}
func TestOpenAIResponses_RejectsUnownedHTTPContinuation(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(
`{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_unknown","input":"hello"}`,
))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(2)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{ID: 101, UserID: 1, GroupID: &groupID, User: &service.User{ID: 1}})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 1, Concurrency: 1})
h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil)
h.Responses(c)
require.Equal(t, http.StatusBadRequest, w.Code)
require.Contains(t, w.Body.String(), "previous_response_id is not available for this user")
}
func TestOpenAIResponses_FunctionCallOutputHTTPGuidanceDoesNotSuggestPreviousResponseReuse(t *testing.T) {
@@ -1412,8 +1460,8 @@ func TestOpenAIResponsesWebSocket_PassthroughTracksModelPerTurn(t *testing.T) {
})
require.Len(t, got.upstreamPayloads, 2)
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
require.Equal(t, "gpt-5.6-terra", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
require.Equal(t, "sol-channel", gjson.GetBytes(got.upstreamPayloads[0], "model").String())
require.Equal(t, "terra-channel", gjson.GetBytes(got.upstreamPayloads[1], "model").String())
require.Len(t, got.clientEvents, 2)
require.Equal(t, "sol", gjson.GetBytes(got.clientEvents[0], "response.model").String())
require.Equal(t, "terra", gjson.GetBytes(got.clientEvents[1], "response.model").String())
@@ -1422,16 +1470,16 @@ func TestOpenAIResponsesWebSocket_PassthroughTracksModelPerTurn(t *testing.T) {
require.Equal(t, "sol", got.logs[0].Model)
require.Equal(t, "sol", got.logs[0].RequestedModel)
require.NotNil(t, got.logs[0].UpstreamModel)
require.Equal(t, "gpt-5.6-sol", *got.logs[0].UpstreamModel)
require.Equal(t, "sol-channel", *got.logs[0].UpstreamModel)
require.NotNil(t, got.logs[0].ModelMappingChain)
require.Equal(t, "sol→sol-channel→gpt-5.6-sol", *got.logs[0].ModelMappingChain)
require.Equal(t, "sol→sol-channel", *got.logs[0].ModelMappingChain)
require.Equal(t, "terra", got.logs[1].Model)
require.Equal(t, "terra", got.logs[1].RequestedModel)
require.NotNil(t, got.logs[1].UpstreamModel)
require.Equal(t, "gpt-5.6-terra", *got.logs[1].UpstreamModel)
require.Equal(t, "terra-channel", *got.logs[1].UpstreamModel)
require.NotNil(t, got.logs[1].ModelMappingChain)
require.Equal(t, "terra→terra-channel→gpt-5.6-terra", *got.logs[1].ModelMappingChain)
require.Equal(t, "terra→terra-channel", *got.logs[1].ModelMappingChain)
require.InDelta(t, got.logs[1].TotalCost*2.5, got.logs[0].TotalCost, 1e-12,
"each turn must be billed with its own channel-mapped model")
}
@@ -1593,6 +1641,29 @@ func TestOpenAIWSTurnBillingModelPreservesImagePricingModel(t *testing.T) {
}
}
func TestOpenAIAccountScheduleModelUsesActualOrSharedResolver(t *testing.T) {
account := &service.Account{
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"model_mapping": map[string]any{"public": "billing"},
"compact_model_mapping": map[string]any{"public": "compact-actual"},
},
}
reported := &service.OpenAIForwardResult{UpstreamModel: "observed-actual"}
require.Equal(t, "observed-actual", openAIAccountScheduleModel(nil, account, "public", true, reported))
require.Equal(t, "compact-actual", openAIAccountScheduleModel(nil, account, "public", true, nil))
require.Equal(t, "billing", openAIAccountScheduleModel(nil, account, "public", false, nil))
c, _ := gin.CreateTestContext(nil)
service.SetOpsUpstreamModel(c, "attempt-actual")
require.Equal(t, "attempt-actual", openAIAccountScheduleModel(c, account, "public", true, nil))
setOpsSelectedAccount(c, account.ID, account.Platform)
require.Equal(t, "attempt-actual", openAIAccountScheduleModel(c, account, "public", true, nil))
}
func TestShouldReportOpenAIWSProxyAccountFailure(t *testing.T) {
t.Run("unsupported client model switch does not penalize account", func(t *testing.T) {
err := fmt.Errorf("wrapped ingress turn: %w", newOpenAIWSUnsupportedModelSwitchError("gpt-unsupported"))
+14 -8
View File
@@ -271,7 +271,11 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
var imageUpstreamErr *service.OpenAIImagesUpstreamError
if errors.As(err, &imageUpstreamErr) {
retryableServerError := service.IsOpenAIImagesRetryableUpstreamError(imageUpstreamErr)
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), !retryableServerError, nil)
if retryableServerError {
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
} else {
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, nil)
}
logEvent := "openai.images.upstream_user_error"
if retryableServerError {
logEvent = "openai.images.upstream_server_error_after_flush"
@@ -287,7 +291,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
}
var failoverErr *service.UpstreamFailoverError
if errors.As(err, &failoverErr) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
if service.OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeForward {
reqLog.Warn("openai.images.upstream_failover_skipped_after_flush",
zap.Int64("account_id", account.ID),
@@ -304,19 +308,21 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai.images.pool_mode_same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_limit", retryLimit),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
zap.Duration("retry_delay", retryDelay),
)
select {
case <-requestCtx.Done():
return
case <-time.After(sameAccountRetryDelay):
case <-time.After(retryDelay):
}
continue
}
@@ -341,7 +347,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
)
continue
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), false, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), false, nil, err)
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
wroteFallback := false
if !upstreamErrorAlreadyCommunicated {
@@ -366,9 +372,9 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
if account.Type == service.AccountTypeOAuth && !account.IsShadow() {
h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(c.Request.Context(), account.ID, result.ResponseHeaders)
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), true, result.FirstTokenMs)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, result.FirstTokenMs)
} else {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(requestModel), true, nil)
h.gatewayService.ReportOpenAIAccountScheduleResult(account, openAIAccountScheduleModel(c, account, requestModel, false, result), true, nil)
}
userAgent := c.GetHeader("User-Agent")
@@ -0,0 +1,94 @@
package handler
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// 非法 service_tier 必须在两个 OpenAI 端点(/v1/responses、/v1/chat/completions)
// 上以 OpenAI 兼容错误结构返回 HTTP 400。这些用例在 handler 的 service_tier
// 校验处短路,不会进入账号选择/重试。
//
// 合法值(fast/priority/flex/auto/default/scale)与省略/null 的接受语义由
// service 层纯校验函数 TestValidateOpenAIServiceTierField 覆盖,避免 handler
// 测试走入真实账号选择/重试路径。
func newServiceTierHandlerTest(t *testing.T) *OpenAIGatewayHandler {
t.Helper()
return &OpenAIGatewayHandler{
gatewayService: &service.OpenAIGatewayService{},
billingCacheService: service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil),
apiKeyService: &service.APIKeyService{},
concurrencyHelper: &ConcurrencyHelper{concurrencyService: service.NewConcurrencyService(
&helperConcurrencyCacheStub{userSeq: []bool{true}},
)},
cfg: &config.Config{},
imageLimiter: &imageConcurrencyLimiter{},
}
}
func runOpenAIHandlerServiceTierTest(t *testing.T, path, body string, handler func(h *OpenAIGatewayHandler, c *gin.Context)) *httptest.ResponseRecorder {
t.Helper()
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(6401)
userID := int64(6402)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
ID: 6403,
GroupID: &groupID,
Group: &service.Group{
ID: groupID,
Platform: service.PlatformOpenAI,
},
User: &service.User{ID: userID, Status: service.StatusActive},
})
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: userID, Concurrency: 1})
handler(newServiceTierHandlerTest(t), c)
return rec
}
func TestOpenAIGatewayHandlerResponses_InvalidServiceTierRejected400(t *testing.T) {
for _, body := range []string{
`{"model":"gpt-5.5","input":"hi","service_tier":"turbo"}`,
`{"model":"gpt-5.5","input":"hi","service_tier":"SPEED"}`,
`{"model":"gpt-5.5","input":"hi","service_tier":""}`,
`{"model":"gpt-5.5","input":"hi","service_tier":123}`,
`{"model":"gpt-5.5","input":"hi","service_tier":{}}`,
} {
rec := runOpenAIHandlerServiceTierTest(t, "/v1/responses", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.Responses(c)
})
require.Equal(t, http.StatusBadRequest, rec.Code, "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid_request_error", "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid service_tier", "body=%s", body)
}
}
func TestOpenAIGatewayHandlerChatCompletions_InvalidServiceTierRejected400(t *testing.T) {
for _, body := range []string{
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":"turbo"}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":"ultra"}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":""}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":["priority"]}`,
} {
rec := runOpenAIHandlerServiceTierTest(t, "/v1/chat/completions", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.ChatCompletions(c)
})
require.Equal(t, http.StatusBadRequest, rec.Code, "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid_request_error", "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid service_tier", "body=%s", body)
}
}
@@ -0,0 +1,149 @@
package handler
import (
"context"
"errors"
"fmt"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
coderws "github.com/coder/websocket"
"github.com/stretchr/testify/require"
)
// issue #6105:入站 Responses WebSocket 的正常结束会被记成账号故障。
//
// 归因发生在 openai_gateway_handler.go 的 ingress 收尾处:只有
// *service.OpenAIWSClientCloseError 且状态码为 1000 被认作正常关闭,其余一律落到
// shouldReportOpenAIWSProxyAccountFailure —— 而它只排除 model-switch 与
// session-preempted 两种。于是客户端干净关闭(底层直接回裸 coderws.CloseError{1000})
// 与客户端中途断开(context.Canceled,收尾用 1001 关闭)都会喂给
// ObserveOpenAIAPIKeyHealthFailure 与 scheduler.ReportResult(success=false),
// 累积到阈值即把上游账号熔断出调度池。
//
// 这些用例钉住判定本身,与既有的 TestShouldReportOpenAIWSProxyAccountFailure 同一层级:
// 调用点位于一个需要真实上游 WS 才能进入的巨型 handler 循环内,仓库既有约定就是直接测判定函数。
// 缺陷主复现之一:客户端干净关闭。底层 conn.Read 的错误被 ReadOpenAIWSClientMessage
// 原样返回,没有任何地方把它包成 *OpenAIWSClientCloseError,所以旧断言看不见它。
func TestOpenAIWSIngressEndedByClient_BareNormalClosureIsNotAccountFailure(t *testing.T) {
err := coderws.CloseError{Code: coderws.StatusNormalClosure, Reason: "client done"}
// 前提:这正是旧判据漏掉它的原因——类型不匹配,不是状态码不匹配。
var closeErr *service.OpenAIWSClientCloseError
require.False(t, errors.As(err, &closeErr),
"裸 coderws.CloseError 不是 *OpenAIWSClientCloseError,旧的 errors.As 必然为假")
// 而按关闭码读,它确实是 1000。
require.Equal(t, coderws.StatusNormalClosure, coderws.CloseStatus(err))
require.True(t, openAIWSIngressEndedByClient(err))
}
// 同一形状被包一层(例如 ingress 把 read 错误裹进上下文)时也必须认得。
func TestOpenAIWSIngressEndedByClient_WrappedBareNormalClosureIsNotAccountFailure(t *testing.T) {
err := fmt.Errorf("ingress turn 3: %w",
coderws.CloseError{Code: coderws.StatusNormalClosure, Reason: "client done"})
require.True(t, openAIWSIngressEndedByClient(err))
}
// 缺陷主复现之二:客户端中途断开。ReadOpenAIWSClientMessage 在 controlCtx.Done()
// 分支用 StatusGoingAway 收尾并把 context.Canceled 作为 cause,所以「只认 1000」
// 这一条判据根本匹配不到它。
func TestOpenAIWSIngressEndedByClient_ClientCancelDuringTurnIsNotAccountFailure(t *testing.T) {
err := service.NewOpenAIWSClientCloseError(
coderws.StatusGoingAway, "websocket request canceled", context.Canceled)
// 前提:状态码是 1001 不是 1000,旧判据必然放行到账号归因。
var closeErr *service.OpenAIWSClientCloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusGoingAway, closeErr.StatusCode())
require.NotEqual(t, coderws.StatusNormalClosure, closeErr.StatusCode())
require.True(t, openAIWSIngressEndedByClient(err))
}
// 既有行为不得回退:网关自己用 1000 收尾(inter-turn idle timeout 就是这条)
// 原本就被认作正常关闭。
func TestOpenAIWSIngressEndedByClient_GatewayNormalClosureStillRecognised(t *testing.T) {
err := service.NewOpenAIWSClientCloseError(
coderws.StatusNormalClosure, "websocket idle timeout", context.DeadlineExceeded)
require.True(t, openAIWSIngressEndedByClient(err))
}
// 收窄证明:1001 本身不足以豁免。网关也会因自身原因用 GoingAway 收场,
// 客户端取消那一支已由 context.Canceled 覆盖,无需整类放行。
func TestOpenAIWSIngressEndedByClient_GoingAwayWithoutCancellationStillReported(t *testing.T) {
err := service.NewOpenAIWSClientCloseError(
coderws.StatusGoingAway, "upstream going away", errors.New("upstream closed session"))
require.False(t, openAIWSIngressEndedByClient(err))
require.True(t, shouldReportOpenAIWSProxyAccountFailure(err), "真实上游故障仍须归因账号")
}
// 契约没有丢:真正的故障仍然惩罚账号。判定组合与调用点一致——
// openAIWSIngressEndedByClient 为假才会走到 shouldReportOpenAIWSProxyAccountFailure。
func TestOpenAIWSIngressEndedByClient_AbnormalClosuresStillReportAccountFailure(t *testing.T) {
cases := []struct {
name string
err error
}{
{
name: "upstream_policy_violation",
err: service.NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation, "upstream websocket authentication failed",
errors.New("upstream rejected credentials")),
},
{
name: "upstream_internal_error",
err: service.NewOpenAIWSClientCloseError(
coderws.StatusInternalError, "upstream websocket proxy failed", nil),
},
{
name: "bare_abnormal_closure",
err: coderws.CloseError{Code: coderws.StatusAbnormalClosure, Reason: "connection reset"},
},
{
name: "generic_read_failure",
err: errors.New("upstream websocket read failed"),
},
{
// 空闲超时之外的 deadline 是真实停滞:不豁免。
name: "deadline_without_normal_close",
err: fmt.Errorf("upstream stalled: %w", context.DeadlineExceeded),
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
require.False(t, openAIWSIngressEndedByClient(tc.err))
require.True(t, shouldReportOpenAIWSProxyAccountFailure(tc.err))
})
}
}
// 不变式:同一条错误,日志侧与归因侧必须给出一致的结论。
// summarizeWSCloseErrorForLog 一直用 coderws.CloseStatus 读关闭码,这正是缺陷时期
// WARN 打印 close_status=1000(StatusNormalClosure) 却同时把账号记为故障的原因。
// 以后任何一侧改了读法,这条会红。
func TestOpenAIWSIngressEndedByClient_MatchesCloseCodeReportedInLog(t *testing.T) {
errs := []error{
coderws.CloseError{Code: coderws.StatusNormalClosure, Reason: "client done"},
fmt.Errorf("ingress turn 3: %w", coderws.CloseError{Code: coderws.StatusNormalClosure}),
service.NewOpenAIWSClientCloseError(coderws.StatusNormalClosure, "websocket idle timeout", context.DeadlineExceeded),
service.NewOpenAIWSClientCloseError(coderws.StatusGoingAway, "websocket request canceled", context.Canceled),
coderws.CloseError{Code: coderws.StatusAbnormalClosure, Reason: "connection reset"},
errors.New("upstream websocket read failed"),
}
for _, err := range errs {
t.Run(err.Error(), func(t *testing.T) {
closeStatus, _ := summarizeWSCloseErrorForLog(err)
if closeStatus == "1000(StatusNormalClosure)" {
require.True(t, openAIWSIngressEndedByClient(err),
"日志按 1000 归类为正常关闭,归因侧不得同时判为账号故障")
}
})
}
}
@@ -0,0 +1,319 @@
package handler
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/Wei-Shaw/sub2api/internal/testutil"
coderws "github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
type openAIWSPassthroughHandlerHarness struct {
clientConn *coderws.Conn
handlerDone <-chan struct{}
moderationRepo *contentModerationHandlerTestRepo
gatewayCache service.GatewayCache
apiKey *service.APIKey
}
func newOpenAIWSPassthroughHandlerHarness(t *testing.T, upstreamURL string) *openAIWSPassthroughHandlerHarness {
t.Helper()
gatewayCache := testutil.NewRedisGatewayCache(t)
settingRepo := &contentModerationHandlerSettingRepo{values: map[string]string{
service.SettingKeyRiskControlEnabled: "true",
service.SettingKeyCyberSessionBlockEnabled: "true",
service.SettingKeyCyberSessionBlockTTLSeconds: "60",
}}
moderationRepo := &contentModerationHandlerTestRepo{}
moderationSvc := service.NewContentModerationService(settingRepo, moderationRepo, nil, nil, nil, nil, nil, nil)
settingSvc := service.NewSettingService(settingRepo, nil)
groupID := int64(4301)
account := service.Account{
ID: 9951,
Name: "openai-ws-passthrough-cyber",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test", "base_url": upstreamURL},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
},
}
cfg := &config.Config{}
cfg.RunMode = config.RunModeSimple
cfg.Default.RateMultiplier = 1
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 3
accountRepo := &openAIWSUsageHandlerAccountRepoStub{account: account}
usageRepo := &openAIWSUsageHandlerUsageLogRepoStub{created: make(chan *service.UsageLog, 2)}
billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
gatewaySvc := service.NewOpenAIGatewayService(
accountRepo, usageRepo, nil, nil, nil, nil, gatewayCache, cfg, nil, nil,
service.NewBillingService(cfg, nil), nil, billingCacheSvc, nil, &service.DeferredService{},
nil, nil, nil, nil, nil, settingSvc, nil,
)
concurrencyCache := &concurrencyCacheMock{
acquireUserSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil },
acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil },
}
h := &OpenAIGatewayHandler{
gatewayService: gatewaySvc,
billingCacheService: billingCacheSvc,
apiKeyService: &service.APIKeyService{},
contentModerationService: moderationSvc,
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(concurrencyCache), SSEPingFormatNone, time.Second),
}
apiKey := &service.APIKey{
ID: 1851,
Name: "ws-cyber-key",
Key: "sk-handler-cyber-test",
GroupID: &groupID,
User: &service.User{ID: 1751, Status: service.StatusActive},
}
handlerDone := make(chan struct{})
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware.ContextKeyAPIKey), apiKey)
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1})
c.Next()
})
router.GET("/openai/v1/responses", func(c *gin.Context) {
h.ResponsesWebSocket(c)
close(handlerDone)
})
handlerServer := httptest.NewServer(router)
t.Cleanup(handlerServer.Close)
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(handlerServer.URL, "http")+"/openai/v1/responses", nil)
cancelDial()
require.NoError(t, err)
t.Cleanup(func() { _ = clientConn.CloseNow() })
return &openAIWSPassthroughHandlerHarness{
clientConn: clientConn,
handlerDone: handlerDone,
moderationRepo: moderationRepo,
gatewayCache: gatewayCache,
apiKey: apiKey,
}
}
func TestOpenAIResponsesWebSocketV2PassthroughCyberMarkIsConsumedAfterTurn(t *testing.T) {
gin.SetMode(gin.TestMode)
upstreamDone := make(chan struct{})
secondUpstreamFrame := make(chan []byte, 1)
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer close(upstreamDone)
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
require.NoError(t, err)
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, _, err = conn.Read(readCtx)
cancelRead()
require.NoError(t, err)
failed := []byte(`{"type":"response.failed","response":{"id":"resp_cyber_handler","model":"gpt-5.1","error":{"code":"cyber_policy","message":"blocked by upstream policy"},"usage":{"input_tokens":11,"output_tokens":3}}}`)
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
err = conn.Write(writeCtx, coderws.MessageText, failed)
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
_, second, err := conn.Read(readCtx)
cancelRead()
if err != nil {
return
}
secondUpstreamFrame <- append([]byte(nil), second...)
completed := []byte(`{"type":"response.completed","response":{"id":"resp_cyber_handler_turn_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`)
writeCtx, cancelWrite = context.WithTimeout(r.Context(), 3*time.Second)
err = conn.Write(writeCtx, coderws.MessageText, completed)
cancelWrite()
require.NoError(t, err)
}))
defer upstreamServer.Close()
harness := newOpenAIWSPassthroughHandlerHarness(t, upstreamServer.URL)
requestPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"cyber-session-1","input":"test"}`
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err := harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(requestPayload))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, event, err := harness.clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "response.failed", gjson.GetBytes(event, "type").String())
require.Eventually(t, func() bool {
logs := harness.moderationRepo.logSnapshot()
return len(logs) == 1 && logs[0].Action == service.ContentModerationActionCyberPolicy &&
strings.Contains(logs[0].Error, "upstream_usage=in:11,out:3")
}, 3*time.Second, 10*time.Millisecond, "handler AfterTurn must call recordCyberPolicyIfMarked and write the risk-control event")
keyCtx, _ := gin.CreateTestContext(httptest.NewRecorder())
keyCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(requestPayload))
blockKey := service.CyberSessionExplicitBlockKey(harness.apiKey.ID, keyCtx, []byte(requestPayload))
require.NotEmpty(t, blockKey)
store, ok := harness.gatewayCache.(service.CyberSessionBlockStore)
require.True(t, ok)
require.Eventually(t, func() bool {
matched, findErr := store.FindCyberSessionBlocked(context.Background(), []string{blockKey})
return findErr == nil && matched == blockKey
}, 3*time.Second, 10*time.Millisecond, "handler AfterTurn must write the cyber session block table")
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"cyber-session-1","input":"follow-up"}`))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = harness.clientConn.Read(readCtx)
cancelRead()
var closeErr coderws.CloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
// closeOpenAIClientWS caps close reasons at 120 bytes; passthrough must expose
// the same client-visible prefix rather than dropping the close frame.
require.Equal(t, "该会话已被网络安全策略屏蔽,请开启新会话 / This session is blocked by cyber-security policy, please ", closeErr.Reason)
select {
case <-harness.handlerDone:
case <-time.After(3 * time.Second):
t.Fatal("websocket handler did not exit")
}
select {
case <-upstreamDone:
case <-time.After(3 * time.Second):
t.Fatal("upstream websocket did not exit")
}
select {
case second := <-secondUpstreamFrame:
t.Fatalf("blocked follow-up reached upstream: %s", second)
default:
}
}
func TestOpenAIResponsesWebSocketV2PassthroughNonCyberTurnAllowsFollowup(t *testing.T) {
gin.SetMode(gin.TestMode)
upstreamDone := make(chan struct{})
secondUpstreamFrame := make(chan []byte, 1)
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer close(upstreamDone)
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
require.NoError(t, err)
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, _, err = conn.Read(readCtx)
cancelRead()
require.NoError(t, err)
firstCompleted := []byte(`{"type":"response.completed","response":{"id":"resp_non_cyber_handler_turn_1","model":"gpt-5.1","usage":{"input_tokens":2,"output_tokens":1}}}`)
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
err = conn.Write(writeCtx, coderws.MessageText, firstCompleted)
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
_, second, err := conn.Read(readCtx)
cancelRead()
require.NoError(t, err)
secondUpstreamFrame <- append([]byte(nil), second...)
secondCompleted := []byte(`{"type":"response.completed","response":{"id":"resp_non_cyber_handler_turn_2","model":"gpt-5.1","usage":{"input_tokens":3,"output_tokens":1}}}`)
writeCtx, cancelWrite = context.WithTimeout(r.Context(), 3*time.Second)
err = conn.Write(writeCtx, coderws.MessageText, secondCompleted)
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
_, _, _ = conn.Read(readCtx)
cancelRead()
}))
defer upstreamServer.Close()
harness := newOpenAIWSPassthroughHandlerHarness(t, upstreamServer.URL)
firstPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"non-cyber-session-1","input":"first"}`
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err := harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(firstPayload))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, firstEvent, err := harness.clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "resp_non_cyber_handler_turn_1", gjson.GetBytes(firstEvent, "response.id").String())
secondPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"non-cyber-session-1","input":"follow-up"}`
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(secondPayload))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
_, secondEvent, err := harness.clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "resp_non_cyber_handler_turn_2", gjson.GetBytes(secondEvent, "response.id").String())
require.Empty(t, harness.moderationRepo.logSnapshot())
keyCtx, _ := gin.CreateTestContext(httptest.NewRecorder())
keyCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(firstPayload))
blockKey := service.CyberSessionExplicitBlockKey(harness.apiKey.ID, keyCtx, []byte(firstPayload))
require.NotEmpty(t, blockKey)
store, ok := harness.gatewayCache.(service.CyberSessionBlockStore)
require.True(t, ok)
matched, findErr := store.FindCyberSessionBlocked(context.Background(), []string{blockKey})
require.NoError(t, findErr)
require.Empty(t, matched)
require.NoError(t, harness.clientConn.Close(coderws.StatusNormalClosure, "done"))
select {
case <-harness.handlerDone:
case <-time.After(3 * time.Second):
t.Fatal("non-cyber websocket handler did not exit")
}
select {
case <-upstreamDone:
case <-time.After(3 * time.Second):
t.Fatal("non-cyber upstream websocket did not exit")
}
select {
case second := <-secondUpstreamFrame:
require.JSONEq(t, secondPayload, string(second))
default:
t.Fatal("non-cyber follow-up did not reach upstream")
}
}
@@ -12,9 +12,40 @@ import (
"github.com/stretchr/testify/require"
)
type blockingOpsResponseWriter struct {
gin.ResponseWriter
writeStarted chan struct{}
writeRelease chan struct{}
}
func (w *blockingOpsResponseWriter) WriteString(s string) (int, error) {
close(w.writeStarted)
<-w.writeRelease
return w.ResponseWriter.WriteString(s)
}
type deterministicOpsCaptureWriterStatePool struct {
states []*opsCaptureWriterState
}
func (p *deterministicOpsCaptureWriterStatePool) Get() any {
if len(p.states) == 0 {
return &opsCaptureWriterState{limit: opsCaptureWriterLimit}
}
last := len(p.states) - 1
state := p.states[last]
p.states = p.states[:last]
return state
}
func (p *deterministicOpsCaptureWriterStatePool) Put(value any) {
if state, ok := value.(*opsCaptureWriterState); ok && state != nil {
p.states = append(p.states, state)
}
}
func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
w := &opsCaptureWriter{}
w.ResponseWriter = nil
assert.NotPanics(t, func() {
assert.Equal(t, 0, w.Status())
@@ -88,3 +119,86 @@ func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) {
require.Equal(t, http.StatusOK, outerStatus)
require.Equal(t, http.StatusOK, recorder.Code)
}
func TestOpsCaptureWriter_StaleLeaseCannotReachReacquiredState(t *testing.T) {
gin.SetMode(gin.TestMode)
pool := &deterministicOpsCaptureWriterStatePool{}
firstRecorder := httptest.NewRecorder()
firstContext, _ := gin.CreateTestContext(firstRecorder)
stale := acquireOpsCaptureWriterFromPool(pool, firstContext.Writer)
releaseOpsCaptureWriter(stale)
secondRecorder := httptest.NewRecorder()
secondContext, _ := gin.CreateTestContext(secondRecorder)
current := acquireOpsCaptureWriterFromPool(pool, secondContext.Writer)
defer releaseOpsCaptureWriter(current)
require.NotSame(t, stale, current)
require.Same(t, stale.state, current.state)
current.WriteHeader(http.StatusInternalServerError)
_, err := current.WriteString("current")
require.NoError(t, err)
require.Equal(t, []byte("current"), current.capturedBytes())
n, err := stale.WriteString("stale")
require.NoError(t, err)
require.Zero(t, n)
require.Nil(t, stale.capturedBytes())
require.Equal(t, []byte("current"), current.capturedBytes())
require.NotContains(t, secondRecorder.Body.String(), "stale")
// Releasing the stale handle must not return an active state to the pool.
releaseOpsCaptureWriter(stale)
thirdRecorder := httptest.NewRecorder()
thirdContext, _ := gin.CreateTestContext(thirdRecorder)
other := acquireOpsCaptureWriterFromPool(pool, thirdContext.Writer)
defer releaseOpsCaptureWriter(other)
require.NotSame(t, current.state, other.state)
}
func TestOpsCaptureWriter_ReleaseWaitsForDelegatedWriteWithoutHoldingStateMutex(t *testing.T) {
gin.SetMode(gin.TestMode)
pool := &deterministicOpsCaptureWriterStatePool{}
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
inner := &blockingOpsResponseWriter{
ResponseWriter: ctx.Writer,
writeStarted: make(chan struct{}),
writeRelease: make(chan struct{}),
}
w := acquireOpsCaptureWriterFromPool(pool, inner)
writeDone := make(chan struct{})
go func() {
defer close(writeDone)
_, _ = w.WriteString("body")
}()
<-inner.writeStarted
if !w.state.mu.TryLock() {
t.Fatal("state mutex remained held across the delegated network write")
}
w.state.mu.Unlock()
releaseDone := make(chan struct{})
go func() {
releaseOpsCaptureWriter(w)
close(releaseDone)
}()
select {
case <-releaseDone:
t.Fatal("release returned while a delegated write was still active")
case <-time.After(20 * time.Millisecond):
}
require.Empty(t, pool.states)
close(inner.writeRelease)
<-writeDone
select {
case <-releaseDone:
case <-time.After(time.Second):
t.Fatal("release did not finish after the delegated write returned")
}
require.Len(t, pool.states, 1)
}
File diff suppressed because it is too large Load Diff
+719 -12
View File
@@ -9,6 +9,7 @@ import (
"testing"
"unicode/utf8"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
@@ -37,15 +38,18 @@ func (r *ingressRejectSettingRepo) Set(context.Context, string, string) error {
type ingressRejectOpsRepo struct {
service.OpsRepository
insertCalls int
entries []*service.OpsInsertErrorLogInput
}
func (r *ingressRejectOpsRepo) InsertErrorLog(context.Context, *service.OpsInsertErrorLogInput) (int64, error) {
func (r *ingressRejectOpsRepo) InsertErrorLog(_ context.Context, entry *service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
r.entries = append(r.entries, entry)
return 0, nil
}
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(context.Context, []*service.OpsInsertErrorLogInput) (int64, error) {
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(_ context.Context, entries []*service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
r.entries = append(r.entries, entries...)
return 0, nil
}
@@ -187,21 +191,23 @@ func TestOpsCaptureWriterPool_ResetOnRelease(t *testing.T) {
writer := acquireOpsCaptureWriter(c.Writer)
require.NotNil(t, writer)
_, err := writer.buf.WriteString("temp-error-body")
c.Writer.WriteHeader(http.StatusInternalServerError)
_, err := writer.WriteString("temp-error-body")
require.NoError(t, err)
require.NotEmpty(t, writer.capturedBytes())
releaseOpsCaptureWriter(writer)
reused := acquireOpsCaptureWriter(c.Writer)
defer releaseOpsCaptureWriter(reused)
require.Zero(t, reused.buf.Len(), "writer should be reset before reuse")
require.Empty(t, reused.capturedBytes(), "writer should be reset before reuse")
}
func TestOpsCaptureWriterPool_DropsLargeBuffers(t *testing.T) {
w := &opsCaptureWriter{}
w.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
require.False(t, shouldPoolOpsCaptureWriter(w))
state := &opsCaptureWriterState{}
state.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
require.False(t, shouldPoolOpsCaptureWriterState(state))
}
func TestEnqueueOpsErrorLog_SanitizesAndBoundsBodyBeforeQueue(t *testing.T) {
@@ -279,6 +285,308 @@ func TestOpsErrorLoggerMiddleware_HardSkipsIngressRejection(t *testing.T) {
require.Zero(t, OpsErrorLogEnqueuedTotal(), "ingress rejection must not enter the error queue")
}
func TestOpsErrorLoggerMiddleware_DedicatedCyberSessionBlockRecordsExactlyOnce(t *testing.T) {
setupOpsErrorLogTestQueue(t, 3)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
h := &OpenAIGatewayHandler{opsService: ops}
apiKey := &service.APIKey{ID: 41, Key: "sk-dedicated-test"}
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, "gpt-test", "session-block-hash")
c.JSON(http.StatusForbidden, gin.H{"error": gin.H{
"type": "permission_error", "code": "session_blocked_by_cyber_policy", "message": "blocked",
}})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusForbidden, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "cyber_policy_session_blocked", job.entry.ErrorType)
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
}
func TestOpsErrorLoggerMiddleware_OrdinaryPermissionStillRecords(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.JSON(http.StatusForbidden, gin.H{"error": gin.H{
"type": "permission_error", "code": "permission_denied", "message": "denied",
}})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "permission_error", job.entry.ErrorType)
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
}
func TestOpsErrorLoggerMiddleware_RecordsRecoveredUpstreamTelemetryOutsideFailureSLA(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
repo := &ingressRejectOpsRepo{}
ops := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
UpstreamStatusCode: http.StatusTooManyRequests,
Message: "earlier attempt was rate limited",
}})
c.JSON(http.StatusOK, gin.H{"status": "completed"})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Nil(t, job.entry.UpstreamErrors, "raw attempts must be released before async queueing")
require.NotNil(t, job.entry.UpstreamErrorsJSON)
queuedEvents, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, queuedEvents, 1)
require.Equal(t, http.StatusTooManyRequests, queuedEvents[0].UpstreamStatusCode)
flushOpsErrorLogBatch([]opsErrorLogJob{job})
require.Equal(t, 1, repo.insertCalls)
require.Len(t, repo.entries, 1)
persisted := repo.entries[0]
require.Equal(t, http.StatusOK, persisted.StatusCode, "recovered telemetry must remain outside failed-request SLA")
require.Equal(t, "upstream", persisted.ErrorPhase)
require.Equal(t, "upstream_error", persisted.ErrorType)
require.Equal(t, "Recovered upstream error 429: earlier attempt was rate limited", persisted.ErrorMessage)
require.NotNil(t, persisted.UpstreamErrorsJSON)
persistedEvents, err := service.ParseOpsUpstreamErrors(*persisted.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, persistedEvents, 1)
require.Equal(t, http.StatusTooManyRequests, persistedEvents[0].UpstreamStatusCode)
}
func TestOpsErrorLoggerMiddleware_RecoveredTelemetryFiltersSkipMonitoringAttempts(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusTooManyRequests, Message: "visible retry"},
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true},
})
c.JSON(http.StatusOK, gin.H{"status": "completed"})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "Recovered upstream error 429: visible retry", job.entry.ErrorMessage)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 1)
require.Equal(t, "visible retry", events[0].Message)
}
func TestOpsErrorLoggerMiddleware_RecoveredTelemetrySkipsAllHiddenAttempts(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
UpstreamStatusCode: http.StatusTooManyRequests,
Message: "hidden retry",
SkipMonitoring: true,
}})
c.JSON(http.StatusOK, gin.H{"status": "completed"})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, int64(0), OpsErrorLogQueueLength())
}
func TestOpsErrorLoggerMiddleware_IntermediateSkipMonitoringDoesNotHideFinalVisibleFailure(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true},
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "visible final"},
})
c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{"type": "upstream_error", "message": "visible final"}})
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusServiceUnavailable, job.entry.StatusCode)
require.Equal(t, "visible final", job.entry.ErrorMessage)
}
func TestOpsErrorLoggerMiddleware_CapturesSplitResponsesFailedSSE(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
setOpsRequestContext(c, "gpt-5.5", true)
c.Status(http.StatusOK)
_, _ = c.Writer.Write([]byte("event: response."))
_, _ = c.Writer.Write([]byte("failed\n"))
_, _ = c.Writer.Write([]byte(`data: {"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"Too many pending requests"}}}`))
_, _ = c.Writer.Write([]byte("\n\n"))
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusTooManyRequests, job.entry.StatusCode)
require.Equal(t, "rate_limit_error", job.entry.ErrorType)
require.Contains(t, job.entry.ErrorMessage, "Too many pending requests")
}
func TestOpsCaptureWriter_CapturesSplitDataOnlyTerminalMarkers(t *testing.T) {
tests := []struct {
name string
prefix string
suffix string
wantType string
wantCode string
wantError string
}{
{
name: "response failed with space",
prefix: `data: {"type":"response.`,
suffix: `failed","response":{"error":{"code":"server_is_overloaded","message":"busy"}}}`,
wantType: "overloaded_error",
wantCode: "server_is_overloaded",
wantError: "busy",
},
{
name: "response failed without space",
prefix: `data:{"type":"response.`,
suffix: `failed","error":{"code":"rate_limit_exceeded","message":"slow down"}}`,
wantType: "rate_limit_error",
wantCode: "rate_limit_exceeded",
wantError: "slow down",
},
{
name: "error with space",
prefix: `data: {"type":"er`,
suffix: `ror","error":{"type":"invalid_request_error","code":"invalid_request","message":"bad input"}}`,
wantType: "invalid_request_error",
wantCode: "invalid_request",
wantError: "bad input",
},
{
name: "error without space",
prefix: `data:{"type":"er`,
suffix: `ror","error":{"type":"authentication_error","code":"authentication_failed","message":"sign in"}}`,
wantType: "authentication_error",
wantCode: "authentication_failed",
wantError: "sign in",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
state.captureResponseChunk([]byte(tt.prefix), http.StatusOK)
require.Empty(t, state.buf.Bytes(), "partial frame must remain in the bounded probe")
state.captureResponseChunk([]byte(tt.suffix+"\n\n"), http.StatusOK)
parsed := parseOpsErrorResponse(state.buf.Bytes())
require.True(t, state.sseCapturing)
require.True(t, parsed.StreamFailure)
require.Equal(t, tt.wantType, parsed.ErrorType)
require.Equal(t, tt.wantCode, parsed.Code)
require.Equal(t, tt.wantError, parsed.Message)
require.LessOrEqual(t, len(state.probe), opsTerminalSSEFrameProbeLimit)
})
}
}
func TestOpsErrorLoggerMiddleware_StreamFailureUsesTerminalErrorOverAttemptContext(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
service.SetOpsUpstreamError(c, http.StatusBadGateway, "Upstream transport error", "earlier attempt failed")
c.Status(http.StatusOK)
_, _ = c.Writer.WriteString("event: er")
_, _ = c.Writer.WriteString("ror\n")
_, _ = c.Writer.WriteString(`data: {"type":"error","error":{"type":"invalid_request_error","code":"context_length_exceeded","message":"input exceeds the context window"}}`)
_, _ = c.Writer.WriteString("\n\n")
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusBadRequest, job.entry.StatusCode)
require.Equal(t, "invalid_request_error", job.entry.ErrorType)
require.NotNil(t, job.entry.UpstreamStatusCode)
require.Equal(t, http.StatusBadRequest, *job.entry.UpstreamStatusCode)
require.NotNil(t, job.entry.UpstreamErrorMessage)
require.Equal(t, "input exceeds the context window", *job.entry.UpstreamErrorMessage)
}
func TestOpsErrorLoggerMiddleware_PrefersContextRequestID(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Header("X-Request-Id", "response-header-id")
c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": "bad input"}})
})
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
request = request.WithContext(context.WithValue(request.Context(), ctxkey.RequestID, "context-request-id"))
router.ServeHTTP(recorder, request)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, "context-request-id", job.entry.RequestID)
}
func TestNormalizeOpsPersistentUserAgentBoundsAndPreservesUTF8(t *testing.T) {
value := strings.Repeat("a", opsErrorLogMaxUserAgentBytes-1) + "你" + strings.Repeat("b", 32)
got := normalizeOpsPersistentUserAgent(" " + value + " ")
@@ -382,6 +690,23 @@ func TestLogOpsStreamError_SkipWhenPassthroughSkipMonitoring(t *testing.T) {
require.Equal(t, int64(0), OpsErrorLogEnqueuedTotal())
}
func TestShouldSkipFinalOpsFailureUsesOnlyFinalAttemptRule(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusBadGateway, Message: "hidden intermediate", SkipMonitoring: true},
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "visible final"},
})
require.False(t, shouldSkipFinalOpsFailure(c))
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusBadGateway, Message: "visible intermediate"},
nil,
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "hidden final", SkipMonitoring: true},
})
require.True(t, shouldSkipFinalOpsFailure(c))
}
// MarkOpsStreamError 采用「首个标记生效」:后续的通用兜底帧不得覆盖根因错误。
func TestMarkOpsStreamError_FirstWins(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -398,6 +723,35 @@ func TestMarkOpsStreamError_FirstWins(t *testing.T) {
require.Equal(t, http.StatusTooManyRequests, se.IntendedStatus)
}
func TestLogOpsStreamError_RecordsOneFailurePerWebSocketTurn(t *testing.T) {
setupOpsErrorLogTestQueue(t, 4)
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
service.SetOpenAIClientTransport(c, service.OpenAIClientTransportWS)
service.BeginOpsStreamTurn(c, 1)
service.MarkOpsStreamFailure(c, "rate_limit_error", "rate_limit_exceeded", "turn one failed", http.StatusTooManyRequests)
service.MarkOpsStreamError(c, "upstream_error", "generic duplicate for turn one", http.StatusBadGateway)
service.BeginOpsStreamTurn(c, 2)
service.MarkOpsStreamFailure(c, "permission_error", "permission_denied", "turn two failed", http.StatusForbidden)
streamErrors := service.GetOpsStreamErrors(c)
require.Len(t, streamErrors, 2)
require.Equal(t, 1, streamErrors[0].Turn)
require.Equal(t, 2, streamErrors[1].Turn)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
logOpsStreamError(c, ops, http.StatusSwitchingProtocols)
require.Equal(t, int64(2), OpsErrorLogQueueLength())
first := <-opsErrorLogQueue
second := <-opsErrorLogQueue
require.Equal(t, "turn one failed", first.entry.ErrorMessage)
require.Equal(t, http.StatusTooManyRequests, first.entry.StatusCode)
require.Equal(t, "turn two failed", second.entry.ErrorMessage)
require.Equal(t, http.StatusForbidden, second.entry.StatusCode)
}
func TestIsKnownOpsErrorType(t *testing.T) {
known := []string{
"invalid_request_error",
@@ -497,6 +851,107 @@ func TestClassifyOpsRoutingCapacityMarkerExcludesMaskedSelectionFailureFromSLA(t
require.Equal(t, "gateway", errorSource)
}
func TestClassifyOpsLocalModelConfigurationRejection(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalModelConfiguration)
phase, isBusinessLimited, errorOwner, errorSource := classifyOpsErrorLog(
c,
"model_not_found",
"Model \"gpt-missing\" is not supported by any configured account in this group",
"",
http.StatusNotFound,
)
require.Equal(t, "routing", phase)
require.True(t, isBusinessLimited)
require.Equal(t, "platform", errorOwner)
require.Equal(t, "gateway", errorSource)
}
func TestClassifyOpsLocalModelConfigurationOverridesStaleUpstreamMarkers(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalModelConfiguration)
c.Set(service.OpsUpstreamStatusCodeKey, http.StatusUnauthorized)
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
Stage: string(service.GatewayFailureStageAccountAuth),
UpstreamStatusCode: http.StatusUnauthorized,
}})
phase, limited, owner, source := classifyOpsErrorLog(c, "model_not_found", "unsupported configured model", "", http.StatusNotFound)
require.Equal(t, "routing", phase)
require.True(t, limited)
require.Equal(t, "platform", owner)
require.Equal(t, "gateway", source)
}
func TestClassifyOpsLocalModelConfigurationRequiresMarkerAndReason(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Set(service.OpsClientBusinessLimitedReasonKey, service.OpsClientBusinessLimitedReasonLocalModelConfiguration)
c.Set(service.OpsUpstreamStatusCodeKey, http.StatusBadGateway)
phase, limited, owner, source := classifyOpsErrorLog(c, "upstream_error", "provider failed", "", http.StatusBadGateway)
require.Equal(t, "upstream", phase)
require.False(t, limited)
require.Equal(t, "provider", owner)
require.Equal(t, "upstream_http", source)
}
func TestOpsErrorLoggerMiddleware_LocalModelConfigurationFields(t *testing.T) {
setupOpsErrorLogTestQueue(t, 1)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/chat/completions", func(c *gin.Context) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalModelConfiguration)
c.Set(opsAccountIDKey, int64(99))
c.Set(opsUpstreamModelKey, "stale-upstream-model")
setActualUpstreamEndpoint(c, "/v1/chat/completions")
c.Set(service.OpsUpstreamStatusCodeKey, http.StatusUnauthorized)
c.Set(service.OpsUpstreamErrorMessageKey, "stale upstream error")
c.Set(service.OpsUpstreamErrorDetailKey, "stale upstream detail")
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{
Stage: string(service.GatewayFailureStageAccountAuth),
UpstreamStatusCode: http.StatusUnauthorized,
Message: "stale auth failure",
}})
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
"type": "model_not_found",
"message": "Model \"gpt-missing\" is not supported by any configured account in this group",
},
})
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
require.JSONEq(t, `{"error":{"type":"model_not_found","message":"Model \"gpt-missing\" is not supported by any configured account in this group"}}`, w.Body.String())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusNotFound, job.entry.StatusCode)
require.Equal(t, "routing", job.entry.ErrorPhase)
require.True(t, job.entry.IsBusinessLimited)
require.Equal(t, "platform", job.entry.ErrorOwner)
require.Equal(t, "gateway", job.entry.ErrorSource)
require.Nil(t, job.entry.AccountID)
require.Nil(t, job.entry.UpstreamStatusCode)
require.Nil(t, job.entry.UpstreamErrors)
require.Nil(t, job.entry.UpstreamErrorMessage)
require.Nil(t, job.entry.UpstreamErrorDetail)
require.Empty(t, job.entry.UpstreamModel)
require.Empty(t, job.entry.UpstreamEndpoint)
}
func TestClassifyOpsAuthClientErrorsExcludedFromSLA(t *testing.T) {
tests := []struct {
name string
@@ -621,7 +1076,11 @@ func TestClassifyOpsAuthClientErrorsExcludedFromSLA(t *testing.T) {
errType := normalizeOpsErrorType(tt.errType, tt.code)
phase, isBusinessLimited, errorOwner, errorSource := classifyOpsErrorLog(c, errType, tt.message, tt.code, tt.status)
require.Equal(t, "api_error", errType)
wantErrType := "api_error"
if tt.errType == "permission_error" {
wantErrType = "permission_error"
}
require.Equal(t, wantErrType, errType)
require.Equal(t, "auth", phase)
require.True(t, isBusinessLimited)
require.Equal(t, "client", errorOwner)
@@ -772,7 +1231,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
message: "This group is restricted to Claude Code clients (/v1/messages only)",
code: "",
status: http.StatusForbidden,
wantErrType: "api_error",
wantErrType: "permission_error",
wantPhase: "request",
},
{
@@ -781,7 +1240,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
message: "Image generation is not enabled for this group",
code: "",
status: http.StatusForbidden,
wantErrType: "api_error",
wantErrType: "permission_error",
wantPhase: "request",
},
{
@@ -808,7 +1267,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
message: "model claude-3-5-sonnet not in whitelist",
code: "",
status: http.StatusForbidden,
wantErrType: "api_error",
wantErrType: "permission_error",
wantPhase: "request",
},
{
@@ -835,7 +1294,7 @@ func TestClassifyOpsLocalBusinessLimitErrorsExcludedFromSLA(t *testing.T) {
message: "openai service_tier=priority is not allowed for model gpt-5.5",
code: "",
status: http.StatusForbidden,
wantErrType: "api_error",
wantErrType: "permission_error",
wantPhase: "request",
},
{
@@ -1146,6 +1605,254 @@ func TestParseOpsErrorResponsePreservesNestedStringCode(t *testing.T) {
require.Equal(t, "API Key 所属分组已删除", parsed.Message)
}
func TestParseOpsErrorResponsePreservesStructuredTopLevelSemantics(t *testing.T) {
tests := []struct {
name string
body string
wantType string
wantCode string
wantMsg string
}{
{
name: "model not found",
body: `{"type":"model_not_found","code":404,"message":"model unavailable"}`,
wantType: "model_not_found",
wantCode: "404",
wantMsg: "model unavailable",
},
{
name: "string error",
body: `{"type":"service_unavailable","code":"temporarily_unavailable","error":"capacity exhausted"}`,
wantType: "service_unavailable",
wantCode: "temporarily_unavailable",
wantMsg: "capacity exhausted",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parsed := parseOpsErrorResponse([]byte(tt.body))
require.Equal(t, tt.wantType, normalizeOpsErrorType(parsed.ErrorType, parsed.Code))
require.Equal(t, tt.wantCode, parsed.Code)
require.Equal(t, tt.wantMsg, parsed.Message)
})
}
}
func TestApplyOpsUpstreamFieldsUsesLastNonNilAttempt(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
service.SetOpsUpstreamError(c, http.StatusUnauthorized, "stale context", "stale detail")
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusTooManyRequests, Message: "first attempt", Detail: "first detail"},
nil,
{UpstreamStatusCode: http.StatusServiceUnavailable, Message: "final attempt", Detail: "final detail"},
nil,
})
entry := &service.OpsInsertErrorLogInput{}
applyOpsUpstreamFieldsFromContext(c, entry)
require.NotNil(t, entry.UpstreamStatusCode)
require.Equal(t, http.StatusServiceUnavailable, *entry.UpstreamStatusCode)
require.NotNil(t, entry.UpstreamErrorMessage)
require.Equal(t, "final attempt", *entry.UpstreamErrorMessage)
require.NotNil(t, entry.UpstreamErrorDetail)
require.Equal(t, "final detail", *entry.UpstreamErrorDetail)
require.Len(t, entry.UpstreamErrors, 4)
}
func TestApplyOpsUpstreamFieldsFinalStatuslessAttemptClearsStaleContext(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
service.SetOpsUpstreamError(c, http.StatusBadGateway, "stale response", "stale body")
c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{
{UpstreamStatusCode: http.StatusBadGateway, Message: "first response"},
{Kind: "request_error", Message: "final transport failure", Detail: "connection reset"},
})
entry := &service.OpsInsertErrorLogInput{}
applyOpsUpstreamFieldsFromContext(c, entry)
require.Nil(t, entry.UpstreamStatusCode)
require.NotNil(t, entry.UpstreamErrorMessage)
require.Equal(t, "final transport failure", *entry.UpstreamErrorMessage)
require.NotNil(t, entry.UpstreamErrorDetail)
require.Equal(t, "connection reset", *entry.UpstreamErrorDetail)
}
func TestOpsCaptureWriter_ProtocolLevelTerminalFrameDetection(t *testing.T) {
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
chunks := []string{
"event : response.failed\r\n",
"data: { \"response\" : { \"error\" : { \"message\" : \"busy\", \"code\" : \"service_unavailable\" } },",
" \"type\" : \"response.failed\" }\r\n\r\n",
}
for _, chunk := range chunks {
state.captureResponseChunk([]byte(chunk), http.StatusOK)
}
parsed := parseOpsErrorResponse(state.buf.Bytes())
require.True(t, state.sseCapturing)
require.True(t, parsed.StreamFailure)
require.Equal(t, "service_unavailable_error", parsed.ErrorType)
require.Equal(t, "service_unavailable", parsed.Code)
require.Equal(t, "busy", parsed.Message)
require.Equal(t, http.StatusServiceUnavailable, inferStreamFailureStatus(nil, parsed))
}
func TestParseOpsSSEFailure_TopLevelErrorsAndUnknownStatus(t *testing.T) {
tests := []struct {
name string
body string
wantType string
wantStatus int
}{
{
name: "top-level permission",
body: "event: error\ndata: {\"message\":\"denied\",\"code\":\"permission_denied\",\"type\":\"error\"}\n\n",
wantType: "permission_error",
wantStatus: http.StatusForbidden,
},
{
name: "top-level unavailable",
body: "data: {\"message\":\"busy\",\"type\":\"error\",\"code\":\"service_unavailable\"}\n\n",
wantType: "service_unavailable_error",
wantStatus: http.StatusServiceUnavailable,
},
{
name: "unknown terminal",
body: "event: response.failed\ndata: {\"type\":\"response.failed\",\"error\":{\"code\":\"new_provider_code\",\"message\":\"failed\"}}\n\n",
wantType: "upstream_error",
wantStatus: http.StatusBadGateway,
},
{
name: "explicit terminal status",
body: "event: error\ndata: {\"type\":\"error\",\"status_code\":429,\"code\":\"new_rate_code\",\"message\":\"slow down\"}\n\n",
wantType: "api_error",
wantStatus: http.StatusTooManyRequests,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parsed := parseOpsErrorResponse([]byte(tt.body))
require.True(t, parsed.StreamFailure)
require.Equal(t, tt.wantType, parsed.ErrorType)
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
service.SetOpsUpstreamError(c, http.StatusUnauthorized, "old attempt", "")
require.Equal(t, tt.wantStatus, inferStreamFailureStatus(c, parsed), "terminal status must not inherit an earlier attempt")
})
}
}
func TestOpsCaptureWriter_OversizedNonTerminalFrameRemainsBounded(t *testing.T) {
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
state.captureResponseChunk([]byte("data: "+strings.Repeat("x", opsTerminalSSEFrameProbeLimit*2)+"\n\n"), http.StatusOK)
require.Empty(t, state.buf.Bytes())
require.LessOrEqual(t, cap(state.probe), opsTerminalSSEFrameProbeLimit)
state.captureResponseChunk([]byte("event: error\ndata: {\"type\":\"error\",\"code\":\"permission_denied\",\"message\":\"denied\"}\n\n"), http.StatusOK)
require.True(t, state.sseCapturing)
require.NotEmpty(t, state.buf.Bytes())
}
func TestOpsCaptureWriter_TerminalMetadataSurvivesBodyCaptureTruncation(t *testing.T) {
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
frame := "event: response.failed\ndata: {\"padding\":\"" + strings.Repeat("x", opsCaptureWriterLimit) + "\",\"type\":\"response.failed\",\"error\":{\"code\":\"service_unavailable\",\"message\":\"busy\"}}\n\n"
state.captureResponseChunk([]byte(frame), http.StatusOK)
state.finalizeResponseCapture()
require.Len(t, state.buf.Bytes(), opsCaptureWriterLimit)
require.True(t, parseOpsErrorResponse(state.buf.Bytes()).StreamFailure, "the bounded parser must fail closed from the terminal event line")
require.True(t, state.terminalFound)
require.Equal(t, "service_unavailable_error", state.terminalError.ErrorType)
require.Equal(t, "busy", state.terminalError.Message)
}
func TestOpsErrorLoggerMiddleware_LargeTerminalFrameUsesEventFallback(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Status(http.StatusOK)
_, _ = c.Writer.WriteString("event: response.failed\n")
_, _ = c.Writer.WriteString("data: {\"authorization\":\"Bearer must-not-persist\",\"padding\":\"" + strings.Repeat("x", opsTerminalSSEFrameProbeLimit*2) + "\"}")
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusBadGateway, job.entry.StatusCode)
require.Equal(t, "upstream_error", job.entry.ErrorType)
require.Equal(t, "upstream stream failed", job.entry.ErrorMessage)
require.NotContains(t, job.entry.ErrorMessage, "must-not-persist")
require.NotContains(t, job.entry.ErrorBody, "must-not-persist")
require.Contains(t, job.entry.ErrorBody, `"payload_truncated":true`)
}
func TestOpsErrorLoggerMiddleware_DetectsTerminalDataAtEOFWithoutBlankLine(t *testing.T) {
setupOpsErrorLogTestQueue(t, 2)
gin.SetMode(gin.TestMode)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.POST("/v1/responses", func(c *gin.Context) {
c.Status(http.StatusOK)
_, _ = c.Writer.WriteString(`data: {"message":"denied","code":"permission_denied","type":"error"}`)
})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil))
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, int64(1), OpsErrorLogQueueLength())
job := <-opsErrorLogQueue
require.Equal(t, http.StatusForbidden, job.entry.StatusCode)
require.Equal(t, "permission_error", job.entry.ErrorType)
require.Equal(t, "denied", job.entry.ErrorMessage)
}
func TestOpsCaptureWriter_DetectsCROnlySSEFrame(t *testing.T) {
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
state.captureResponseChunk([]byte("data: {\"type\":\"error\",\"code\":\"service_unavailable\",\"message\":\"busy\"}\r\r"), http.StatusOK)
require.True(t, state.sseCapturing)
parsed := parseOpsErrorResponse(state.buf.Bytes())
require.True(t, parsed.StreamFailure)
require.Equal(t, "service_unavailable_error", parsed.ErrorType)
}
func TestSanitizeOpsSSEDataForPersistence_RedactsJSONFields(t *testing.T) {
body := []byte("event: error\ndata: {\"type\":\"error\",\"authorization\":\"Bearer secret\",\ndata: \"nested\":{\"api_key\":\"sk-secret\"}}\n\n")
sanitized := sanitizeOpsSSEDataForPersistence(body)
require.NotContains(t, sanitized, "Bearer secret")
require.NotContains(t, sanitized, "sk-secret")
require.Contains(t, sanitized, `"authorization":"[REDACTED]"`)
require.Contains(t, sanitized, `"api_key":"[REDACTED]"`)
}
func TestSanitizeOpsSSEDataForPersistence_DropsTruncatedJSONFragment(t *testing.T) {
body := []byte("event: error\ndata: {\"type\":\"error\",\"authorization\":\"Bearer leaked")
sanitized := sanitizeOpsSSEDataForPersistence(body)
require.NotContains(t, sanitized, "Bearer leaked")
require.Contains(t, sanitized, `data: {"payload_truncated":true}`)
}
func BenchmarkOpsCaptureWriterSuccessfulSSEFrames(b *testing.B) {
frame := []byte("event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n")
state := &opsCaptureWriterState{limit: opsCaptureWriterLimit}
b.ReportAllocs()
for i := 0; i < b.N; i++ {
state.captureResponseChunk(frame, http.StatusOK)
}
if state.buf.Len() != 0 {
b.Fatal("successful frames must not be captured")
}
}
func TestSetOpsEndpointContext_SetsContextKeys(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -15,11 +15,11 @@ func TestOpsCaptureWriterDoesNotCopyIngressRejectBody(t *testing.T) {
context, _ := gin.CreateTestContext(httptest.NewRecorder())
writer := acquireOpsCaptureWriter(context.Writer)
defer releaseOpsCaptureWriter(writer)
writer.ctx = context
writer.setContext(context)
context.Writer = writer
middleware2.MarkIngressRejected(context, middleware2.IngressRejectInvalidAPIKey)
context.Status(http.StatusUnauthorized)
_, err := context.Writer.WriteString(`{"code":"INVALID_API_KEY","message":"Invalid API key"}`)
require.NoError(t, err)
require.Zero(t, writer.buf.Len())
require.Empty(t, writer.capturedBytes())
}

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