From f9fac7acbe8084eb65a9a04cb89a9b6c4110f36d Mon Sep 17 00:00:00 2001 From: shaw Date: Thu, 11 Jun 2026 22:57:54 +0800 Subject: [PATCH] =?UTF-8?q?test(safety-net):=20Phase-0=20=E5=9B=9E?= =?UTF-8?q?=E5=BD=92=E5=AE=89=E5=85=A8=E7=BD=91=E2=80=94=E2=80=94=E7=89=B9?= =?UTF-8?q?=E5=BE=81=E5=8C=96/=E4=B8=8D=E5=8F=98=E9=87=8F=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=20+=20=E5=9F=BA=E5=87=86=E5=9F=BA=E7=BA=BF=20+=20CI?= =?UTF-8?q?=20=E9=97=A8=E7=A6=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 插件化改造前置:把'绝不回归'变成机器可验证的硬约束。 - 透传特征化:anthropic 逐字节/SSE 逐事件、openai、gemini 路径与错误透传 - 计费不变量:5m/1h 双档缓存价、倍率叠加顺序、overages、端到端金额与配额增量 - 调度不变量:粘性会话、failover 10/3/3 完整循环、并发槽配平、等待队列 - 热路径基准基线 + scripts/bench-baseline.sh(allocs 严格/ns 15%/基准缺失 FAIL) - Makefile test-invariants 目标 + CI invariants job;修复失效的 test-e2e 脚本引用 - .gitignore:放行 backend/scripts 与 docs/plugin-architecture(既有白名单模式) 注:本分支按整体门禁验证(部分测试夹具引用后续提交的签名,中间提交不保证独立编译)。 --- .github/workflows/backend-ci.yml | 17 + .gitignore | 4 + backend/Makefile | 17 +- ...gateway_intercept_characterization_test.go | 367 +++++++++++ .../scheduling_invariants_failover_test.go | 614 ++++++++++++++++++ .../scheduling_invariants_slots_test.go | 289 +++++++++ ...ntigravity_gemini_characterization_test.go | 188 ++++++ .../service/billing_invariants_cost_test.go | 291 +++++++++ .../service/billing_invariants_e2e_test.go | 546 ++++++++++++++++ .../billing_invariants_overages_test.go | 141 ++++ .../billing_invariants_preflight_test.go | 226 +++++++ .../service/gateway_forward_benchmark_test.go | 83 +++ ...teway_passthrough_characterization_test.go | 411 ++++++++++++ ...teway_passthrough_characterization_test.go | 181 ++++++ .../service/scheduling_invariants_test.go | 480 ++++++++++++++ backend/scripts/bench-baseline.sh | 102 +++ backend/testdata/bench/baseline.env.txt | 4 + backend/testdata/bench/baseline.txt | 48 ++ 18 files changed, 4007 insertions(+), 2 deletions(-) create mode 100644 backend/internal/handler/gateway_intercept_characterization_test.go create mode 100644 backend/internal/handler/scheduling_invariants_failover_test.go create mode 100644 backend/internal/handler/scheduling_invariants_slots_test.go create mode 100644 backend/internal/service/antigravity_gemini_characterization_test.go create mode 100644 backend/internal/service/billing_invariants_cost_test.go create mode 100644 backend/internal/service/billing_invariants_e2e_test.go create mode 100644 backend/internal/service/billing_invariants_overages_test.go create mode 100644 backend/internal/service/billing_invariants_preflight_test.go create mode 100644 backend/internal/service/gateway_forward_benchmark_test.go create mode 100644 backend/internal/service/gateway_passthrough_characterization_test.go create mode 100644 backend/internal/service/openai_gateway_passthrough_characterization_test.go create mode 100644 backend/internal/service/scheduling_invariants_test.go create mode 100644 backend/scripts/bench-baseline.sh create mode 100644 backend/testdata/bench/baseline.env.txt create mode 100644 backend/testdata/bench/baseline.txt diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml index fb4d0ce652..396b4db4ae 100644 --- a/.github/workflows/backend-ci.yml +++ b/.github/workflows/backend-ci.yml @@ -28,6 +28,23 @@ jobs: working-directory: backend run: make test-integration + # 插件化改造回归安全网:特征化/不变量测试作为独立命名的合并门禁 + # (这些测试也包含在 test-unit 中,独立 job 是为了回归信号清晰可辨) + # 详见 .claude/plugin-refactor/INVARIANTS.md + invariants: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + - uses: actions/setup-go@v6 + with: + go-version-file: backend/go.mod + check-latest: false + cache: true + cache-dependency-path: backend/go.sum + - name: Regression invariants (plugin-refactor safety net) + working-directory: backend + run: make test-invariants + frontend: runs-on: ubuntu-latest steps: diff --git a/.gitignore b/.gitignore index 0f4e3b0625..5ded43b37e 100644 --- a/.gitignore +++ b/.gitignore @@ -133,6 +133,10 @@ docs/* !docs/ADMIN_PAYMENT_INTEGRATION_API.md !docs/legal/ !docs/legal/*.md +!docs/plugin-architecture/ +!docs/plugin-architecture/*.md +!backend/scripts/ +!backend/scripts/*.sh .serena/ .codex/ frontend/coverage/ diff --git a/backend/Makefile b/backend/Makefile index 7084ccb93e..1a5bcf4d2f 100644 --- a/backend/Makefile +++ b/backend/Makefile @@ -1,4 +1,4 @@ -.PHONY: build generate test test-unit test-integration test-e2e +.PHONY: build generate test test-unit test-integration test-e2e test-invariants new-module VERSION ?= $(shell tr -d '\r\n' < ./cmd/server/VERSION) LDFLAGS ?= -s -w -X main.Version=$(VERSION) @@ -20,8 +20,21 @@ test-unit: test-integration: go test -tags=integration ./... +# 插件化改造回归安全网(特征化/不变量测试,详见 .claude/plugin-refactor/INVARIANTS.md) +test-invariants: + go test -tags=unit -count=1 -run 'Characterization|Invariant' ./internal/... + +# 生成插件模块骨架:make new-module ID=job.foo +# (生成后按提示在 internal/modules/standard/imports.go 手工加 import 行) +new-module: +ifndef ID + $(error ID is required: make new-module ID=job.foo) +endif + go run ./tools/newmodule -id "$(ID)" + +# scripts/e2e-test.sh 已不存在,test-e2e 与 test-e2e-local 等价(env 驱动的真实服务 e2e) test-e2e: - ./scripts/e2e-test.sh + go test -tags=e2e -v -timeout=300s ./internal/integration/... test-e2e-local: go test -tags=e2e -v -timeout=300s ./internal/integration/... diff --git a/backend/internal/handler/gateway_intercept_characterization_test.go b/backend/internal/handler/gateway_intercept_characterization_test.go new file mode 100644 index 0000000000..fe970b18a3 --- /dev/null +++ b/backend/internal/handler/gateway_intercept_characterization_test.go @@ -0,0 +1,367 @@ +//go:build unit + +// Phase-0 TASK-002 特征化测试:网关拦截链路(INVARIANTS I-7.1 / I-7.2)。 +// 通过完整的 GatewayHandler.Messages / CountTokens 入口驱动,断言客户端最终 +// 收到的状态码与错误体: +// - I-7.1 内容审核拦截:403 + content_policy_violation 错误体;审核服务失败 = 放行(fail-open); +// - I-7.2 Claude Code 版本检查:低版本 400 + 升级提示;/count_tokens 路径豁免。 +// +// 复用 gateway_handler_warmup_intercept_unit_test.go 的 newTestGatewayHandler 夹具: +// antigravity 账号 intercept_warmup_requests=true 时,Warmup 请求在转发上游前被 +// mock 拦截返回 200,使"请求被放行"成为可观测结果(无需真实上游)。 +// +// 注意:Claude Code 版本上下限走 service 包内全局 60s TTL 缓存 +// (setting_service.go versionBoundsCache)。本文件所有版本测试固定使用同一组 +// 上下限(min=1.5.0,无 max),避免同进程内缓存交叉污染。 +package handler + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + middleware "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" +) + +// passCharSettingRepo 是 service.SettingRepository 的内存实现。 +type passCharSettingRepo struct { + values map[string]string +} + +func (r *passCharSettingRepo) Get(_ context.Context, key string) (*service.Setting, error) { + if v, ok := r.values[key]; ok { + return &service.Setting{Key: key, Value: v}, nil + } + return nil, service.ErrSettingNotFound +} + +func (r *passCharSettingRepo) GetValue(_ context.Context, key string) (string, error) { + if v, ok := r.values[key]; ok { + return v, nil + } + return "", service.ErrSettingNotFound +} + +func (r *passCharSettingRepo) Set(_ context.Context, key, value string) error { + if r.values == nil { + r.values = map[string]string{} + } + r.values[key] = value + return nil +} + +func (r *passCharSettingRepo) GetMultiple(_ context.Context, keys []string) (map[string]string, error) { + out := map[string]string{} + for _, key := range keys { + if v, ok := r.values[key]; ok { + out[key] = v + } + } + return out, nil +} + +func (r *passCharSettingRepo) SetMultiple(_ context.Context, settings map[string]string) error { + for k, v := range settings { + _ = r.Set(context.Background(), k, v) + } + return nil +} + +func (r *passCharSettingRepo) GetAll(_ context.Context) (map[string]string, error) { + out := make(map[string]string, len(r.values)) + for k, v := range r.values { + out[k] = v + } + return out, nil +} + +func (r *passCharSettingRepo) Delete(_ context.Context, key string) error { + delete(r.values, key) + return nil +} + +// passCharModerationRepo 是 service.ContentModerationRepository 的空操作实现。 +type passCharModerationRepo struct{} + +func (r *passCharModerationRepo) CreateLog(context.Context, *service.ContentModerationLog) error { + return nil +} + +func (r *passCharModerationRepo) ListLogs(context.Context, service.ContentModerationLogFilter) ([]service.ContentModerationLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} + +func (r *passCharModerationRepo) CountFlaggedByUserSince(context.Context, int64, time.Time) (int, error) { + return 0, nil +} + +func (r *passCharModerationRepo) CleanupExpiredLogs(context.Context, time.Time, time.Time) (*service.ContentModerationCleanupResult, error) { + return &service.ContentModerationCleanupResult{}, nil +} + +// passCharInterceptFixture 构造完整的 Messages 链路夹具(拦截预热账号 + 上下文注入)。 +func passCharInterceptFixture(t *testing.T) (*GatewayHandler, func()) { + t.Helper() + + groupID := int64(7001) + accountID := int64(7101) + + group := &service.Group{ + ID: groupID, + Hydrated: true, + Platform: service.PlatformAnthropic, + Status: service.StatusActive, + } + account := &service.Account{ + ID: accountID, + Name: "pass-char-intercept", + Platform: service.PlatformAntigravity, + Type: service.AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "tok_char", + "intercept_warmup_requests": true, + }, + Extra: map[string]any{ + "mixed_scheduling": true, + }, + Concurrency: 1, + Priority: 1, + Status: service.StatusActive, + Schedulable: true, + AccountGroups: []service.AccountGroup{{AccountID: accountID, GroupID: groupID}}, + } + return newTestGatewayHandler(t, group, []*service.Account{account}) +} + +// passCharNewMessagesContext 构造带认证上下文的 /v1/messages 请求。 +func passCharNewMessagesContext(t *testing.T, path string, body []byte) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + + groupID := int64(7001) + group := &service.Group{ + ID: groupID, + Hydrated: true, + Platform: service.PlatformAnthropic, + Status: service.StatusActive, + } + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req = req.WithContext(context.WithValue(req.Context(), ctxkey.Group, group)) + c.Request = req + + apiKey := &service.APIKey{ + ID: 7301, + UserID: 7401, + GroupID: &groupID, + Status: service.StatusActive, + User: &service.User{ + ID: 7401, + Concurrency: 10, + Balance: 100, + }, + Group: group, + } + c.Set(string(middleware.ContextKeyAPIKey), apiKey) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.UserID, Concurrency: 10}) + return c, rec +} + +// passCharWarmupBody 返回触发预热拦截的最小请求体。 +func passCharWarmupBody() []byte { + return []byte(`{ + "model": "claude-sonnet-4-5", + "max_tokens": 256, + "messages": [{"role":"user","content":[{"type":"text","text":"Warmup"}]}] + }`) +} + +// passCharClaudeCodeWarmupBody 返回能通过 Claude Code 客户端校验且触发预热拦截的请求体。 +func passCharClaudeCodeWarmupBody() []byte { + deviceID := strings.Repeat("a", 64) + return []byte(`{ + "model": "claude-sonnet-4-5", + "max_tokens": 256, + "system": [{"type":"text","text":"You are Claude Code, Anthropic's official CLI for Claude."}], + "metadata": {"user_id": "user_` + deviceID + `_account__session_01234567-89ab-cdef-0123-456789abcdef"}, + "messages": [{"role":"user","content":[{"type":"text","text":"Warmup"}]}] + }`) +} + +// passCharSetClaudeCodeHeaders 设置 Claude Code 客户端校验所需的请求头。 +func passCharSetClaudeCodeHeaders(c *gin.Context, ua string) { + c.Request.Header.Set("User-Agent", ua) + c.Request.Header.Set("X-App", "cli") + c.Request.Header.Set("anthropic-beta", "claude-code-20250219") + c.Request.Header.Set("anthropic-version", "2023-06-01") +} + +// passCharModerationService 构造指向 mock 审核 API 的内容审核服务。 +func passCharModerationService(t *testing.T, moderationBaseURL string) *service.ContentModerationService { + t.Helper() + cfgJSON, err := json.Marshal(map[string]any{ + "enabled": true, + "mode": "pre_block", + "base_url": moderationBaseURL, + "api_keys": []string{"sk-audit"}, + "sample_rate": 100, + "all_groups": true, + "auto_ban_enabled": false, + "email_on_hit": false, + "retry_count": 0, + "timeout_ms": 2000, + }) + require.NoError(t, err) + + settingRepo := &passCharSettingRepo{values: map[string]string{ + service.SettingKeyRiskControlEnabled: "true", + service.SettingKeyContentModerationConfig: string(cfgJSON), + }} + return service.NewContentModerationService(settingRepo, &passCharModerationRepo{}, nil, nil, nil, nil, nil) +} + +// passCharRequireWarmupMock 断言响应为预热拦截 mock(请求被放行并完成)。 +func passCharRequireWarmupMock(t *testing.T, rec *httptest.ResponseRecorder) { + t.Helper() + require.Equal(t, http.StatusOK, rec.Code) + var resp map[string]any + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, "msg_mock_warmup", resp["id"]) +} + +// TestGatewayCharacterization_ContentModerationBlock 固化 I-7.1 拦截侧: +// 审核 API 判定命中(pre_block 模式)时,客户端收到 403 + +// {"type":"error","error":{"type":"content_policy_violation","message":<配置的拦截文案>}}。 +func TestGatewayCharacterization_ContentModerationBlock(t *testing.T) { + moderationSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"results":[{"flagged":true,"category_scores":{"sexual":0.99}}]}`)) + })) + defer moderationSrv.Close() + + h, cleanup := passCharInterceptFixture(t) + defer cleanup() + h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL)) + + c, rec := passCharNewMessagesContext(t, "/v1/messages", passCharWarmupBody()) + h.Messages(c) + + require.Equal(t, http.StatusForbidden, rec.Code, "内容审核命中应返回 403") + require.JSONEq(t, + `{"type":"error","error":{"type":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`, + rec.Body.String()) +} + +// TestGatewayCharacterization_ContentModerationFailOpen 固化 I-7.1 fail-open 侧: +// 审核 API 自身失败(HTTP 500)时请求必须放行——客户端收到正常业务响应 +// (此处为预热拦截 mock 200),而不是 403。 +func TestGatewayCharacterization_ContentModerationFailOpen(t *testing.T) { + moderationSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "moderation backend exploded", http.StatusInternalServerError) + })) + defer moderationSrv.Close() + + h, cleanup := passCharInterceptFixture(t) + defer cleanup() + h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL)) + + c, rec := passCharNewMessagesContext(t, "/v1/messages", passCharWarmupBody()) + h.Messages(c) + + passCharRequireWarmupMock(t, rec) +} + +// passCharVersionBoundSettingService 返回固定 min=1.5.0(无 max)的 SettingService。 +// 注意:版本上下限有进程内全局 60s 缓存,本文件所有版本测试统一使用该配置。 +func passCharVersionBoundSettingService() *service.SettingService { + repo := &passCharSettingRepo{values: map[string]string{ + service.SettingKeyMinClaudeCodeVersion: "1.5.0", + }} + return service.NewSettingService(repo, &config.Config{}) +} + +// TestGatewayCharacterization_ClaudeCodeVersionCheck 固化 I-7.2 主路径: +// Claude Code 客户端(UA claude-cli/x.y.z + 完整客户端特征)在 /v1/messages 上: +// - CLI 版本低于最低要求 → 400 invalid_request_error + 升级提示; +// - CLI 版本满足要求 → 请求放行(预热拦截 mock 200)。 +func TestGatewayCharacterization_ClaudeCodeVersionCheck(t *testing.T) { + t.Run("低版本_400升级提示", func(t *testing.T) { + h, cleanup := passCharInterceptFixture(t) + defer cleanup() + h.settingService = passCharVersionBoundSettingService() + + c, rec := passCharNewMessagesContext(t, "/v1/messages", passCharClaudeCodeWarmupBody()) + passCharSetClaudeCodeHeaders(c, "claude-cli/1.0.0 (external)") + h.Messages(c) + + require.Equal(t, http.StatusBadRequest, rec.Code) + var resp struct { + Type string `json:"type"` + Error struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, "error", resp.Type) + require.Equal(t, "invalid_request_error", resp.Error.Type) + require.Contains(t, resp.Error.Message, "Your Claude Code version (1.0.0) is below the minimum required version (1.5.0)") + require.Contains(t, resp.Error.Message, "npm update -g @anthropic-ai/claude-code", "错误信息应包含升级指引") + }) + + t.Run("版本达标_请求放行", func(t *testing.T) { + h, cleanup := passCharInterceptFixture(t) + defer cleanup() + h.settingService = passCharVersionBoundSettingService() + + c, rec := passCharNewMessagesContext(t, "/v1/messages", passCharClaudeCodeWarmupBody()) + passCharSetClaudeCodeHeaders(c, "claude-cli/2.0.0 (external)") + h.Messages(c) + + passCharRequireWarmupMock(t, rec) + }) +} + +// TestGatewayCharacterization_ClaudeCodeVersionCheck_CountTokensExempt 固化 I-7.2 豁免侧: +// /v1/messages/count_tokens 不做版本检查——低版本 Claude Code CLI 不会收到版本 400, +// 请求继续走 count_tokens 业务(antigravity 账号当前返回 404 not_found_error)。 +func TestGatewayCharacterization_ClaudeCodeVersionCheck_CountTokensExempt(t *testing.T) { + h, cleanup := passCharInterceptFixture(t) + defer cleanup() + h.settingService = passCharVersionBoundSettingService() + + body := []byte(`{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`) + c, rec := passCharNewMessagesContext(t, "/v1/messages/count_tokens", body) + passCharSetClaudeCodeHeaders(c, "claude-cli/1.0.0 (external)") + h.CountTokens(c) + + var resp struct { + Type string `json:"type"` + Error struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.NotContains(t, resp.Error.Message, "below the minimum required version", + "count_tokens 必须豁免版本检查") + + // 固化当前实际行为:antigravity 账号不支持 count_tokens,返回 404 not_found_error。 + require.Equal(t, http.StatusNotFound, rec.Code) + require.Equal(t, "not_found_error", resp.Error.Type) +} diff --git a/backend/internal/handler/scheduling_invariants_failover_test.go b/backend/internal/handler/scheduling_invariants_failover_test.go new file mode 100644 index 0000000000..dc7fdee16e --- /dev/null +++ b/backend/internal/handler/scheduling_invariants_failover_test.go @@ -0,0 +1,614 @@ +//go:build unit + +// Phase-0 TASK-004 failover 不变量测试(INVARIANTS I-5.1 / I-5.2 / I-5.3、I-6.1 部分)。 +// +// 固化内容: +// - I-5.1 三平台换号上限默认值(anthropic=10 / gemini=3 / openai=3,双断言: +// 构造函数默认值 + 硬编码当前值);anthropic 平台通过完整 Messages 入口驱动 +// 真实 failover 循环:连续上游 500 时恰好尝试 maxAccountSwitches+1 个账号后 +// 返回 502 + 映射错误体; +// - I-5.2 失败账号进入排除集合不被重选(同一账号不会被尝试两次); +// - I-5.3 可重试错误(池模式)先同账号重试 maxSameAccountRetries 次再换号(完整链); +// - I-6.1 转发失败路径下账号/用户槽位获取-释放严格配平(配平计数器); +// - chat-completions 兼容路径耗尽错误体("All available accounts exhausted")。 +// +// 复用同包既有夹具:fakeSchedulerCache / fakeGroupRepo +// (gateway_handler_warmup_intercept_unit_test.go)、mockTempUnscheduler / +// newTestFailoverErr(failover_loop_test.go)。 +// 本文件新增的包级辅助类型/函数一律带 schedInv 前缀。 +package handler + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + middleware "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" +) + +// --------------------------------------------------------------------------- +// schedInv 夹具 +// --------------------------------------------------------------------------- + +// schedInvUpstream 是固定响应的 HTTPUpstream 桩,记录每次调用的 accountID。 +type schedInvUpstream struct { + mu sync.Mutex + status int + body string + panicOn bool + attempts []int64 +} + +var _ service.HTTPUpstream = (*schedInvUpstream)(nil) + +func (u *schedInvUpstream) Do(req *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + u.mu.Lock() + u.attempts = append(u.attempts, accountID) + shouldPanic := u.panicOn + status := u.status + body := u.body + u.mu.Unlock() + + if req != nil && req.Body != nil { + _, _ = io.Copy(io.Discard, req.Body) + _ = req.Body.Close() + } + if shouldPanic { + panic("schedInv: simulated upstream panic") + } + return &http.Response{ + StatusCode: status, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(body)), + }, nil +} + +func (u *schedInvUpstream) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return u.Do(req, proxyURL, accountID, accountConcurrency) +} + +func (u *schedInvUpstream) attemptedAccounts() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.attempts...) +} + +// schedInvCountingCache 是带真实容量语义的 ConcurrencyCache, +// 记录账号/用户槽位与等待计数的获取-释放配平。 +type schedInvCountingCache struct { + mu sync.Mutex + + accountHeld map[int64]map[string]struct{} + userHeld map[int64]map[string]struct{} + + accountAcquired int + accountReleased int + // accountReleasedUnknown 记录对未持有 requestID 的重复/无效释放(幂等释放语义) + accountReleasedUnknown int + userAcquired int + userReleased int + userReleasedUnknown int + + userWaitInc int + userWaitDec int + accountWaitInc int + accountWaitDec int +} + +var _ service.ConcurrencyCache = (*schedInvCountingCache)(nil) + +func schedInvNewCountingCache() *schedInvCountingCache { + return &schedInvCountingCache{ + accountHeld: make(map[int64]map[string]struct{}), + userHeld: make(map[int64]map[string]struct{}), + } +} + +func (c *schedInvCountingCache) AcquireAccountSlot(_ context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + held := c.accountHeld[accountID] + if maxConcurrency > 0 && len(held) >= maxConcurrency { + return false, nil + } + if held == nil { + held = make(map[string]struct{}) + c.accountHeld[accountID] = held + } + held[requestID] = struct{}{} + c.accountAcquired++ + return true, nil +} + +func (c *schedInvCountingCache) ReleaseAccountSlot(_ context.Context, accountID int64, requestID string) error { + c.mu.Lock() + defer c.mu.Unlock() + if held := c.accountHeld[accountID]; held != nil { + if _, ok := held[requestID]; ok { + delete(held, requestID) + c.accountReleased++ + return nil + } + } + c.accountReleasedUnknown++ + return nil +} + +func (c *schedInvCountingCache) GetAccountConcurrency(_ context.Context, accountID int64) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.accountHeld[accountID]), nil +} + +func (c *schedInvCountingCache) GetAccountConcurrencyBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) { + c.mu.Lock() + defer c.mu.Unlock() + result := make(map[int64]int, len(accountIDs)) + for _, id := range accountIDs { + result[id] = len(c.accountHeld[id]) + } + return result, nil +} + +func (c *schedInvCountingCache) IncrementAccountWaitCount(_ context.Context, _ int64, _ int) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.accountWaitInc++ + return true, nil +} + +func (c *schedInvCountingCache) DecrementAccountWaitCount(_ context.Context, _ int64) error { + c.mu.Lock() + defer c.mu.Unlock() + c.accountWaitDec++ + return nil +} + +func (c *schedInvCountingCache) GetAccountWaitingCount(_ context.Context, _ int64) (int, error) { + return 0, nil +} + +func (c *schedInvCountingCache) AcquireUserSlot(_ context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + held := c.userHeld[userID] + if maxConcurrency > 0 && len(held) >= maxConcurrency { + return false, nil + } + if held == nil { + held = make(map[string]struct{}) + c.userHeld[userID] = held + } + held[requestID] = struct{}{} + c.userAcquired++ + return true, nil +} + +func (c *schedInvCountingCache) ReleaseUserSlot(_ context.Context, userID int64, requestID string) error { + c.mu.Lock() + defer c.mu.Unlock() + if held := c.userHeld[userID]; held != nil { + if _, ok := held[requestID]; ok { + delete(held, requestID) + c.userReleased++ + return nil + } + } + c.userReleasedUnknown++ + return nil +} + +func (c *schedInvCountingCache) GetUserConcurrency(_ context.Context, userID int64) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.userHeld[userID]), nil +} + +func (c *schedInvCountingCache) IncrementWaitCount(_ context.Context, _ int64, _ int) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.userWaitInc++ + return true, nil +} + +func (c *schedInvCountingCache) DecrementWaitCount(_ context.Context, _ int64) error { + c.mu.Lock() + defer c.mu.Unlock() + c.userWaitDec++ + return nil +} + +func (c *schedInvCountingCache) GetAccountsLoadBatch(_ context.Context, accounts []service.AccountWithConcurrency) (map[int64]*service.AccountLoadInfo, error) { + c.mu.Lock() + defer c.mu.Unlock() + result := make(map[int64]*service.AccountLoadInfo, len(accounts)) + for _, acc := range accounts { + held := len(c.accountHeld[acc.ID]) + loadRate := 0 + if acc.MaxConcurrency > 0 { + loadRate = held * 100 / acc.MaxConcurrency + } + result[acc.ID] = &service.AccountLoadInfo{ + AccountID: acc.ID, + CurrentConcurrency: held, + LoadRate: loadRate, + } + } + return result, nil +} + +func (c *schedInvCountingCache) GetUsersLoadBatch(_ context.Context, users []service.UserWithConcurrency) (map[int64]*service.UserLoadInfo, error) { + result := make(map[int64]*service.UserLoadInfo, len(users)) + for _, u := range users { + result[u.ID] = &service.UserLoadInfo{UserID: u.ID} + } + return result, nil +} + +func (c *schedInvCountingCache) CleanupExpiredAccountSlots(_ context.Context, _ int64) error { return nil } +func (c *schedInvCountingCache) CleanupStaleProcessSlots(_ context.Context, _ string) error { return nil } + +// schedInvRequireBalanced 断言所有槽位与等待计数获取-释放严格配平。 +func schedInvRequireBalanced(t *testing.T, cc *schedInvCountingCache) { + t.Helper() + cc.mu.Lock() + defer cc.mu.Unlock() + for accountID, held := range cc.accountHeld { + require.Empty(t, held, "账号 %d 仍有未释放的槽位", accountID) + } + for userID, held := range cc.userHeld { + require.Empty(t, held, "用户 %d 仍有未释放的槽位", userID) + } + require.Equal(t, cc.accountAcquired, cc.accountReleased, "账号槽位获取/释放必须配平") + require.Equal(t, cc.userAcquired, cc.userReleased, "用户槽位获取/释放必须配平") + require.Equal(t, cc.userWaitInc, cc.userWaitDec, "用户等待计数增减必须配平") + require.Equal(t, cc.accountWaitInc, cc.accountWaitDec, "账号等待计数增减必须配平") +} + +// schedInvGroup 构造测试分组。 +func schedInvGroup(groupID int64) *service.Group { + return &service.Group{ + ID: groupID, + Hydrated: true, + Platform: service.PlatformAnthropic, + Status: service.StatusActive, + } +} + +// schedInvPassthroughAccount 构造 anthropic API Key 透传账号(Forward 路径不依赖真实上游域名)。 +func schedInvPassthroughAccount(id, groupID int64, extraCreds map[string]any) *service.Account { + creds := map[string]any{"api_key": fmt.Sprintf("sk-sched-inv-%d", id)} + for k, v := range extraCreds { + creds[k] = v + } + return &service.Account{ + ID: id, + Name: fmt.Sprintf("sched-inv-acc-%d", id), + Platform: service.PlatformAnthropic, + Type: service.AccountTypeAPIKey, + Credentials: creds, + Extra: map[string]any{"anthropic_passthrough": true}, + Concurrency: 5, + Priority: 1, + Status: service.StatusActive, + Schedulable: true, + AccountGroups: []service.AccountGroup{{AccountID: id, GroupID: groupID}}, + } +} + +// schedInvNewHandler 构造完整 GatewayHandler:真实 failover 循环 + 计数并发缓存 + 上游桩。 +// cfg 传 nil 给 NewGatewayHandler,以固化默认换号上限(anthropic=10 / gemini=3)。 +func schedInvNewHandler(t *testing.T, group *service.Group, accounts []*service.Account, upstream service.HTTPUpstream, cc *schedInvCountingCache) (*GatewayHandler, func()) { + t.Helper() + + // 隔离环境变量:避免本机 SUB2API_DEBUG_GATEWAY_BODY 在测试期间写调试文件。 + t.Setenv("SUB2API_DEBUG_GATEWAY_BODY", "") + + schedulerSnapshot := service.NewSchedulerSnapshotService(&fakeSchedulerCache{accounts: accounts}, nil, nil, nil, nil) + concurrencySvc := service.NewConcurrencyService(cc) + + gwSvc := service.NewGatewayService( + nil, // accountRepo(scheduler snapshot 命中,不需要) + &fakeGroupRepo{group: group}, + nil, nil, nil, nil, nil, // usageLogRepo / usageBillingRepo / userRepo / userSubRepo / userGroupRateRepo + nil, // cache(粘性会话关闭) + // 空 cfg(非 nil):API Key 透传分支会校验默认 base_url, + // validateUpstreamBaseURL 在 cfg=nil 时会空指针。 + &config.Config{}, + schedulerSnapshot, + concurrencySvc, + nil, // billingService + &service.RateLimitService{}, // rateLimitService(零值:500 仅记录日志,无副作用) + nil, // billingCacheService + nil, // identityService + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + ) + + // RunModeSimple 跳过计费检查,避免引入 repo/cache 依赖。 + billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, + &config.Config{RunMode: config.RunModeSimple}, nil) + + h := NewGatewayHandler( + gwSvc, + nil, // geminiCompatService + nil, // antigravityGatewayService + nil, // userService + concurrencySvc, + billingCacheSvc, + nil, // usageService + nil, // apiKeyService + nil, // usageRecordWorkerPool + nil, // errorPassthroughService + nil, // preFlightHooks + nil, // userMsgQueueService + nil, // cfg → 默认换号上限 + nil, // settingService + ) + return h, func() { billingCacheSvc.Stop() } +} + +// schedInvNewMessagesContext 构造带认证上下文的 /v1/messages 请求。 +// 返回的 cancel 用于模拟请求结束时的 context 取消(生产环境由 net/http 完成)。 +func schedInvNewMessagesContext(t *testing.T, group *service.Group, body []byte) (*gin.Context, *httptest.ResponseRecorder, context.CancelFunc) { + t.Helper() + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + req := httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + ctx := context.WithValue(req.Context(), ctxkey.Group, group) + ctx, cancel := context.WithCancel(ctx) + c.Request = req.WithContext(ctx) + + apiKey := &service.APIKey{ + ID: 8301, + UserID: 8401, + GroupID: &group.ID, + Status: service.StatusActive, + User: &service.User{ + ID: 8401, + Concurrency: 10, + Balance: 100, + }, + Group: group, + } + c.Set(string(middleware.ContextKeyAPIKey), apiKey) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.UserID, Concurrency: 10}) + return c, rec, cancel +} + +func schedInvMessagesBody() []byte { + return []byte(`{ + "model": "claude-sonnet-4-5", + "max_tokens": 256, + "messages": [{"role":"user","content":[{"type":"text","text":"scheduling invariants probe"}]}] + }`) +} + +// --------------------------------------------------------------------------- +// I-5.1 三平台换号上限默认值(双断言) +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_FailoverSwitchLimit_DefaultValues(t *testing.T) { + t.Run("anthropic与gemini默认上限", func(t *testing.T) { + h := NewGatewayHandler(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + require.Equal(t, 10, h.maxAccountSwitches, "anthropic 默认换号上限必须为 10") + require.Equal(t, 3, h.maxAccountSwitchesGemini, "gemini 默认换号上限必须为 3") + }) + + t.Run("openai默认上限", func(t *testing.T) { + oh := NewOpenAIGatewayHandler(nil, nil, nil, nil, nil, nil, nil, nil, nil) + require.Equal(t, 3, oh.maxAccountSwitches, "openai 默认换号上限必须为 3") + }) + + t.Run("配置可覆盖默认上限", func(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.MaxAccountSwitches = 5 + cfg.Gateway.MaxAccountSwitchesGemini = 2 + h := NewGatewayHandler(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, cfg, nil) + require.Equal(t, 5, h.maxAccountSwitches) + require.Equal(t, 2, h.maxAccountSwitchesGemini) + }) +} + +// --------------------------------------------------------------------------- +// I-5.1 + I-5.2 + I-6.1:anthropic 完整 failover 循环 +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_FailoverAnthropic_FullLoopExhaustion 通过完整 Messages +// 入口驱动真实 failover 循环:12 个账号 + 恒定上游 500。 +// 固化语义: +// 1. 恰好尝试 maxAccountSwitches+1 = 11 个账号(初次 + 10 次换号); +// 2. 失败账号进入排除集合,同一账号不被尝试两次; +// 3. 耗尽后客户端收到 502 + {"type":"error","error":{"type":"upstream_error",...}} +// (/v1/messages 路径按 mapUpstreamError 映射 500→502; +// "server_error/All available accounts exhausted" 错误体属于 +// chat-completions/responses 兼容路径,见下方独立用例); +// 4. 转发失败路径下账号/用户槽位获取-释放严格配平。 +func TestSchedulingInvariant_FailoverAnthropic_FullLoopExhaustion(t *testing.T) { + groupID := int64(9001) + group := schedInvGroup(groupID) + + accounts := make([]*service.Account, 0, 12) + for i := int64(1); i <= 12; i++ { + accounts = append(accounts, schedInvPassthroughAccount(9100+i, groupID, nil)) + } + + upstream := &schedInvUpstream{ + status: http.StatusInternalServerError, + body: `{"type":"error","error":{"type":"api_error","message":"schedInv upstream boom"}}`, + } + cc := schedInvNewCountingCache() + h, cleanup := schedInvNewHandler(t, group, accounts, upstream, cc) + defer cleanup() + + c, rec, cancel := schedInvNewMessagesContext(t, group, schedInvMessagesBody()) + defer cancel() + + h.Messages(c) + + // 1. 换号次数恰为上限:尝试账号数 = maxAccountSwitches + 1(常量引用 + 硬编码双断言) + attempts := upstream.attemptedAccounts() + require.Len(t, attempts, h.maxAccountSwitches+1, "尝试账号数必须等于 maxAccountSwitches+1") + require.Len(t, attempts, 11, "anthropic 平台:1 次初始尝试 + 10 次换号") + + // 2. 排除集合语义:同一账号不被重选 + seen := make(map[int64]struct{}, len(attempts)) + for _, accountID := range attempts { + _, dup := seen[accountID] + require.False(t, dup, "账号 %d 被尝试了两次:失败账号必须进入排除集合", accountID) + seen[accountID] = struct{}{} + } + + // 3. 耗尽后错误语义:502 + upstream_error 映射错误体(固化当前实际行为) + require.Equal(t, http.StatusBadGateway, rec.Code) + require.JSONEq(t, + `{"type":"error","error":{"type":"upstream_error","message":"Upstream service temporarily unavailable"}}`, + rec.Body.String()) + + // 4. I-6.1 转发失败路径:槽位严格配平 + require.Equal(t, 11, cc.accountAcquired, "每次尝试获取一次账号槽位") + require.Equal(t, 1, cc.userAcquired, "整个请求只获取一次用户槽位") + schedInvRequireBalanced(t, cc) +} + +// --------------------------------------------------------------------------- +// I-5.3 池模式可重试错误:同账号重试 3 次后换号(完整链) +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_FailoverSameAccountRetry_FullChain 使用池模式账号 +// (pool_mode_retry_status_codes=[500])驱动完整链路: +// 上游恒定 500 → 同账号共尝试 1+maxSameAccountRetries=4 次 → 加入排除集合换号 → +// 无其他账号 → 502。同时验证每次重试均独立获取/释放账号槽位(配平)。 +func TestSchedulingInvariant_FailoverSameAccountRetry_FullChain(t *testing.T) { + groupID := int64(9002) + group := schedInvGroup(groupID) + + poolAccount := schedInvPassthroughAccount(9201, groupID, map[string]any{ + "pool_mode": true, + "pool_mode_retry_status_codes": []any{float64(500)}, + }) + + upstream := &schedInvUpstream{ + status: http.StatusInternalServerError, + body: `{"type":"error","error":{"type":"api_error","message":"schedInv pool boom"}}`, + } + cc := schedInvNewCountingCache() + h, cleanup := schedInvNewHandler(t, group, []*service.Account{poolAccount}, upstream, cc) + defer cleanup() + + c, rec, cancel := schedInvNewMessagesContext(t, group, schedInvMessagesBody()) + defer cancel() + + h.Messages(c) + + // 同账号共尝试 1 + maxSameAccountRetries 次(常量引用 + 硬编码双断言) + attempts := upstream.attemptedAccounts() + require.Len(t, attempts, 1+maxSameAccountRetries, "可重试错误:1 次初始 + maxSameAccountRetries 次同账号重试") + require.Len(t, attempts, 4) + for _, accountID := range attempts { + require.Equal(t, poolAccount.ID, accountID, "重试必须发生在同一账号上") + } + + // 重试耗尽 + 换号后无可用账号 → 502 + require.Equal(t, http.StatusBadGateway, rec.Code) + require.JSONEq(t, + `{"type":"error","error":{"type":"upstream_error","message":"Upstream service temporarily unavailable"}}`, + rec.Body.String()) + + // 每轮重试独立获取/释放账号槽位 + require.Equal(t, 4, cc.accountAcquired) + schedInvRequireBalanced(t, cc) +} + +// --------------------------------------------------------------------------- +// I-5.1 gemini 换号上限:FailoverState 循环契约 +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_FailoverGemini_SwitchLimitLoopContract 按 Messages +// gemini 分支的真实接线(NewFailoverState(h.maxAccountSwitchesGemini, ...)) +// 驱动 FailoverState:连续上游失败时恰好允许 3+1=4 个账号尝试后耗尽。 +// (gemini 平台 Forward 的服务层内部 500 重试带秒级退避,完整 e2e 不可在 +// 单测时间预算内执行,故此处固化 handler 循环契约层语义。) +func TestSchedulingInvariant_FailoverGemini_SwitchLimitLoopContract(t *testing.T) { + h := NewGatewayHandler(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + require.Equal(t, 3, h.maxAccountSwitchesGemini) + + mock := &mockTempUnscheduler{} + fs := NewFailoverState(h.maxAccountSwitchesGemini, false) + + attempts := 0 + for i := 1; ; i++ { + require.LessOrEqual(t, i, 10, "防御:循环不应超过 10 次") + attempts++ + action := fs.HandleFailoverError(context.Background(), mock, int64(i), service.PlatformGemini, newTestFailoverErr(500, false, false)) + if action == FailoverExhausted { + break + } + require.Equal(t, FailoverContinue, action) + } + + require.Equal(t, h.maxAccountSwitchesGemini+1, attempts, "gemini:尝试账号数 = 上限 + 1") + require.Equal(t, 4, attempts) + require.Len(t, fs.FailedAccountIDs, 4, "所有失败账号都进入排除集合") +} + +// --------------------------------------------------------------------------- +// chat-completions 兼容路径耗尽错误体 +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_FailoverExhausted_ChatCompletionsErrorBody 固化 +// chat-completions 兼容路径的耗尽错误体(INVARIANTS 记录的 +// "All available accounts exhausted" 属于此路径而非 /v1/messages 路径)。 +// 注意当前实际行为:lastErr 非空时状态码透传上游状态码(500→500), +// lastErr 为空时才回退 502。 +func TestSchedulingInvariant_FailoverExhausted_ChatCompletionsErrorBody(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("lastErr存在_状态码透传上游", func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + h := &GatewayHandler{} + h.handleCCFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusInternalServerError, + ResponseBody: []byte(`{"error":{"message":"boom"}}`), + }, false) + + require.Equal(t, http.StatusInternalServerError, rec.Code, + "当前行为:lastErr 存在时透传上游状态码(非固定 502)") + require.JSONEq(t, + `{"error":{"type":"server_error","message":"All available accounts exhausted"}}`, + rec.Body.String()) + }) + + t.Run("lastErr为空_回退502", func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + h := &GatewayHandler{} + h.handleCCFailoverExhausted(c, nil, false) + + require.Equal(t, http.StatusBadGateway, rec.Code) + require.JSONEq(t, + `{"error":{"type":"server_error","message":"All available accounts exhausted"}}`, + rec.Body.String()) + }) +} diff --git a/backend/internal/handler/scheduling_invariants_slots_test.go b/backend/internal/handler/scheduling_invariants_slots_test.go new file mode 100644 index 0000000000..7044309e94 --- /dev/null +++ b/backend/internal/handler/scheduling_invariants_slots_test.go @@ -0,0 +1,289 @@ +//go:build unit + +// Phase-0 TASK-004 并发槽不变量测试(INVARIANTS I-6.1 / I-6.2)。 +// +// 固化内容: +// - I-6.1 用户槽/账号槽获取-释放严格配平: +// * 正常路径(上游 200,完整 Messages 链路); +// * early-return 路径(预热拦截:选号后转发前直接返回); +// * panic 路径(defer 释放用户槽 + context 取消兜底回收账号槽); +// * wrapReleaseOnDone 恰好一次语义(显式调用 / context 取消两种触发); +// * service 层 AcquireResult.ReleaseFunc 重复调用的当前行为特征化; +// - I-6.2 AcquireUserSlotWithWait 等待队列:满 → 等待 → 释放 → 唤醒; +// 等待超时返回 ConcurrencyError{IsTimeout}。 +// +// 转发失败路径的配平已由 scheduling_invariants_failover_test.go 的完整循环用例覆盖。 +// 复用本包 schedInv* 夹具(scheduling_invariants_failover_test.go)。 +package handler + +import ( + "context" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// schedInvNewGinContext 构造仅含基础请求的 gin 上下文(供 ConcurrencyHelper 使用)。 +func schedInvNewGinContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + return c, rec +} + +// --------------------------------------------------------------------------- +// I-6.2 用户槽等待队列:满 → 等待 → 释放 → 唤醒 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_UserSlotWaitQueue_FullWaitReleaseWakeup(t *testing.T) { + cc := schedInvNewCountingCache() + helper := NewConcurrencyHelper(service.NewConcurrencyService(cc), SSEPingFormatNone, time.Hour) + + const userID = int64(701) + + // 1. 占满唯一槽位 + c1, _ := schedInvNewGinContext(t) + var ss1 bool + release1, err := helper.AcquireUserSlotWithWait(c1, userID, 1, false, &ss1) + require.NoError(t, err) + require.NotNil(t, release1) + + // 2. 第二个请求进入等待(槽位已满) + type waitResult struct { + release func() + err error + } + done := make(chan waitResult, 1) + go func() { + c2, _ := schedInvNewGinContext(t) + var ss2 bool + release2, err2 := helper.AcquireUserSlotWithWait(c2, userID, 1, false, &ss2) + done <- waitResult{release: release2, err: err2} + }() + + // 槽位未释放期间,等待者不应返回 + select { + case r := <-done: + t.Fatalf("槽位未释放时等待者不应获取成功: err=%v", r.err) + case <-time.After(400 * time.Millisecond): + } + + // 3. 释放 → 4. 等待者被唤醒并成功获取 + release1() + select { + case r := <-done: + require.NoError(t, r.err, "释放后等待者必须能获取槽位") + require.NotNil(t, r.release) + r.release() + case <-time.After(10 * time.Second): + t.Fatal("释放槽位后等待者在 10s 内未被唤醒") + } + + require.Equal(t, 2, cc.userAcquired, "两次成功获取") + schedInvRequireBalanced(t, cc) +} + +// TestSchedulingInvariant_AccountSlotWait_TimeoutReturnsConcurrencyError 固化: +// 账号槽等待超时后返回 ConcurrencyError{IsTimeout:true},且不产生未配平的获取。 +func TestSchedulingInvariant_AccountSlotWait_TimeoutReturnsConcurrencyError(t *testing.T) { + cc := schedInvNewCountingCache() + svc := service.NewConcurrencyService(cc) + helper := NewConcurrencyHelper(svc, SSEPingFormatNone, time.Hour) + + const accountID = int64(801) + + // 占满唯一槽位 + holder, err := svc.AcquireAccountSlot(context.Background(), accountID, 1) + require.NoError(t, err) + require.True(t, holder.Acquired) + + c, _ := schedInvNewGinContext(t) + var ss bool + release, err := helper.AcquireAccountSlotWithWaitTimeout(c, accountID, 1, 300*time.Millisecond, false, &ss) + require.Nil(t, release) + require.Error(t, err) + var concurrencyErr *ConcurrencyError + require.ErrorAs(t, err, &concurrencyErr) + require.True(t, concurrencyErr.IsTimeout, "等待超时必须返回 IsTimeout 的 ConcurrencyError") + require.Equal(t, "account", concurrencyErr.SlotType) + + holder.ReleaseFunc() + schedInvRequireBalanced(t, cc) +} + +// --------------------------------------------------------------------------- +// I-6.1 正常路径与 early-return 路径配平(完整 Messages 链路) +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_SlotBalance_NormalSuccessPath 固化:上游 200 正常完成时, +// 账号槽与用户槽各获取/释放一次,严格配平。 +func TestSchedulingInvariant_SlotBalance_NormalSuccessPath(t *testing.T) { + groupID := int64(9003) + group := schedInvGroup(groupID) + account := schedInvPassthroughAccount(9301, groupID, nil) + + upstream := &schedInvUpstream{ + status: http.StatusOK, + body: `{"id":"msg_sched_inv","type":"message","role":"assistant","model":"claude-sonnet-4-5",` + + `"content":[{"type":"text","text":"ok"}],"stop_reason":"end_turn",` + + `"usage":{"input_tokens":3,"output_tokens":2}}`, + } + cc := schedInvNewCountingCache() + h, cleanup := schedInvNewHandler(t, group, []*service.Account{account}, upstream, cc) + defer cleanup() + + c, rec, cancel := schedInvNewMessagesContext(t, group, schedInvMessagesBody()) + defer cancel() + + h.Messages(c) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), "msg_sched_inv", "上游 200 响应体应透传给客户端") + require.Len(t, upstream.attemptedAccounts(), 1) + require.Equal(t, 1, cc.accountAcquired) + require.Equal(t, 1, cc.userAcquired) + schedInvRequireBalanced(t, cc) +} + +// TestSchedulingInvariant_SlotBalance_InterceptEarlyReturnPath 固化 early-return 路径: +// 预热拦截在选号成功后、转发上游前直接返回 mock 响应,账号槽必须被释放。 +func TestSchedulingInvariant_SlotBalance_InterceptEarlyReturnPath(t *testing.T) { + groupID := int64(9004) + group := schedInvGroup(groupID) + account := schedInvPassthroughAccount(9401, groupID, map[string]any{ + "intercept_warmup_requests": true, + }) + + upstream := &schedInvUpstream{status: http.StatusOK, body: `{}`} + cc := schedInvNewCountingCache() + h, cleanup := schedInvNewHandler(t, group, []*service.Account{account}, upstream, cc) + defer cleanup() + + warmupBody := []byte(`{ + "model": "claude-sonnet-4-5", + "max_tokens": 256, + "messages": [{"role":"user","content":[{"type":"text","text":"Warmup"}]}] + }`) + c, rec, cancel := schedInvNewMessagesContext(t, group, warmupBody) + defer cancel() + + h.Messages(c) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), "msg_mock_warmup", "预热请求应被拦截返回 mock 响应") + require.Empty(t, upstream.attemptedAccounts(), "early-return 路径不应触达上游") + require.Equal(t, 1, cc.accountAcquired, "选号阶段获取过账号槽") + schedInvRequireBalanced(t, cc) +} + +// --------------------------------------------------------------------------- +// I-6.1 panic 路径:defer 释放用户槽 + context 取消兜底回收账号槽 +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_SlotBalance_PanicPath 固化当前 panic 语义: +// - 用户槽通过 defer userReleaseFunc() 在 panic 展开时立即释放; +// - 账号槽的释放调用位于 Forward 之后(非 defer),panic 时依赖 +// wrapReleaseOnDone 注册的 context.AfterFunc 在请求 context 取消时兜底回收 +// (生产环境由 net/http 在连接结束时取消请求 context)。 +func TestSchedulingInvariant_SlotBalance_PanicPath(t *testing.T) { + groupID := int64(9005) + group := schedInvGroup(groupID) + account := schedInvPassthroughAccount(9501, groupID, nil) + + upstream := &schedInvUpstream{status: http.StatusOK, body: `{}`, panicOn: true} + cc := schedInvNewCountingCache() + h, cleanup := schedInvNewHandler(t, group, []*service.Account{account}, upstream, cc) + defer cleanup() + + c, _, cancel := schedInvNewMessagesContext(t, group, schedInvMessagesBody()) + defer cancel() + + func() { + defer func() { + require.NotNil(t, recover(), "上游 panic 应穿透 Messages(由外层 gin recovery 兜底)") + }() + h.Messages(c) + }() + + // 用户槽:defer 在 panic 展开时已释放 + cc.mu.Lock() + userReleased := cc.userReleased + cc.mu.Unlock() + require.Equal(t, 1, userReleased, "panic 展开时 defer 必须释放用户槽") + + // 账号槽:panic 跳过了显式释放调用,由请求 context 取消兜底回收 + cancel() + require.Eventually(t, func() bool { + cc.mu.Lock() + defer cc.mu.Unlock() + return cc.accountReleased == cc.accountAcquired + }, 5*time.Second, 10*time.Millisecond, "context 取消后账号槽必须被兜底回收") + schedInvRequireBalanced(t, cc) +} + +// --------------------------------------------------------------------------- +// I-6.1 wrapReleaseOnDone 恰好一次语义 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_WrapReleaseOnDone_ExactlyOnce(t *testing.T) { + t.Run("重复显式调用只释放一次", func(t *testing.T) { + var calls atomic.Int32 + release := wrapReleaseOnDone(context.Background(), func() { calls.Add(1) }) + release() + release() + require.Equal(t, int32(1), calls.Load(), "重复调用不得重复释放") + }) + + t.Run("context取消自动触发释放", func(t *testing.T) { + var calls atomic.Int32 + ctx, cancel := context.WithCancel(context.Background()) + release := wrapReleaseOnDone(ctx, func() { calls.Add(1) }) + cancel() + require.Eventually(t, func() bool { return calls.Load() == 1 }, + 2*time.Second, 5*time.Millisecond, "context 取消必须自动触发释放") + // 取消后再显式调用也不重复释放 + release() + require.Equal(t, int32(1), calls.Load()) + }) + + t.Run("nil释放函数返回nil", func(t *testing.T) { + require.Nil(t, wrapReleaseOnDone(context.Background(), nil)) + }) +} + +// --------------------------------------------------------------------------- +// service 层 ReleaseFunc 重复调用特征化 +// --------------------------------------------------------------------------- + +// TestSchedulingCharacterization_ServiceReleaseFuncNotOnceGuarded 固化当前实际行为: +// ConcurrencyService 返回的 AcquireResult.ReleaseFunc 本身没有 once 保护, +// 重复调用会向缓存重复发出释放请求;配平依赖两点: +// 1. 缓存释放按 requestID 幂等(Redis ZREM 不存在的成员是 no-op); +// 2. handler 层统一经 wrapReleaseOnDone 包装后才暴露。 +func TestSchedulingCharacterization_ServiceReleaseFuncNotOnceGuarded(t *testing.T) { + cc := schedInvNewCountingCache() + svc := service.NewConcurrencyService(cc) + + result, err := svc.AcquireAccountSlot(context.Background(), 901, 5) + require.NoError(t, err) + require.True(t, result.Acquired) + + result.ReleaseFunc() + result.ReleaseFunc() + + cc.mu.Lock() + defer cc.mu.Unlock() + require.Equal(t, 1, cc.accountReleased, "第一次释放生效") + require.Equal(t, 1, cc.accountReleasedUnknown, + "当前行为:第二次释放仍会调用缓存(按 requestID 幂等,不破坏配平)") + require.Empty(t, cc.accountHeld[901]) +} diff --git a/backend/internal/service/antigravity_gemini_characterization_test.go b/backend/internal/service/antigravity_gemini_characterization_test.go new file mode 100644 index 0000000000..7bca8cbef2 --- /dev/null +++ b/backend/internal/service/antigravity_gemini_characterization_test.go @@ -0,0 +1,188 @@ +//go:build unit + +// Phase-0 TASK-002 特征化测试:gemini/antigravity 路径流式/非流式透传(INVARIANTS I-1.6)。 +// ForwardGemini 的外部可观测语义: +// - 流式:上游 v1internal SSE(data: {"response":{...}})被解包后逐事件转发给客户端; +// - 非流式:上游流式响应被收集合并为单个 JSON(文本片段拼接)后返回; +// - 非 failover 上游错误(如 404):解包后的错误体 + 上游状态码原样返回。 +package service + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// passCharAntigravityService 构造最小可运行的 AntigravityGatewayService。 +func passCharAntigravityService(upstream HTTPUpstream) *AntigravityGatewayService { + return &AntigravityGatewayService{ + settingService: NewSettingService(&antigravitySettingRepoStub{}, &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}), + tokenProvider: &AntigravityTokenProvider{}, + httpUpstream: upstream, + } +} + +// passCharAntigravityAccount 返回带模型映射的 antigravity OAuth 账号夹具。 +func passCharAntigravityAccount(id int64, mapping map[string]any) *Account { + return &Account{ + ID: id, + Name: "pass-char-antigravity", + Platform: PlatformAntigravity, + Type: AccountTypeOAuth, + Status: StatusActive, + Concurrency: 1, + Schedulable: true, + Credentials: map[string]any{ + "access_token": "ag-token", + "model_mapping": mapping, + }, + } +} + +func passCharGeminiBody(t *testing.T) []byte { + t.Helper() + body, err := json.Marshal(map[string]any{ + "contents": []map[string]any{ + {"role": "user", "parts": []map[string]any{{"text": "hello"}}}, + }, + }) + require.NoError(t, err) + return body +} + +// TestGatewayCharacterization_GeminiStreamUnwrapsV1Internal 固化 I-1.6 流式侧: +// 上游 data 行的 v1internal 包裹({"response":{...}})被解包,客户端按事件收到内层 JSON 原文。 +func TestGatewayCharacterization_GeminiStreamUnwrapsV1Internal(t *testing.T) { + gin.SetMode(gin.TestMode) + + innerChunk1 := `{"candidates":[{"content":{"role":"model","parts":[{"text":"Hel"}]}}]}` + innerChunk2 := `{"candidates":[{"content":{"role":"model","parts":[{"text":"lo ✓"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3,"cachedContentTokenCount":2}}` + upstreamSSE := "data: {\"response\":" + innerChunk1 + "}\n\n" + + "data: {\"response\":" + innerChunk2 + "}\n\n" + + upstream := &httpUpstreamStub{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"rid-char-gemini-stream"}, + }, + Body: io.NopCloser(bytes.NewReader([]byte(upstreamSSE))), + }} + svc := passCharAntigravityService(upstream) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-2.5-flash:streamGenerateContent", bytes.NewReader(passCharGeminiBody(t))) + + account := passCharAntigravityAccount(9201, map[string]any{"gemini-2.5-flash": "gemini-3-pro-high"}) + + result, err := svc.ForwardGemini(context.Background(), c, account, "gemini-2.5-flash", "streamGenerateContent", true, passCharGeminiBody(t), false) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Stream) + + wantEvents := []string{ + "data: " + innerChunk1, + "data: " + innerChunk2, + } + require.Equal(t, wantEvents, passCharSplitSSEEvents(rec.Body.String()), + "客户端应按事件收到解包后的内层 JSON 原文") + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-gemini-stream", rec.Header().Get("x-request-id")) + + require.Equal(t, "gemini-2.5-flash", result.Model) + require.Equal(t, "gemini-3-pro-high", result.UpstreamModel, "计费模型使用映射后的模型") + require.Equal(t, 6, result.Usage.InputTokens, "input = promptTokenCount - cachedContentTokenCount") + require.Equal(t, 3, result.Usage.OutputTokens) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) +} + +// TestGatewayCharacterization_GeminiNonStreamCollectsStream 固化 I-1.6 非流式侧: +// 客户端请求非流式时,网关收集上游流式 chunk,将文本片段拼接进最后一个含 parts 的 +// 响应中,作为单个 JSON(HTTP 200, application/json)返回。 +func TestGatewayCharacterization_GeminiNonStreamCollectsStream(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstreamSSE := `data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"Hel"}]}}]}}` + "\n\n" + + `data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"text":"lo ✓"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3,"cachedContentTokenCount":2}}}` + "\n\n" + + upstream := &httpUpstreamStub{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"rid-char-gemini-collect"}, + }, + Body: io.NopCloser(bytes.NewReader([]byte(upstreamSSE))), + }} + svc := passCharAntigravityService(upstream) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-2.5-flash:generateContent", bytes.NewReader(passCharGeminiBody(t))) + + account := passCharAntigravityAccount(9202, map[string]any{"gemini-2.5-flash": "gemini-3-pro-high"}) + + result, err := svc.ForwardGemini(context.Background(), c, account, "gemini-2.5-flash", "generateContent", false, passCharGeminiBody(t), false) + require.NoError(t, err) + require.NotNil(t, result) + require.False(t, result.Stream) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "application/json", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-gemini-collect", rec.Header().Get("x-request-id")) + + // 合并语义:以最后一个含 parts 的 chunk 为基底,第一个 text part 替换为全部文本片段的拼接。 + require.JSONEq(t, `{ + "candidates":[{ + "content":{"role":"model","parts":[{"text":"Hello ✓"}]}, + "finishReason":"STOP" + }], + "usageMetadata":{"promptTokenCount":8,"candidatesTokenCount":3,"cachedContentTokenCount":2} + }`, rec.Body.String()) + + require.Equal(t, 6, result.Usage.InputTokens) + require.Equal(t, 3, result.Usage.OutputTokens) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) +} + +// TestGatewayCharacterization_GeminiUpstreamErrorUnwrappedPassthrough 固化 I-1.6 错误侧: +// 非 failover 上游错误(404)时,客户端收到上游状态码 + 解包后的错误体原文。 +func TestGatewayCharacterization_GeminiUpstreamErrorUnwrappedPassthrough(t *testing.T) { + gin.SetMode(gin.TestMode) + + innerErr := `{"error":{"code":404,"message":"model not found: gemini-3-pro-high","status":"NOT_FOUND"}}` + upstream := &httpUpstreamStub{resp: &http.Response{ + StatusCode: http.StatusNotFound, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "X-Request-Id": []string{"rid-char-gemini-err"}, + }, + Body: io.NopCloser(bytes.NewReader([]byte(`{"response":` + innerErr + `}`))), + }} + svc := passCharAntigravityService(upstream) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-2.5-flash:generateContent", bytes.NewReader(passCharGeminiBody(t))) + + account := passCharAntigravityAccount(9203, map[string]any{"gemini-2.5-flash": "gemini-3-pro-high"}) + + result, err := svc.ForwardGemini(context.Background(), c, account, "gemini-2.5-flash", "generateContent", false, passCharGeminiBody(t), false) + require.Error(t, err) + require.Nil(t, result) + + require.Equal(t, http.StatusNotFound, rec.Code, "非 failover 上游错误状态码原样返回") + require.Equal(t, innerErr, rec.Body.String(), "错误体应为解包后的内层 JSON 原文") + require.Equal(t, "application/json", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-gemini-err", rec.Header().Get("x-request-id")) +} diff --git a/backend/internal/service/billing_invariants_cost_test.go b/backend/internal/service/billing_invariants_cost_test.go new file mode 100644 index 0000000000..fcd3107e94 --- /dev/null +++ b/backend/internal/service/billing_invariants_cost_test.go @@ -0,0 +1,291 @@ +//go:build unit + +// TASK-003 计费精度不变量测试(INVARIANTS.md ② 计费精度)。 +// +// 本文件覆盖纯计算层(BillingService)的不变量: +// - I-2.2: 5m 与 1h 缓存写入差异化计价(两档金额断言 + 无明细回退 + breakdown 开关门控) +// - I-2.4: 倍率叠加顺序 serviceTier → rateMultiplier;rateMultiplier=0 时 +// ActualCost=0 但 TotalCost 保留 +// +// 所有期望值均为人工核算的固定值(算式在行内注释),容差 1e-10。 +// 这些是 characterization 测试:固化当前 main 上的实际行为,插件化改造期间 +// 任何金额漂移都应使其失败。 +package service + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +// billInvManualPricingService 构造带手工定价表的 BillingService(无动态价格源)。 +func billInvManualPricingService(pricing *ModelPricing) *BillingService { + return &BillingService{ + cfg: &config.Config{}, + fallbackPrices: map[string]*ModelPricing{ + "claude-sonnet-4": pricing, + }, + } +} + +// TestBillingInvariant_CacheTier5mVs1h 锁定 I-2.2:cache_creation 分 +// CacheCreation5mPrice / CacheCreation1hPrice 两档分别计价。 +// 定价:input $3/MTok, output $15/MTok, cache_read $0.3/MTok, +// 5m 写入 $4/MTok, 1h 写入 $5/MTok(SupportsCacheBreakdown=true)。 +func TestBillingInvariant_CacheTier5mVs1h(t *testing.T) { + breakdownPricing := &ModelPricing{ + InputPricePerToken: 3e-6, + OutputPricePerToken: 15e-6, + CacheReadPricePerToken: 0.3e-6, + CacheCreationPricePerToken: 3.75e-6, + SupportsCacheBreakdown: true, + CacheCreation5mPrice: 4e-6, + CacheCreation1hPrice: 5e-6, + } + noBreakdownPricing := &ModelPricing{ + InputPricePerToken: 3e-6, + OutputPricePerToken: 15e-6, + CacheReadPricePerToken: 0.3e-6, + CacheCreationPricePerToken: 3.75e-6, + SupportsCacheBreakdown: false, + } + + tests := []struct { + name string + pricing *ModelPricing + tokens UsageTokens + // 期望值(固定值 + 1e-10 容差) + wantInput float64 + wantOutput float64 + wantCacheCreation float64 + wantCacheRead float64 + wantTotal float64 + }{ + { + name: "5m与1h两档分别计价", + pricing: breakdownPricing, + tokens: UsageTokens{ + InputTokens: 1000, + OutputTokens: 500, + // 上游通常同时给出总量与 ephemeral 明细;有明细时按明细计价 + CacheCreationTokens: 12000, + CacheCreation5mTokens: 8000, + CacheCreation1hTokens: 4000, + CacheReadTokens: 2000, + }, + wantInput: 0.003, // 1000 × $3/MTok = 0.003 + wantOutput: 0.0075, // 500 × $15/MTok = 0.0075 + // 8000 × 4e-6 = 0.032; 4000 × 5e-6 = 0.020; 合计 0.052 + wantCacheCreation: 0.052, + wantCacheRead: 0.0006, // 2000 × $0.3/MTok = 0.0006 + wantTotal: 0.0631, // 0.003+0.0075+0.052+0.0006 + }, + { + name: "仅1h写入按1h单价", + pricing: breakdownPricing, + tokens: UsageTokens{ + CacheCreationTokens: 10000, + CacheCreation1hTokens: 10000, + }, + wantCacheCreation: 0.05, // 10000 × 5e-6 = 0.05 + wantTotal: 0.05, + }, + { + name: "无ephemeral明细时全部回退5m单价", + pricing: breakdownPricing, + tokens: UsageTokens{ + CacheCreationTokens: 6000, // 5m/1h 明细均为 0 + }, + wantCacheCreation: 0.024, // 6000 × 4e-6 = 0.024(回退 5m 档) + wantTotal: 0.024, + }, + { + name: "breakdown关闭时按标准单价计CacheCreationTokens并忽略5m1h明细", + pricing: noBreakdownPricing, + tokens: UsageTokens{ + CacheCreationTokens: 6000, + CacheCreation5mTokens: 4000, + CacheCreation1hTokens: 2000, + }, + wantCacheCreation: 0.0225, // 6000 × 3.75e-6 = 0.0225(明细被忽略) + wantTotal: 0.0225, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := billInvManualPricingService(tt.pricing) + cost, err := svc.CalculateCost("claude-sonnet-4", tt.tokens, 1.0) + require.NoError(t, err) + require.InDelta(t, tt.wantInput, cost.InputCost, 1e-10) + require.InDelta(t, tt.wantOutput, cost.OutputCost, 1e-10) + require.InDelta(t, tt.wantCacheCreation, cost.CacheCreationCost, 1e-10) + require.InDelta(t, tt.wantCacheRead, cost.CacheReadCost, 1e-10) + require.InDelta(t, tt.wantTotal, cost.TotalCost, 1e-10) + require.InDelta(t, tt.wantTotal, cost.ActualCost, 1e-10) // 倍率 1.0 + }) + } +} + +// TestBillingInvariant_CacheTierBreakdownGating 锁定 I-2.2 的动态定价门控: +// 仅当 LiteLLM 的 1h 单价存在且严格大于 5m 单价时启用两档计费 +// (防止上游数据错误导致少收费,见 billing_service.go GetModelPricing)。 +func TestBillingInvariant_CacheTierBreakdownGating(t *testing.T) { + tokens := UsageTokens{ + CacheCreationTokens: 1000, + CacheCreation5mTokens: 600, + CacheCreation1hTokens: 400, + } + + tests := []struct { + name string + price5m, price1h float64 + wantBreakdown bool + wantCacheCreation float64 + }{ + { + name: "1h大于5m启用两档", + price5m: 4e-6, price1h: 5e-6, + wantBreakdown: true, + // 600 × 4e-6 = 0.0024; 400 × 5e-6 = 0.002; 合计 0.0044 + wantCacheCreation: 0.0044, + }, + { + name: "1h等于5m禁用两档防少收费", + price5m: 4e-6, price1h: 4e-6, + wantBreakdown: false, + // 回退标准单价:1000 × 4e-6 = 0.004 + wantCacheCreation: 0.004, + }, + { + name: "1h缺失禁用两档", + price5m: 4e-6, price1h: 0, + wantBreakdown: false, + // 1000 × 4e-6 = 0.004 + wantCacheCreation: 0.004, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := NewBillingService(&config.Config{}, &PricingService{ + pricingData: map[string]*LiteLLMModelPricing{ + "tier-gate-model": { + InputCostPerToken: 3e-6, + OutputCostPerToken: 15e-6, + CacheCreationInputTokenCost: tt.price5m, + CacheCreationInputTokenCostAbove1hr: tt.price1h, + CacheReadInputTokenCost: 0.3e-6, + }, + }, + }) + + pricing, err := svc.GetModelPricing("tier-gate-model") + require.NoError(t, err) + require.Equal(t, tt.wantBreakdown, pricing.SupportsCacheBreakdown) + + cost, err := svc.CalculateCost("tier-gate-model", tokens, 1.0) + require.NoError(t, err) + require.InDelta(t, tt.wantCacheCreation, cost.CacheCreationCost, 1e-10) + }) + } +} + +// TestBillingInvariant_ServiceTierThenRateMultiplier 锁定 I-2.4 的应用顺序: +// 各分项成本先应用 serviceTier 倍率(或 priority 显式价),TotalCost 为分项之和 +// (含 tier、不含 rateMultiplier),最后 ActualCost = TotalCost × rateMultiplier。 +func TestBillingInvariant_ServiceTierThenRateMultiplier(t *testing.T) { + tests := []struct { + name string + model string + serviceTier string + rateMultiplier float64 + tokens UsageTokens + wantInput float64 + wantOutput float64 + wantCacheWrite float64 + wantCacheRead float64 + wantTotal float64 + wantActual float64 + }{ + { + // gpt-5.4 fallback 价:in $2.5/MTok, out $15/MTok, cache_write $2.5/MTok, + // cache_read $0.25/MTok;flex 无显式价 → 0.5 倍 tier multiplier。 + name: "flex_tier后再乘rateMultiplier", + model: "gpt-5.4", + serviceTier: "flex", + rateMultiplier: 1.5, + tokens: UsageTokens{InputTokens: 1000, OutputTokens: 500, CacheCreationTokens: 400, CacheReadTokens: 2000}, + wantInput: 0.00125, // 1000×2.5e-6=0.0025 ×0.5 + wantOutput: 0.00375, // 500×15e-6=0.0075 ×0.5 + wantCacheWrite: 0.0005, // 400×2.5e-6=0.001 ×0.5 + wantCacheRead: 0.00025, // 2000×0.25e-6=0.0005 ×0.5 + wantTotal: 0.00575, // 分项之和(含 tier、不含 rate) + wantActual: 0.008625, // 0.00575 × 1.5 + }, + { + // gpt-5.4 有显式 priority 价:in $5/MTok, out $30/MTok, cache_read $0.5/MTok。 + // 注意(characterization):cache_write 无 priority 价,且显式 priority 价 + // 生效时 tierMultiplier 固定 1.0,因此 cache_write 仍按基础价 $2.5/MTok + // 计费、不做 2 倍上浮——与下方"无显式价回退 2 倍"的行为不同。 + name: "priority显式价后再乘rateMultiplier", + model: "gpt-5.4", + serviceTier: "priority", + rateMultiplier: 2.0, + tokens: UsageTokens{InputTokens: 1000, OutputTokens: 500, CacheCreationTokens: 1000, CacheReadTokens: 2000}, + wantInput: 0.005, // 1000 × 5e-6(priority 价) + wantOutput: 0.015, // 500 × 30e-6(priority 价) + wantCacheWrite: 0.0025, // 1000 × 2.5e-6(基础价,无 priority 上浮) + wantCacheRead: 0.001, // 2000 × 0.5e-6(priority 价) + wantTotal: 0.0235, + wantActual: 0.047, // 0.0235 × 2.0 + }, + { + // claude-sonnet-4 fallback 无 priority 显式价 → 回退 2.0 倍 tier multiplier, + // 此路径下 cache_write 也参与 2 倍上浮。 + name: "priority无显式价回退2倍tier后再乘rateMultiplier", + model: "claude-sonnet-4", + serviceTier: "priority", + rateMultiplier: 1.5, + tokens: UsageTokens{InputTokens: 1000, OutputTokens: 500, CacheCreationTokens: 2000, CacheReadTokens: 3000}, + wantInput: 0.006, // 1000×3e-6=0.003 ×2 + wantOutput: 0.015, // 500×15e-6=0.0075 ×2 + wantCacheWrite: 0.015, // 2000×3.75e-6=0.0075 ×2 + wantCacheRead: 0.0018, // 3000×0.3e-6=0.0009 ×2 + wantTotal: 0.0378, + wantActual: 0.0567, // 0.0378 × 1.5 + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := newTestBillingService() + cost, err := svc.CalculateCostWithServiceTier(tt.model, tt.tokens, tt.rateMultiplier, tt.serviceTier) + require.NoError(t, err) + require.InDelta(t, tt.wantInput, cost.InputCost, 1e-10) + require.InDelta(t, tt.wantOutput, cost.OutputCost, 1e-10) + require.InDelta(t, tt.wantCacheWrite, cost.CacheCreationCost, 1e-10) + require.InDelta(t, tt.wantCacheRead, cost.CacheReadCost, 1e-10) + require.InDelta(t, tt.wantTotal, cost.TotalCost, 1e-10) + require.InDelta(t, tt.wantActual, cost.ActualCost, 1e-10) + }) + } +} + +// TestBillingInvariant_ZeroAndNegativeRateMultiplier 锁定 I-2.4: +// rateMultiplier=0(免费账号)时 ActualCost=0 但 TotalCost 保留; +// 负数倍率按 0 处理(防缓存/迁移残留导致按 1x 误扣)。 +func TestBillingInvariant_ZeroAndNegativeRateMultiplier(t *testing.T) { + svc := newTestBillingService() + tokens := UsageTokens{InputTokens: 1000, OutputTokens: 500, CacheCreationTokens: 2000, CacheReadTokens: 3000} + // claude-sonnet-4: 0.003 + 0.0075 + 2000×3.75e-6(=0.0075) + 3000×0.3e-6(=0.0009) = 0.0189 + const wantTotal = 0.0189 + + for _, multiplier := range []float64{0, -2} { + cost, err := svc.CalculateCost("claude-sonnet-4", tokens, multiplier) + require.NoError(t, err) + require.InDelta(t, wantTotal, cost.TotalCost, 1e-10, "multiplier=%v 时 TotalCost 应保留原始费用", multiplier) + require.InDelta(t, 0.0, cost.ActualCost, 1e-10, "multiplier=%v 时 ActualCost 应为 0", multiplier) + } +} diff --git a/backend/internal/service/billing_invariants_e2e_test.go b/backend/internal/service/billing_invariants_e2e_test.go new file mode 100644 index 0000000000..9b23a32528 --- /dev/null +++ b/backend/internal/service/billing_invariants_e2e_test.go @@ -0,0 +1,546 @@ +//go:build unit + +// TASK-003 计费/配额端到端不变量测试(INVARIANTS.md ②③)。 +// +// 本文件通过 GatewayService.RecordUsage 走完整后扣链路,锁定: +// - I-2.6: 完整 usage 输入 → 最终扣费金额 + 余额扣减/订阅用量增量二选一分支 +// (含 image 计价路径与 ImageOutputTokens 独立计价,呼应 I-2.3 端到端联动) +// - I-2.2: 5m/1h 两档缓存写入的端到端金额 +// - I-3.3: API Key quota_used 增量金额(统一命令路径 + legacy 直写路径) +// - I-3.4: Account 级配额增量金额 = TotalCost × AccountRateMultiplier +// (注意:用 TotalCost 而非 ActualCost,分组倍率不影响账号配额消耗) +// +// 断言对象为 UsageBillingRepository.Apply 收到的 UsageBillingCommand(生产原子 +// 扣费路径的唯一输入)与 UsageLog 金额字段。所有期望值均人工核算(算式见注释)。 +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +// billInvAccountRepoStub 捕获 legacy 路径下的账号配额增量。 +type billInvAccountRepoStub struct { + AccountRepository + + quotaIncrCalls int + lastQuotaAmount float64 +} + +func (s *billInvAccountRepoStub) IncrementQuotaUsed(ctx context.Context, id int64, amount float64) error { + s.quotaIncrCalls++ + s.lastQuotaAmount = amount + return nil +} + +// billInvSubRepoStub 捕获 legacy 路径下的订阅用量增量金额。 +type billInvSubRepoStub struct { + UserSubscriptionRepository + + incrementCalls int + lastCost float64 +} + +func (s *billInvSubRepoStub) IncrementUsage(ctx context.Context, id int64, costUSD float64) error { + s.incrementCalls++ + s.lastCost = costUSD + return nil +} + +func billInvF64Ptr(v float64) *float64 { return &v } + +// billInvNewGatewayService 构造带可注入 BillingService / 计费仓储的网关服务。 +func billInvNewGatewayService( + billingSvc *BillingService, + accountRepo AccountRepository, + usageRepo UsageLogRepository, + billingRepo UsageBillingRepository, + userRepo UserRepository, + subRepo UserSubscriptionRepository, +) *GatewayService { + cfg := &config.Config{} + cfg.Default.RateMultiplier = 1.0 + if billingSvc == nil { + billingSvc = NewBillingService(cfg, nil) + } + svc := NewGatewayService( + accountRepo, + nil, // groupRepo + usageRepo, + billingRepo, + userRepo, + subRepo, + nil, // userGroupRateRepo + nil, // cache + cfg, + nil, // schedulerSnapshot + nil, // concurrencyService + billingSvc, + nil, // rateLimitService + &BillingCacheService{}, + nil, // identityService + nil, // httpUpstream + &DeferredService{}, + nil, // claudeTokenProvider + nil, // sessionLimitCache + nil, // rpmCache + nil, // digestStore + nil, // settingService + nil, // tlsFPProfileService + nil, // channelService + nil, // resolver + nil, // balanceNotifyService + nil, // userPlatformQuotaRepo + ) + return svc +} + +// TestBillingInvariant_EndToEndUsageBillingCommand 表驱动锁定 I-2.6 / I-3.3 / I-3.4: +// 每行 = (定价/分组倍率/账号倍率/usage 输入) → (UsageBillingCommand 各项金额 + UsageLog 金额)。 +func TestBillingInvariant_EndToEndUsageBillingCommand(t *testing.T) { + groupID := int64(11) + type expected struct { + balanceCost float64 + subscriptionCost float64 + apiKeyQuotaCost float64 + apiKeyRateLimitCost float64 + accountQuotaCost float64 + subscriptionID *int64 + logTotal float64 + logActual float64 + billingType int8 + } + + // 公共 usage:claude-sonnet-4 fallback 价(in $3/MTok, out $15/MTok, + // cache_write $3.75/MTok, cache_read $0.3/MTok)。 + // TotalCost = 1000×3e-6 + 500×15e-6 + 2000×3.75e-6 + 3000×0.3e-6 + // = 0.003 + 0.0075 + 0.0075 + 0.0009 = 0.0189 + fourDimUsage := ClaudeUsage{ + InputTokens: 1000, + OutputTokens: 500, + CacheCreationInputTokens: 2000, + CacheReadInputTokens: 3000, + } + + tests := []struct { + name string + pricingData map[string]*LiteLLMModelPricing // 可选:动态定价 + model string + usage ClaudeUsage + imageCount int + imageSize string + group *Group + account *Account + apiKeyQuota float64 + rateLimit5h float64 + sub *UserSubscription + want expected + }{ + { + name: "余额模式_四维token组合计价", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 1.0}, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + // API Key Quota=100 → quota 增量 = ActualCost + apiKeyQuota: 100, + want: expected{ + balanceCost: 0.0189, // ActualCost = 0.0189 × 1.0 + apiKeyQuotaCost: 0.0189, // I-3.3: = ActualCost + logTotal: 0.0189, + logActual: 0.0189, + billingType: BillingTypeBalance, + }, + }, + { + name: "余额模式_分组倍率2x只影响Actual不影响Total", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 2.0}, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + apiKeyQuota: 100, + want: expected{ + balanceCost: 0.0378, // 0.0189 × 2.0 + apiKeyQuotaCost: 0.0378, // I-3.3: 跟随 ActualCost + logTotal: 0.0189, // TotalCost 不含分组倍率 + logActual: 0.0378, + billingType: BillingTypeBalance, + }, + }, + { + name: "订阅模式_订阅用量增量替代余额扣减", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ + ID: groupID, Platform: PlatformAnthropic, + RateMultiplier: 1.5, SubscriptionType: SubscriptionTypeSubscription, + }, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + sub: &UserSubscription{ID: 42}, + want: expected{ + balanceCost: 0, // 二选一:订阅模式不扣余额 + subscriptionCost: 0.02835, // 0.0189 × 1.5 + subscriptionID: i64p(42), + logTotal: 0.0189, + logActual: 0.02835, + billingType: BillingTypeSubscription, + }, + }, + { + name: "订阅模式_免费分组倍率0时订阅增量为0但Total保留", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ + ID: groupID, Platform: PlatformAnthropic, + RateMultiplier: 0, SubscriptionType: SubscriptionTypeSubscription, + }, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + sub: &UserSubscription{ID: 42}, + want: expected{ + balanceCost: 0, + subscriptionCost: 0, // ActualCost = 0 + subscriptionID: i64p(42), // TotalCost > 0 仍记录订阅归属 + logTotal: 0.0189, + logActual: 0, + billingType: BillingTypeSubscription, + }, + }, + { + name: "账号配额增量用TotalCost乘账号倍率而非ActualCost", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 2.0}, + account: &Account{ + ID: 702, Type: AccountTypeAPIKey, + Extra: map[string]any{"quota_limit": 100.0}, + RateMultiplier: billInvF64Ptr(0.5), + }, + want: expected{ + balanceCost: 0.0378, // 余额按分组倍率 2.0 + accountQuotaCost: 0.00945, // I-3.4: 0.0189(Total) × 0.5(账号倍率) + logTotal: 0.0189, + logActual: 0.0378, + billingType: BillingTypeBalance, + }, + }, + { + name: "OAuth账号不计账号配额增量", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 1.0}, + account: &Account{ + ID: 703, Type: AccountTypeOAuth, + Extra: map[string]any{"quota_limit": 100.0}, + }, + want: expected{ + balanceCost: 0.0189, + accountQuotaCost: 0, // 仅 apikey/bedrock 账号计配额 + logTotal: 0.0189, + logActual: 0.0189, + billingType: BillingTypeBalance, + }, + }, + { + name: "APIKey无限额时quota增量为0", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 1.0}, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + apiKeyQuota: 0, // Quota=0 表示不限额 + want: expected{ + balanceCost: 0.0189, + apiKeyQuotaCost: 0, + logTotal: 0.0189, + logActual: 0.0189, + billingType: BillingTypeBalance, + }, + }, + { + name: "APIKey限速窗口增量等于ActualCost", + model: "claude-sonnet-4", + usage: fourDimUsage, + group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 1.0}, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + rateLimit5h: 10, + want: expected{ + balanceCost: 0.0189, + apiKeyRateLimitCost: 0.0189, + logTotal: 0.0189, + logActual: 0.0189, + billingType: BillingTypeBalance, + }, + }, + { + name: "图片按张计价_分组2K价格", + model: "gemini-3-pro-image", + imageCount: 2, + imageSize: "2K", + group: &Group{ + ID: groupID, Platform: PlatformGemini, RateMultiplier: 1.0, + ImagePrice2K: billInvF64Ptr(0.19), + }, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + want: expected{ + balanceCost: 0.38, // 2 × $0.19 + logTotal: 0.38, + logActual: 0.38, + billingType: BillingTypeBalance, + }, + }, + { + name: "图片独立倍率只影响Actual", + model: "gemini-3-pro-image", + imageCount: 2, + imageSize: "2K", + group: &Group{ + ID: groupID, Platform: PlatformGemini, RateMultiplier: 1.0, + ImagePrice2K: billInvF64Ptr(0.19), + ImageRateIndependent: true, + ImageRateMultiplier: 2.0, + }, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + want: expected{ + balanceCost: 0.76, // 0.38 × 2.0(图片独立倍率) + logTotal: 0.38, + logActual: 0.76, + billingType: BillingTypeBalance, + }, + }, + { + name: "ImageOutputTokens按独立单价从输出中拆分计价", + pricingData: map[string]*LiteLLMModelPricing{ + "img-token-model": { + InputCostPerToken: 3e-6, + OutputCostPerToken: 15e-6, + OutputCostPerImageToken: 30e-6, + }, + }, + model: "img-token-model", + usage: ClaudeUsage{ + InputTokens: 100, + OutputTokens: 200, // 其中 50 为图片输出 token + ImageOutputTokens: 50, + }, + group: &Group{ID: groupID, Platform: PlatformGemini, RateMultiplier: 1.0}, + account: &Account{ID: 701, Type: AccountTypeOAuth}, + want: expected{ + // input: 100×3e-6 = 0.0003 + // 文本输出: (200-50)×15e-6 = 0.00225 + // 图片输出: 50×30e-6 = 0.0015 + // Total = 0.00405 + balanceCost: 0.00405, + logTotal: 0.00405, + logActual: 0.00405, + billingType: BillingTypeBalance, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}} + var billingSvc *BillingService + if tt.pricingData != nil { + cfg := &config.Config{} + billingSvc = NewBillingService(cfg, &PricingService{pricingData: tt.pricingData}) + } + svc := billInvNewGatewayService(billingSvc, nil, usageRepo, billingRepo, + &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}) + quotaSvc := &openAIRecordUsageAPIKeyQuotaStub{} + + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "bill-inv-" + tt.name, + Usage: tt.usage, + Model: tt.model, + ImageCount: tt.imageCount, + ImageSize: tt.imageSize, + Duration: time.Second, + }, + APIKey: &APIKey{ + ID: 501, + Quota: tt.apiKeyQuota, + RateLimit5h: tt.rateLimit5h, + GroupID: i64p(tt.group.ID), + Group: tt.group, + }, + User: &User{ID: 601}, + Account: tt.account, + Subscription: tt.sub, + APIKeyService: quotaSvc, + }) + require.NoError(t, err) + + cmd := billingRepo.lastCmd + require.NotNil(t, cmd, "应走统一计费命令路径") + require.InDelta(t, tt.want.balanceCost, cmd.BalanceCost, 1e-10, "BalanceCost") + require.InDelta(t, tt.want.subscriptionCost, cmd.SubscriptionCost, 1e-10, "SubscriptionCost") + require.InDelta(t, tt.want.apiKeyQuotaCost, cmd.APIKeyQuotaCost, 1e-10, "APIKeyQuotaCost") + require.InDelta(t, tt.want.apiKeyRateLimitCost, cmd.APIKeyRateLimitCost, 1e-10, "APIKeyRateLimitCost") + require.InDelta(t, tt.want.accountQuotaCost, cmd.AccountQuotaCost, 1e-10, "AccountQuotaCost") + if tt.want.subscriptionID != nil { + require.NotNil(t, cmd.SubscriptionID) + require.Equal(t, *tt.want.subscriptionID, *cmd.SubscriptionID) + } else { + require.Nil(t, cmd.SubscriptionID) + } + + log := usageRepo.lastLog + require.NotNil(t, log) + require.InDelta(t, tt.want.logTotal, log.TotalCost, 1e-10, "UsageLog.TotalCost") + require.InDelta(t, tt.want.logActual, log.ActualCost, 1e-10, "UsageLog.ActualCost") + require.Equal(t, tt.want.billingType, log.BillingType) + }) + } +} + +// TestBillingInvariant_CacheTier5mVs1hEndToEnd 锁定 I-2.2 端到端:上游返回的 +// 5m/1h ephemeral 明细流经 RecordUsage 后按两档单价分别计费(PR3061 漂移点②)。 +func TestBillingInvariant_CacheTier5mVs1hEndToEnd(t *testing.T) { + cfg := &config.Config{} + billingSvc := NewBillingService(cfg, &PricingService{ + pricingData: map[string]*LiteLLMModelPricing{ + "claude-sonnet-4": { + InputCostPerToken: 3e-6, + OutputCostPerToken: 15e-6, + CacheCreationInputTokenCost: 3.75e-6, // 5m 档 + CacheCreationInputTokenCostAbove1hr: 6e-6, // 1h 档(> 5m → 启用两档) + CacheReadInputTokenCost: 0.3e-6, + }, + }, + }) + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}} + svc := billInvNewGatewayService(billingSvc, nil, usageRepo, billingRepo, + &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}) + + groupID := int64(11) + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "bill-inv-cache-tier-e2e", + Usage: ClaudeUsage{ + InputTokens: 1000, + OutputTokens: 500, + CacheCreationInputTokens: 12000, // 总量 = 5m + 1h + CacheCreation5mTokens: 8000, + CacheCreation1hTokens: 4000, + CacheReadInputTokens: 2000, + }, + Model: "claude-sonnet-4", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 501, GroupID: i64p(groupID), Group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 1.0}}, + User: &User{ID: 601}, + Account: &Account{ID: 701, Type: AccountTypeOAuth}, + }) + require.NoError(t, err) + + // cache_creation = 8000×3.75e-6 + 4000×6e-6 = 0.03 + 0.024 = 0.054 + // Total = 1000×3e-6 + 500×15e-6 + 0.054 + 2000×0.3e-6 + // = 0.003 + 0.0075 + 0.054 + 0.0006 = 0.0651 + log := usageRepo.lastLog + require.NotNil(t, log) + require.InDelta(t, 0.054, log.CacheCreationCost, 1e-10) + require.InDelta(t, 0.0651, log.TotalCost, 1e-10) + require.InDelta(t, 0.0651, log.ActualCost, 1e-10) + + require.NotNil(t, billingRepo.lastCmd) + require.InDelta(t, 0.0651, billingRepo.lastCmd.BalanceCost, 1e-10) +} + +// TestBillingInvariant_LegacyPathIncrements 锁定 legacy 兜底路径(统一计费仓储 +// 不可用时)的各级增量金额与生产命令路径一致: +// - 余额扣减 = ActualCost(I-2.6) +// - API Key quota_used 增量 = ActualCost(I-3.3) +// - 账号配额增量 = TotalCost × AccountRateMultiplier(I-3.4) +// - 订阅模式:IncrementUsage(ActualCost),不扣余额(I-2.6 二选一) +func TestBillingInvariant_LegacyPathIncrements(t *testing.T) { + groupID := int64(11) + fourDimUsage := ClaudeUsage{ + InputTokens: 1000, + OutputTokens: 500, + CacheCreationInputTokens: 2000, + CacheReadInputTokens: 3000, + } + + t.Run("余额模式各级增量", func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &billInvSubRepoStub{} + accountRepo := &billInvAccountRepoStub{} + quotaSvc := &openAIRecordUsageAPIKeyQuotaStub{} + // billingRepo == nil → applyUsageBilling 回退 postUsageBilling 直写 + svc := billInvNewGatewayService(nil, accountRepo, usageRepo, nil, userRepo, subRepo) + + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "bill-inv-legacy-balance", + Usage: fourDimUsage, + Model: "claude-sonnet-4", + Duration: time.Second, + }, + APIKey: &APIKey{ + ID: 501, Quota: 100, + GroupID: i64p(groupID), + Group: &Group{ID: groupID, Platform: PlatformAnthropic, RateMultiplier: 2.0}, + }, + User: &User{ID: 601}, + Account: &Account{ + ID: 702, Type: AccountTypeAPIKey, + Extra: map[string]any{"quota_limit": 100.0}, + RateMultiplier: billInvF64Ptr(0.5), + }, + APIKeyService: quotaSvc, + }) + require.NoError(t, err) + + // ActualCost = 0.0189 × 2.0 = 0.0378 + require.Equal(t, 1, userRepo.deductCalls) + require.InDelta(t, 0.0378, userRepo.lastAmount, 1e-10, "余额扣减 = ActualCost") + require.Equal(t, 1, quotaSvc.quotaCalls) + require.InDelta(t, 0.0378, quotaSvc.lastAmount, 1e-10, "API Key quota 增量 = ActualCost") + // 账号配额 = TotalCost(0.0189) × 账号倍率(0.5) = 0.00945 + require.Equal(t, 1, accountRepo.quotaIncrCalls) + require.InDelta(t, 0.00945, accountRepo.lastQuotaAmount, 1e-10, "账号配额增量 = Total × 账号倍率") + // 余额模式不应触发订阅增量 + require.Equal(t, 0, subRepo.incrementCalls) + }) + + t.Run("订阅模式增量替代余额扣减", func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &billInvSubRepoStub{} + svc := billInvNewGatewayService(nil, &billInvAccountRepoStub{}, usageRepo, nil, userRepo, subRepo) + + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "bill-inv-legacy-subscription", + Usage: fourDimUsage, + Model: "claude-sonnet-4", + Duration: time.Second, + }, + APIKey: &APIKey{ + ID: 501, + GroupID: i64p(groupID), + Group: &Group{ + ID: groupID, Platform: PlatformAnthropic, + RateMultiplier: 1.5, SubscriptionType: SubscriptionTypeSubscription, + }, + }, + User: &User{ID: 601}, + Account: &Account{ID: 701, Type: AccountTypeOAuth}, + Subscription: &UserSubscription{ID: 42}, + }) + require.NoError(t, err) + + // 订阅增量 = ActualCost = 0.0189 × 1.5 = 0.02835;余额不扣 + require.Equal(t, 1, subRepo.incrementCalls) + require.InDelta(t, 0.02835, subRepo.lastCost, 1e-10, "订阅用量增量 = ActualCost") + require.Equal(t, 0, userRepo.deductCalls, "订阅模式不应扣余额(二选一)") + }) +} diff --git a/backend/internal/service/billing_invariants_overages_test.go b/backend/internal/service/billing_invariants_overages_test.go new file mode 100644 index 0000000000..4d147bbb02 --- /dev/null +++ b/backend/internal/service/billing_invariants_overages_test.go @@ -0,0 +1,141 @@ +//go:build unit + +// TASK-003 overages 不变量测试(INVARIANTS.md I-2.5,PR3061 漂移点①)。 +// +// overages(AI Credits 超量请求)由 accounts.extra["allow_overages"] 控制,仅 +// antigravity 平台生效。本文件锁定: +// - 开关解析的允许/拒绝语义(平台门控 + 字段缺失/类型错误按拒绝处理) +// - 拒绝路径:上游 429 quota_exhausted 时不注入 enabledCreditTypes、不发起 +// credits 重试(开关关闭 / 积分已耗尽两种拒绝原因) +// +// 允许路径(注入 credits 并继续请求)已由 antigravity_credits_overages_test.go +// 覆盖。超额部分的"计价"没有独立分支:credits 重试成功后的响应走与普通请求 +// 完全相同的 RecordUsage 计费链路(由本任务 I-2.6 端到端测试锁定金额),因此 +// 这里只需锁定允许/拒绝语义。 +package service + +import ( + "bytes" + "context" + "io" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// TestBillingInvariant_OveragesFlagSemantics 锁定 allow_overages 开关解析语义。 +func TestBillingInvariant_OveragesFlagSemantics(t *testing.T) { + tests := []struct { + name string + account Account + want bool + }{ + { + name: "antigravity平台显式true允许", + account: Account{Platform: PlatformAntigravity, Extra: map[string]any{"allow_overages": true}}, + want: true, + }, + { + name: "antigravity平台显式false拒绝", + account: Account{Platform: PlatformAntigravity, Extra: map[string]any{"allow_overages": false}}, + want: false, + }, + { + name: "字段缺失默认拒绝", + account: Account{Platform: PlatformAntigravity, Extra: map[string]any{}}, + want: false, + }, + { + name: "Extra为nil默认拒绝", + account: Account{Platform: PlatformAntigravity}, + want: false, + }, + { + name: "非bool类型按拒绝处理", + account: Account{Platform: PlatformAntigravity, Extra: map[string]any{"allow_overages": "true"}}, + want: false, + }, + { + name: "非antigravity平台即使为true也拒绝", + account: Account{Platform: PlatformAnthropic, Extra: map[string]any{"allow_overages": true}}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, tt.account.IsOveragesEnabled()) + }) + } +} + +// TestBillingInvariant_OveragesDeniedNoCreditsInjection 锁定拒绝路径: +// 上游返回 429 quota_exhausted 时,若 overages 关闭或积分已耗尽, +// handleSmartRetry 不得注入 enabledCreditTypes 发起 credits 重试, +// 而是落入默认重试逻辑(smartRetryActionContinue)。 +func TestBillingInvariant_OveragesDeniedNoCreditsInjection(t *testing.T) { + quotaExhaustedBody := []byte(`{"error":{"status":"RESOURCE_EXHAUSTED","message":"QUOTA_EXHAUSTED"}}`) + + tests := []struct { + name string + extra map[string]any + }{ + { + name: "开关关闭时不注入credits", + extra: map[string]any{}, // allow_overages 缺失 → 拒绝 + }, + { + name: "积分已耗尽时不注入credits", + extra: map[string]any{ + "allow_overages": true, + modelRateLimitsKey: map[string]any{ + creditsExhaustedKey: map[string]any{ + "rate_limited_at": time.Now().UTC().Format(time.RFC3339), + "rate_limit_reset_at": time.Now().Add(5 * time.Hour).UTC().Format(time.RFC3339), + }, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &mockSmartRetryUpstream{} + account := &Account{ + ID: 201, + Name: "acc-201", + Type: AccountTypeOAuth, + Platform: PlatformAntigravity, + Extra: tt.extra, + } + resp := &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{}, + Body: io.NopCloser(bytes.NewReader(quotaExhaustedBody)), + } + params := antigravityRetryLoopParams{ + ctx: context.Background(), + prefix: "[bill-inv]", + account: account, + accessToken: "token", + action: "generateContent", + body: []byte(`{"model":"claude-sonnet-4-5","request":{}}`), + httpUpstream: upstream, + accountRepo: &stubAntigravityAccountRepo{}, + requestedModel: "claude-sonnet-4-5", + handleError: func(ctx context.Context, prefix string, account *Account, statusCode int, headers http.Header, body []byte, requestedModel string, groupID int64, sessionHash string, isStickySession bool) *handleModelRateLimitResult { + return nil + }, + } + + svc := &AntigravityGatewayService{} + result := svc.handleSmartRetry(params, resp, quotaExhaustedBody, "https://ag-1.test", 0, []string{"https://ag-1.test"}) + + require.NotNil(t, result) + require.Equal(t, smartRetryActionContinue, result.action, "拒绝 overages 后应落入默认重试逻辑") + require.Empty(t, upstream.requestBodies, "不得发起 credits 注入重试请求") + }) + } +} diff --git a/backend/internal/service/billing_invariants_preflight_test.go b/backend/internal/service/billing_invariants_preflight_test.go new file mode 100644 index 0000000000..0b6c44a0e7 --- /dev/null +++ b/backend/internal/service/billing_invariants_preflight_test.go @@ -0,0 +1,226 @@ +//go:build unit + +// TASK-003 preflight 配额/余额拒绝不变量测试(INVARIANTS.md I-3.6)。 +// +// CheckBillingEligibility 是网关入口的统一计费资格检查;handler 在 +// AcquireUserSlotWithWait 等待结束后会用同一函数做"二次检查" +// (internal/handler/gateway_handler.go:255),因此这里锁定其无状态语义: +// - 余额模式:余额 <= 0 → ErrInsufficientBalance +// - 订阅模式:日/周/月用量达到分组限额 → ErrDaily/Weekly/MonthlyLimitExceeded; +// 订阅过期或非 active → ErrSubscriptionInvalid +// - user×platform 配额:日限额耗尽 → ErrUserPlatformDailyQuotaExhausted(429); +// 订阅模式豁免该检查 +// - 并发等待后二次检查:第一次放行后用量/余额变化,再次调用即拒绝 +// - simple 运行模式跳过所有计费检查 +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +// billInvBillingCacheStub 只实现 CheckBillingEligibility 路径会触达的读方法, +// 其余方法走嵌入接口(不会被调用)。 +type billInvBillingCacheStub struct { + BillingCache + + balance float64 + sub *SubscriptionCacheData + quotaEntry *UserPlatformQuotaCacheEntry +} + +func (s *billInvBillingCacheStub) GetUserBalance(ctx context.Context, userID int64) (float64, error) { + return s.balance, nil +} + +func (s *billInvBillingCacheStub) GetSubscriptionCache(ctx context.Context, userID, groupID int64) (*SubscriptionCacheData, error) { + return s.sub, nil +} + +func (s *billInvBillingCacheStub) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*UserPlatformQuotaCacheEntry, bool, error) { + if s.quotaEntry == nil { + return nil, false, nil + } + return s.quotaEntry, true, nil +} + +// billInvQuotaRepoStub 仅用于让 userPlatformQuotaRepo 非 nil(cache HIT 路径不 +// 会触达 DB),嵌入接口的其余方法不会被调用。 +type billInvQuotaRepoStub struct { + UserPlatformQuotaRepository +} + +func (s *billInvQuotaRepoStub) GetByUserPlatform(ctx context.Context, userID int64, platform string) (*UserPlatformQuotaRecord, error) { + return nil, nil +} + +func billInvNewBillingCacheService(t *testing.T, cache BillingCache, cfg *config.Config) *BillingCacheService { + t.Helper() + if cfg == nil { + cfg = &config.Config{} + } + svc := NewBillingCacheService(cache, nil, nil, nil, nil, nil, cfg, &billInvQuotaRepoStub{}) + t.Cleanup(svc.Stop) + return svc +} + +// billInvQuotaEntryV1 构造当前窗口内的 SchemaV1 user×platform 配额缓存条目。 +func billInvQuotaEntryV1(dailyLimit, dailyUsage float64) *UserPlatformQuotaCacheEntry { + now := time.Now() + return &UserPlatformQuotaCacheEntry{ + SchemaVersion: UserPlatformQuotaCacheSchemaV1, + DailyLimitUSD: &dailyLimit, + DailyUsageUSD: dailyUsage, + DailyWindowStart: &now, + WeeklyWindowStart: &now, + MonthlyWindowStart: &now, + } +} + +// TestBillingInvariant_PreflightBalanceEligibility 锁定余额模式 preflight 语义。 +func TestBillingInvariant_PreflightBalanceEligibility(t *testing.T) { + user := &User{ID: 601} + + t.Run("余额耗尽拒绝", func(t *testing.T) { + svc := billInvNewBillingCacheService(t, &billInvBillingCacheStub{balance: 0}, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, "") + require.ErrorIs(t, err, ErrInsufficientBalance) + }) + + t.Run("余额为正放行", func(t *testing.T) { + svc := billInvNewBillingCacheService(t, &billInvBillingCacheStub{balance: 5.0}, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, "") + require.NoError(t, err) + }) + + t.Run("并发等待后二次检查反映最新余额", func(t *testing.T) { + cache := &billInvBillingCacheStub{balance: 0.01} + svc := billInvNewBillingCacheService(t, cache, nil) + + // 第一次检查(获取并发槽前):余额尚存 → 放行 + require.NoError(t, svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, "")) + + // 等待期间其他请求把余额扣到 0 → 等待结束后的二次检查必须拒绝 + cache.balance = 0 + err := svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, "") + require.ErrorIs(t, err, ErrInsufficientBalance) + }) +} + +// TestBillingInvariant_PreflightSubscriptionLimits 锁定订阅模式 preflight 语义: +// 日/周/月任一窗口用量达到分组限额即拒绝;订阅非 active 或已过期拒绝。 +func TestBillingInvariant_PreflightSubscriptionLimits(t *testing.T) { + user := &User{ID: 601} + subscription := &UserSubscription{ID: 42} + group := &Group{ + ID: 7, + SubscriptionType: SubscriptionTypeSubscription, + DailyLimitUSD: billInvF64Ptr(10), + WeeklyLimitUSD: billInvF64Ptr(50), + MonthlyLimitUSD: billInvF64Ptr(100), + } + activeFuture := time.Now().Add(24 * time.Hour) + + tests := []struct { + name string + sub *SubscriptionCacheData + wantErr error + }{ + { + name: "限额内放行", + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: activeFuture, DailyUsage: 9.99, WeeklyUsage: 49.99, MonthlyUsage: 99.99}, + wantErr: nil, + }, + { + name: "日限额达到拒绝", + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: activeFuture, DailyUsage: 10}, + wantErr: ErrDailyLimitExceeded, + }, + { + name: "周限额达到拒绝", + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: activeFuture, WeeklyUsage: 50}, + wantErr: ErrWeeklyLimitExceeded, + }, + { + name: "月限额达到拒绝", + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: activeFuture, MonthlyUsage: 100}, + wantErr: ErrMonthlyLimitExceeded, + }, + { + name: "订阅过期拒绝", + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: time.Now().Add(-time.Minute)}, + wantErr: ErrSubscriptionInvalid, + }, + { + name: "订阅非active拒绝", + sub: &SubscriptionCacheData{Status: "cancelled", ExpiresAt: activeFuture}, + wantErr: ErrSubscriptionInvalid, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := billInvNewBillingCacheService(t, &billInvBillingCacheStub{sub: tt.sub}, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, group, subscription, "") + if tt.wantErr == nil { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, tt.wantErr) + } + }) + } +} + +// TestBillingInvariant_PreflightUserPlatformQuota 锁定 user×platform 配额的 +// preflight 语义:余额模式下日限额耗尽 → 429 拒绝;订阅模式豁免该检查。 +func TestBillingInvariant_PreflightUserPlatformQuota(t *testing.T) { + user := &User{ID: 601} + + t.Run("日配额耗尽拒绝", func(t *testing.T) { + cache := &billInvBillingCacheStub{ + balance: 5.0, // 余额充足,确保拒绝来自 platform quota + quotaEntry: billInvQuotaEntryV1(5.0, 5.0), + } + svc := billInvNewBillingCacheService(t, cache, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, PlatformAnthropic) + require.ErrorIs(t, err, ErrUserPlatformDailyQuotaExhausted) + }) + + t.Run("日配额未满放行", func(t *testing.T) { + cache := &billInvBillingCacheStub{ + balance: 5.0, + quotaEntry: billInvQuotaEntryV1(5.0, 4.99), + } + svc := billInvNewBillingCacheService(t, cache, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, nil, nil, PlatformAnthropic) + require.NoError(t, err) + }) + + t.Run("订阅模式豁免platform配额检查", func(t *testing.T) { + group := &Group{ + ID: 7, + SubscriptionType: SubscriptionTypeSubscription, + DailyLimitUSD: billInvF64Ptr(10), + } + cache := &billInvBillingCacheStub{ + sub: &SubscriptionCacheData{Status: SubscriptionStatusActive, ExpiresAt: time.Now().Add(24 * time.Hour)}, + quotaEntry: billInvQuotaEntryV1(5.0, 999), // platform 配额早已超限 + } + svc := billInvNewBillingCacheService(t, cache, nil) + err := svc.CheckBillingEligibility(context.Background(), user, nil, group, &UserSubscription{ID: 42}, PlatformAnthropic) + require.NoError(t, err, "订阅模式下 user×platform 配额不应生效") + }) +} + +// TestBillingInvariant_PreflightSimpleModeBypass 锁定 simple 运行模式跳过所有 +// 计费检查(余额为 0 也放行)。 +func TestBillingInvariant_PreflightSimpleModeBypass(t *testing.T) { + cfg := &config.Config{RunMode: config.RunModeSimple} + svc := billInvNewBillingCacheService(t, &billInvBillingCacheStub{balance: 0}, cfg) + err := svc.CheckBillingEligibility(context.Background(), &User{ID: 601}, nil, nil, nil, PlatformAnthropic) + require.NoError(t, err) +} diff --git a/backend/internal/service/gateway_forward_benchmark_test.go b/backend/internal/service/gateway_forward_benchmark_test.go new file mode 100644 index 0000000000..b1e96e172d --- /dev/null +++ b/backend/internal/service/gateway_forward_benchmark_test.go @@ -0,0 +1,83 @@ +//go:build unit + +// Phase-0 TASK-005 热路径基准:覆盖 anthropic API Key 透传账号的完整 +// Forward 路径(非流式 + 流式 SSE)。与 scripts/bench-baseline.sh 配套使用, +// 基线存于 testdata/bench/baseline.txt;对比策略:allocs/op 严格(不允许增加)、 +// ns/op 宽松(容忍 15% 抖动)。复用 TASK-002 特征化测试的夹具(passChar* 前缀)。 +package service + +import ( + "context" + "io" + "log" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" +) + +func benchFwdRun(b *testing.B, header http.Header, upstreamBody string, reqBody string, stream bool) { + b.Helper() + gin.SetMode(gin.TestMode) + // 透传分支每请求有一条 stdlib log,静音以降低基准 I/O 噪声 + // (日志成本在 baseline 与 compare 两侧同等消除,不影响对比有效性)。 + prevLogOut := log.Writer() + log.SetOutput(io.Discard) + b.Cleanup(func() { log.SetOutput(prevLogOut) }) + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + account := passCharAnthropicAccount(9100) + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: header.Clone(), + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }, + } + svc := passCharGatewayService(cfg, upstream) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(reqBody)), + Model: "claude-sonnet-4-20250514", + Stream: stream, + } + if _, err := svc.Forward(context.Background(), c, account, parsed); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGatewayForward_AnthropicNonStreamPassthrough(b *testing.B) { + upstreamJSON := `{"id":"msg_bench_01","type":"message","content":[{"type":"text","text":"hello"}],` + + `"usage":{"input_tokens":12,"output_tokens":7,"cache_read_input_tokens":3}}` + benchFwdRun(b, + http.Header{"Content-Type": []string{"application/json"}}, + upstreamJSON, + `{"model":"claude-sonnet-4-20250514","messages":[{"role":"user","content":"hi"}]}`, + false, + ) +} + +func BenchmarkGatewayForward_AnthropicStreamPassthrough(b *testing.B) { + events := []string{ + "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_bench_02\",\"usage\":{\"input_tokens\":9,\"cached_tokens\":3}}}", + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hel\"}}", + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"lo\"}}", + "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":5}}", + "event: message_stop\ndata: {\"type\":\"message_stop\"}", + } + benchFwdRun(b, + http.Header{"Content-Type": []string{"text/event-stream"}}, + strings.Join(events, "\n\n")+"\n\n", + `{"model":"claude-sonnet-4-20250514","stream":true,"messages":[{"role":"user","content":"hi"}]}`, + true, + ) +} diff --git a/backend/internal/service/gateway_passthrough_characterization_test.go b/backend/internal/service/gateway_passthrough_characterization_test.go new file mode 100644 index 0000000000..a8703f06d1 --- /dev/null +++ b/backend/internal/service/gateway_passthrough_characterization_test.go @@ -0,0 +1,411 @@ +//go:build unit + +// 本文件是 Phase-0 TASK-002 的特征化测试(characterization tests)。 +// 它把"客户端最终收到的内容"固化为基线:固定 mock 上游,逐字节/逐事件断言 +// 状态码、响应体与响应头。任何后续重构若改变这些外部可观测行为,测试会立即失败。 +// +// 覆盖的不变量(见 .claude/plugin-refactor/INVARIANTS.md): +// - I-1.1 anthropic 非流式 2xx body/状态码逐字节透传(升级自 JSONEq 级断言) +// - I-1.4 响应 header 白名单过滤(默认白名单 / additional_allowed / force_remove) +// - I-1.7 上游 4xx/5xx 错误按 ErrorPassthroughService 规则透传或转换 +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/model" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// passCharSplitSSEEvents 按 SSE 事件边界(空行 "\n\n")分帧。 +// 特征化测试只按事件比对,不按网络 chunk 比对;帧首尾多余的换行 +// (合法的 SSE 分帧差异,如额外空行)会被剥离,但帧内内容保持原文。 +func passCharSplitSSEEvents(body string) []string { + frames := strings.Split(body, "\n\n") + out := make([]string, 0, len(frames)) + for _, f := range frames { + f = strings.Trim(f, "\r\n") + if f == "" { + continue + } + out = append(out, f) + } + return out +} + +// passCharAnthropicAccount 返回 anthropic API Key 透传账号夹具。 +func passCharAnthropicAccount(id int64) *Account { + return &Account{ + ID: id, + Name: "pass-char-anthropic", + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "upstream-anthropic-key", + "base_url": "https://api.anthropic.com", + }, + Extra: map[string]any{ + "anthropic_passthrough": true, + }, + Status: StatusActive, + Schedulable: true, + } +} + +// passCharGinContext 构造带 /v1/messages POST 请求的 gin 测试上下文。 +func passCharGinContext(t *testing.T, path string) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, path, nil) + return c, rec +} + +func passCharGatewayService(cfg *config.Config, upstream HTTPUpstream) *GatewayService { + return &GatewayService{ + cfg: cfg, + responseHeaderFilter: compileResponseHeaderFilter(cfg), + httpUpstream: upstream, + rateLimitService: &RateLimitService{}, + deferredService: &DeferredService{}, + } +} + +// TestGatewayCharacterization_AnthropicNonStreamPassthroughByteExact 固化 I-1.1: +// anthropic 透传账号非流式转发时,上游 2xx 响应体逐字节透传、状态码透传。 +// 上游 body 故意包含字段顺序、多余空白、Unicode、HTML 字符与未知字段, +// 任何 JSON 重新序列化都会破坏逐字节相等。 +func TestGatewayCharacterization_AnthropicNonStreamPassthroughByteExact(t *testing.T) { + upstreamJSON := "{\"id\":\"msg_char_01\", \"type\":\"message\",\n" + + "\"unknown_field\":{\"nested\":[1,2.50,\"<&>\"]},\n" + + "\"content\":[{\"type\":\"text\",\"text\":\"héllo ✓ &\"}],\n" + + "\"usage\":{\"input_tokens\":12,\"output_tokens\":7}}" + + for _, status := range []int{http.StatusOK, http.StatusAccepted} { + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: status, + Header: http.Header{ + "Content-Type": []string{"application/json; charset=utf-8"}, + "x-request-id": []string{"rid-char-nonstream"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamJSON)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := passCharGatewayService(cfg, upstream) + + c, rec := passCharGinContext(t, "/v1/messages") + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(`{"model":"claude-sonnet-4-20250514","messages":[{"role":"user","content":"hi"}]}`)), + Model: "claude-sonnet-4-20250514", + } + + result, err := svc.Forward(context.Background(), c, passCharAnthropicAccount(9001), parsed) + require.NoError(t, err) + require.NotNil(t, result) + + require.Equal(t, status, rec.Code, "上游 2xx 状态码必须透传") + require.Equal(t, upstreamJSON, rec.Body.String(), "上游 2xx body 必须逐字节透传") + require.Equal(t, "application/json; charset=utf-8", rec.Header().Get("Content-Type"), "Content-Type 必须透传") + require.Equal(t, "rid-char-nonstream", rec.Header().Get("x-request-id"), "x-request-id 必须透传") + require.Equal(t, 12, result.Usage.InputTokens) + require.Equal(t, 7, result.Usage.OutputTokens) + } +} + +// TestGatewayCharacterization_AnthropicStreamPassthroughEventSequence 固化 I-1.1/I-1.2 流式侧: +// anthropic 透传账号流式转发时,客户端收到的 SSE 事件序列(按 \n\n 分帧)与上游完全一致, +// 含 event: 行、usage 事件与 cached_tokens 兼容字段,且网关不改写事件内容。 +func TestGatewayCharacterization_AnthropicStreamPassthroughEventSequence(t *testing.T) { + upstreamEvents := []string{ + "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_char_02\",\"usage\":{\"input_tokens\":9,\"cached_tokens\":3}}}", + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hel\"}}", + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"lo ✓\"}}", + "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":5}}", + "event: message_stop\ndata: {\"type\":\"message_stop\"}", + } + upstreamSSE := strings.Join(upstreamEvents, "\n\n") + "\n\n" + + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "x-request-id": []string{"rid-char-stream"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := passCharGatewayService(cfg, upstream) + + c, rec := passCharGinContext(t, "/v1/messages") + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(`{"model":"claude-sonnet-4-20250514","stream":true,"messages":[{"role":"user","content":"hi"}]}`)), + Model: "claude-sonnet-4-20250514", + Stream: true, + } + + result, err := svc.Forward(context.Background(), c, passCharAnthropicAccount(9002), parsed) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Stream) + + gotEvents := passCharSplitSSEEvents(rec.Body.String()) + require.Equal(t, upstreamEvents, gotEvents, "客户端收到的 SSE 事件序列必须与上游逐事件一致") + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-stream", rec.Header().Get("x-request-id")) + require.Equal(t, "no-cache", rec.Header().Get("Cache-Control")) + require.Equal(t, "no", rec.Header().Get("X-Accel-Buffering")) + + require.Equal(t, 9, result.Usage.InputTokens) + require.Equal(t, 5, result.Usage.OutputTokens) + require.Equal(t, 3, result.Usage.CacheReadInputTokens, "cached_tokens 应被解析进计费 usage") +} + +// TestGatewayCharacterization_ResponseHeaderFilter 固化 I-1.4: +// 响应 header 白名单过滤语义。 +// - 关闭(security.response_headers.enabled=false,默认):只用内置白名单, +// additional_allowed / force_remove 均不生效; +// - 开启:additional_allowed 追加放行,force_remove 即使是默认白名单键也强制移除。 +func TestGatewayCharacterization_ResponseHeaderFilter(t *testing.T) { + upstreamHeaders := func() http.Header { + return http.Header{ + "Content-Type": []string{"application/json"}, + "Cache-Control": []string{"no-store"}, + "x-request-id": []string{"rid-char-headers"}, + "X-Ratelimit-Remaining-Requests": []string{"99"}, + "Retry-After": []string{"17"}, + "Www-Authenticate": []string{"Bearer realm=api"}, + "Set-Cookie": []string{"secret=upstream"}, + "X-Custom-Upstream": []string{"custom-value"}, + } + } + + tests := []struct { + name string + headersCfg config.ResponseHeaderConfig + wantPresent map[string]string + wantAbsentKeys []string + }{ + { + name: "默认关闭_只用内置白名单且additional与force_remove不生效", + headersCfg: config.ResponseHeaderConfig{ + Enabled: false, + AdditionalAllowed: []string{"x-custom-upstream"}, // 关闭时不生效 + ForceRemove: []string{"x-request-id"}, // 关闭时不生效 + }, + wantPresent: map[string]string{ + "Content-Type": "application/json", + "Cache-Control": "no-store", + "x-request-id": "rid-char-headers", + "X-Ratelimit-Remaining-Requests": "99", + "Retry-After": "17", + "Www-Authenticate": "Bearer realm=api", + }, + wantAbsentKeys: []string{"Set-Cookie", "X-Custom-Upstream"}, + }, + { + name: "开启_additional追加放行且force_remove覆盖默认白名单", + headersCfg: config.ResponseHeaderConfig{ + Enabled: true, + AdditionalAllowed: []string{"x-custom-upstream"}, + ForceRemove: []string{"x-request-id"}, + }, + wantPresent: map[string]string{ + "Content-Type": "application/json", + "X-Custom-Upstream": "custom-value", + "Retry-After": "17", + }, + wantAbsentKeys: []string{"Set-Cookie", "x-request-id"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: upstreamHeaders(), + Body: io.NopCloser(strings.NewReader(`{"id":"msg_h","usage":{"input_tokens":1,"output_tokens":1}}`)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + cfg.Security.ResponseHeaders = tt.headersCfg + svc := passCharGatewayService(cfg, upstream) + + c, rec := passCharGinContext(t, "/v1/messages") + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(`{"model":"claude-sonnet-4-20250514","messages":[{"role":"user","content":"hi"}]}`)), + Model: "claude-sonnet-4-20250514", + } + + _, err := svc.Forward(context.Background(), c, passCharAnthropicAccount(9003), parsed) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code) + + for key, want := range tt.wantPresent { + require.Equal(t, want, rec.Header().Get(key), "响应头 %s 应保留", key) + } + for _, key := range tt.wantAbsentKeys { + require.Empty(t, rec.Header().Get(key), "响应头 %s 应被过滤", key) + } + }) + } +} + +func passCharIntPtr(v int) *int { return &v } +func passCharStrPtr(v string) *string { return &v } + +// passCharErrorPassthroughService 构造带固定规则的错误透传服务(不连缓存)。 +func passCharErrorPassthroughService(rules []*model.ErrorPassthroughRule) *ErrorPassthroughService { + return NewErrorPassthroughService(&mockErrorPassthroughRepo{rules: rules}, nil) +} + +// TestGatewayCharacterization_UpstreamErrorPassthroughRules 固化 I-1.7: +// 上游非 failover 错误(如 404)经 ErrorPassthroughService.MatchRule 决定透传或转换: +// - 规则命中 + passthrough_code/passthrough_body:透传上游状态码 + 提取的上游 message; +// - 规则命中 + 自定义 response_code/custom_message:使用自定义值; +// - 规则未命中:默认 502 + "Upstream request failed"。 +// +// 三种情况下错误体均为 {"type":"error","error":{"type":"upstream_error","message":...}}。 +func TestGatewayCharacterization_UpstreamErrorPassthroughRules(t *testing.T) { + passthroughRule := &model.ErrorPassthroughRule{ + ID: 1, + Name: "404-keyword-passthrough", + Enabled: true, + Priority: 1, + ErrorCodes: []int{404}, + Keywords: []string{"quota_exceeded_marker"}, + MatchMode: model.MatchModeAll, + Platforms: []string{model.PlatformAnthropic}, + PassthroughCode: true, + PassthroughBody: true, + } + overrideRule := &model.ErrorPassthroughRule{ + ID: 2, + Name: "404-keyword-override", + Enabled: true, + Priority: 1, + ErrorCodes: []int{404}, + Keywords: []string{"quota_exceeded_marker"}, + MatchMode: model.MatchModeAll, + Platforms: []string{model.PlatformAnthropic}, + PassthroughCode: false, + ResponseCode: passCharIntPtr(http.StatusTooManyRequests), + PassthroughBody: false, + CustomMessage: passCharStrPtr("custom upstream busy"), + } + + upstreamHitBody := `{"type":"error","error":{"type":"not_found_error","message":"quota_exceeded_marker: model not found"}}` + upstreamMissBody := `{"type":"error","error":{"type":"not_found_error","message":"plain not found"}}` + + tests := []struct { + name string + rules []*model.ErrorPassthroughRule + upstreamBody string + wantStatus int + wantMessage string + }{ + { + name: "规则命中_透传上游状态码与消息", + rules: []*model.ErrorPassthroughRule{passthroughRule}, + upstreamBody: upstreamHitBody, + wantStatus: http.StatusNotFound, + wantMessage: "quota_exceeded_marker: model not found", + }, + { + name: "规则命中_自定义状态码与消息", + rules: []*model.ErrorPassthroughRule{overrideRule}, + upstreamBody: upstreamHitBody, + wantStatus: http.StatusTooManyRequests, + wantMessage: "custom upstream busy", + }, + { + name: "规则未命中_默认502转换", + rules: []*model.ErrorPassthroughRule{passthroughRule}, + upstreamBody: upstreamMissBody, + wantStatus: http.StatusBadGateway, + wantMessage: "Upstream request failed", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusNotFound, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(tt.upstreamBody)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := &GatewayService{ + cfg: cfg, + responseHeaderFilter: compileResponseHeaderFilter(cfg), + httpUpstream: upstream, + } + + c, rec := passCharGinContext(t, "/v1/messages") + BindErrorPassthroughService(c, passCharErrorPassthroughService(tt.rules)) + + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(`{"model":"claude-sonnet-4-20250514","messages":[{"role":"user","content":"hi"}]}`)), + Model: "claude-sonnet-4-20250514", + } + + result, err := svc.Forward(context.Background(), c, passCharAnthropicAccount(9004), parsed) + require.Error(t, err, "上游错误时 Forward 必须返回 error(响应已写给客户端)") + require.Nil(t, result) + + require.Equal(t, tt.wantStatus, rec.Code) + require.JSONEq(t, + `{"type":"error","error":{"type":"upstream_error","message":"`+tt.wantMessage+`"}}`, + rec.Body.String()) + }) + } +} + +// TestGatewayCharacterization_Upstream400BodyPassthrough 固化 handleErrorResponse 对 400 的 +// 特殊语义:未命中透传规则时,上游 400 响应体原样透传给客户端(状态码 400 + 原 body)。 +func TestGatewayCharacterization_Upstream400BodyPassthrough(t *testing.T) { + upstreamBody := `{"type":"error","error":{"type":"invalid_request_error","message":"max_tokens: required"}}` + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := &GatewayService{ + cfg: cfg, + responseHeaderFilter: compileResponseHeaderFilter(cfg), + httpUpstream: upstream, + } + + c, rec := passCharGinContext(t, "/v1/messages") + parsed := &ParsedRequest{ + Body: NewRequestBodyRef([]byte(`{"model":"claude-sonnet-4-20250514","messages":[{"role":"user","content":"hi"}]}`)), + Model: "claude-sonnet-4-20250514", + } + + result, err := svc.Forward(context.Background(), c, passCharAnthropicAccount(9005), parsed) + require.Error(t, err) + require.Nil(t, result) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Equal(t, upstreamBody, rec.Body.String(), "上游 400 body 应原样透传") +} diff --git a/backend/internal/service/openai_gateway_passthrough_characterization_test.go b/backend/internal/service/openai_gateway_passthrough_characterization_test.go new file mode 100644 index 0000000000..dd5e1ba4d4 --- /dev/null +++ b/backend/internal/service/openai_gateway_passthrough_characterization_test.go @@ -0,0 +1,181 @@ +//go:build unit + +// Phase-0 TASK-002 特征化测试:OpenAI 路径流式/非流式透传(INVARIANTS I-1.5)。 +// 固定 mock 上游,断言客户端最终收到的字节/SSE 事件序列/状态码/响应头。 +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +// passCharOpenAIAccount 返回 OpenAI API Key 自动透传账号夹具。 +func passCharOpenAIAccount(id int64) *Account { + return &Account{ + ID: id, + Name: "pass-char-openai", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-upstream-openai", + "base_url": "https://upstream.example.com", + }, + Extra: map[string]any{ + "use_responses_api": true, + "openai_passthrough": true, + }, + Status: StatusActive, + Schedulable: true, + } +} + +func passCharOpenAIService(cfg *config.Config, upstream HTTPUpstream) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: cfg, + responseHeaderFilter: compileResponseHeaderFilter(cfg), + httpUpstream: upstream, + } +} + +func passCharOpenAIGinContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + return c, rec +} + +// TestGatewayCharacterization_OpenAIStreamPassthrough 固化 I-1.5 流式侧: +// OpenAI 透传账号流式转发时,客户端收到的 SSE 内容与上游逐事件一致(按 \n\n 分帧), +// 包含 event: 行、preamble 事件(response.created)与终止事件(response.completed/[DONE])。 +func TestGatewayCharacterization_OpenAIStreamPassthrough(t *testing.T) { + upstreamEvents := []string{ + "event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_char_1\"}}", + "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"Hel\"}", + "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"lo ✓\"}", + "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_char_1\",\"usage\":{\"input_tokens\":11,\"output_tokens\":4,\"input_tokens_details\":{\"cached_tokens\":2}}}}", + "data: [DONE]", + } + upstreamSSE := strings.Join(upstreamEvents, "\n\n") + "\n\n" + + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "x-request-id": []string{"rid-char-openai-stream"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := passCharOpenAIService(cfg, upstream) + + c, rec := passCharOpenAIGinContext(t) + body := []byte(`{"model":"gpt-5","stream":true,"instructions":"be brief","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}]}`) + + result, err := svc.Forward(context.Background(), c, passCharOpenAIAccount(9101), body) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Stream) + + gotEvents := passCharSplitSSEEvents(rec.Body.String()) + require.Equal(t, upstreamEvents, gotEvents, "客户端收到的 SSE 事件序列必须与上游逐事件一致") + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-openai-stream", rec.Header().Get("x-request-id")) + require.Equal(t, "no-cache", rec.Header().Get("Cache-Control")) + + // 透传模式下请求侧只替换认证:上游应收到原始 body + Bearer 上游 key,不残留入站鉴权。 + require.NotNil(t, upstream.lastReq) + require.Equal(t, string(body), string(upstream.lastBody), "透传模式上游请求体应与客户端原始 body 一致") + require.Equal(t, "Bearer sk-upstream-openai", upstream.lastReq.Header.Get("Authorization")) + + require.Equal(t, 11, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.OutputTokens) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) +} + +// TestGatewayCharacterization_OpenAINonStreamPassthroughByteExact 固化 I-1.5 非流式侧: +// 上游 2xx JSON 响应体逐字节透传、状态码/Content-Type 透传,x-codex-* 配额头强制放行。 +func TestGatewayCharacterization_OpenAINonStreamPassthroughByteExact(t *testing.T) { + upstreamJSON := "{\"id\":\"resp_char_2\", \"object\":\"response\",\n" + + "\"output\":[{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"héllo <&> ✓\"}]}],\n" + + "\"usage\":{\"input_tokens\":21,\"output_tokens\":8,\"input_tokens_details\":{\"cached_tokens\":5}}}" + + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "x-request-id": []string{"rid-char-openai-nonstream"}, + "Set-Cookie": []string{"secret=upstream"}, + "X-Codex-Primary-Used-Percent": []string{"42.5"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamJSON)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := passCharOpenAIService(cfg, upstream) + + c, rec := passCharOpenAIGinContext(t) + body := []byte(`{"model":"gpt-5","stream":false,"instructions":"be brief","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}]}`) + + result, err := svc.Forward(context.Background(), c, passCharOpenAIAccount(9102), body) + require.NoError(t, err) + require.NotNil(t, result) + require.False(t, result.Stream) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, upstreamJSON, rec.Body.String(), "上游 2xx body 必须逐字节透传") + require.Equal(t, "application/json", rec.Header().Get("Content-Type")) + require.Equal(t, "rid-char-openai-nonstream", rec.Header().Get("x-request-id")) + require.Equal(t, "42.5", rec.Header().Get("X-Codex-Primary-Used-Percent"), "x-codex-* 配额头透传模式强制放行") + require.Empty(t, rec.Header().Get("Set-Cookie"), "Set-Cookie 应被响应头过滤移除") + + require.Equal(t, 21, result.Usage.InputTokens) + require.Equal(t, 8, result.Usage.OutputTokens) + require.Equal(t, 5, result.Usage.CacheReadInputTokens) +} + +// TestGatewayCharacterization_OpenAIPassthroughUpstreamErrorBodyVerbatim 固化 OpenAI 透传 +// 模式的错误语义:非容量类 4xx(如 400)保持原样代理——上游状态码 + 原始错误体透传, +// Forward 返回 error(供 handler 记日志,但响应已写完)。 +func TestGatewayCharacterization_OpenAIPassthroughUpstreamErrorBodyVerbatim(t *testing.T) { + upstreamErrJSON := `{"error":{"type":"invalid_request_error","message":"Unsupported parameter: 'foo'","param":"foo"}}` + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "x-request-id": []string{"rid-char-openai-err"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamErrJSON)), + }, + } + cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}} + svc := passCharOpenAIService(cfg, upstream) + + c, rec := passCharOpenAIGinContext(t) + body := []byte(`{"model":"gpt-5","stream":false,"instructions":"be brief","input":"hi","foo":1}`) + + result, err := svc.Forward(context.Background(), c, passCharOpenAIAccount(9103), body) + require.Error(t, err) + require.Nil(t, result) + + require.Equal(t, http.StatusBadRequest, rec.Code, "透传模式上游错误状态码原样代理") + require.Equal(t, upstreamErrJSON, rec.Body.String(), "透传模式上游错误体逐字节透传") + require.Equal(t, "application/json", rec.Header().Get("Content-Type")) +} diff --git a/backend/internal/service/scheduling_invariants_test.go b/backend/internal/service/scheduling_invariants_test.go new file mode 100644 index 0000000000..548fa82259 --- /dev/null +++ b/backend/internal/service/scheduling_invariants_test.go @@ -0,0 +1,480 @@ +//go:build unit + +// Phase-0 TASK-004 调度不变量测试(INVARIANTS I-4.1 / I-4.2 / I-4.3 / I-4.5)。 +// +// 本文件固化粘性会话与选号过滤的外部可观测行为: +// - I-4.1 绑定后二次请求命中同账号;session hash 输入优先级 +// (metadata.user_id → cacheable content → IP+UA+APIKeyID+system+messages); +// - I-4.2 粘性绑定 TTL = 1 小时(常量引用 + 硬编码双断言 + miniredis 过期语义); +// - I-4.3 sticky_escape:粘性账号不可用时能逃逸并重选成功(只断言硬语义, +// 不断言触发阈值细节); +// - I-4.5 账号级模型映射白名单(model_mapping)与渠道模型映射/定价限制对选号的影响。 +// +// 断言纪律:不断言选号的具体排序结果(负载感知排序属可演化策略), +// 只断言"命中/逃逸/排除/拒绝"等硬语义。 +// +// 复用同包既有夹具:mockAccountRepoForPlatform / mockGroupRepoForGateway / +// mockConcurrencyCache(gateway_multiplatform_test.go)、newTestChannelService / +// makeStandardRepo(channel_service_test.go)、mustParseSessionHashRequest / +// anthropicSessionBody / msg(generate_session_hash_test.go)。 +// 本文件新增的包级辅助类型/函数一律带 schedInv 前缀。 +package service + +import ( + "context" + "errors" + "fmt" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +// --------------------------------------------------------------------------- +// schedInv 夹具 +// --------------------------------------------------------------------------- + +// schedInvSetCall 记录一次 SetSessionAccountID / RefreshSessionTTL 调用。 +type schedInvSetCall struct { + groupID int64 + hash string + accountID int64 + ttl time.Duration +} + +// schedInvGatewayCache 是 GatewayCache 的内存实现,记录 TTL 与删除调用。 +type schedInvGatewayCache struct { + mu sync.Mutex + bindings map[string]int64 + setCalls []schedInvSetCall + refreshCalls []schedInvSetCall + deleteCalls []string +} + +var _ GatewayCache = (*schedInvGatewayCache)(nil) + +func schedInvNewGatewayCache() *schedInvGatewayCache { + return &schedInvGatewayCache{bindings: make(map[string]int64)} +} + +func schedInvCacheKey(groupID int64, hash string) string { + return fmt.Sprintf("%d:%s", groupID, hash) +} + +func (c *schedInvGatewayCache) GetSessionAccountID(_ context.Context, groupID int64, sessionHash string) (int64, error) { + c.mu.Lock() + defer c.mu.Unlock() + if id, ok := c.bindings[schedInvCacheKey(groupID, sessionHash)]; ok { + return id, nil + } + return 0, errors.New("schedInv: session binding not found") +} + +func (c *schedInvGatewayCache) SetSessionAccountID(_ context.Context, groupID int64, sessionHash string, accountID int64, ttl time.Duration) error { + c.mu.Lock() + defer c.mu.Unlock() + c.bindings[schedInvCacheKey(groupID, sessionHash)] = accountID + c.setCalls = append(c.setCalls, schedInvSetCall{groupID: groupID, hash: sessionHash, accountID: accountID, ttl: ttl}) + return nil +} + +func (c *schedInvGatewayCache) RefreshSessionTTL(_ context.Context, groupID int64, sessionHash string, ttl time.Duration) error { + c.mu.Lock() + defer c.mu.Unlock() + c.refreshCalls = append(c.refreshCalls, schedInvSetCall{groupID: groupID, hash: sessionHash, ttl: ttl}) + return nil +} + +func (c *schedInvGatewayCache) DeleteSessionAccountID(_ context.Context, groupID int64, sessionHash string) error { + c.mu.Lock() + defer c.mu.Unlock() + delete(c.bindings, schedInvCacheKey(groupID, sessionHash)) + c.deleteCalls = append(c.deleteCalls, schedInvCacheKey(groupID, sessionHash)) + return nil +} + +func (c *schedInvGatewayCache) binding(groupID int64, hash string) int64 { + c.mu.Lock() + defer c.mu.Unlock() + return c.bindings[schedInvCacheKey(groupID, hash)] +} + +// schedInvRedisGatewayCache 是 miniredis 后端的 GatewayCache, +// 与 repository/gateway_cache.go 的键语义一致,用于验证 TTL 过期行为。 +type schedInvRedisGatewayCache struct { + rdb *redis.Client +} + +var _ GatewayCache = (*schedInvRedisGatewayCache)(nil) + +func (c *schedInvRedisGatewayCache) key(groupID int64, hash string) string { + return fmt.Sprintf("sticky_session:%d:%s", groupID, hash) +} + +func (c *schedInvRedisGatewayCache) GetSessionAccountID(ctx context.Context, groupID int64, sessionHash string) (int64, error) { + return c.rdb.Get(ctx, c.key(groupID, sessionHash)).Int64() +} + +func (c *schedInvRedisGatewayCache) SetSessionAccountID(ctx context.Context, groupID int64, sessionHash string, accountID int64, ttl time.Duration) error { + return c.rdb.Set(ctx, c.key(groupID, sessionHash), accountID, ttl).Err() +} + +func (c *schedInvRedisGatewayCache) RefreshSessionTTL(ctx context.Context, groupID int64, sessionHash string, ttl time.Duration) error { + return c.rdb.Expire(ctx, c.key(groupID, sessionHash), ttl).Err() +} + +func (c *schedInvRedisGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error { + return c.rdb.Del(ctx, c.key(groupID, sessionHash)).Err() +} + +// schedInvAccountRepo 构造账号 repo mock(复用 mockAccountRepoForPlatform)。 +func schedInvAccountRepo(accounts ...Account) *mockAccountRepoForPlatform { + repo := &mockAccountRepoForPlatform{ + accounts: accounts, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + return repo +} + +// schedInvLoadAwareService 构造启用负载感知调度(Layer1.5/Layer2)的 GatewayService。 +func schedInvLoadAwareService(repo *mockAccountRepoForPlatform, cache GatewayCache) *GatewayService { + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + return &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } +} + +func schedInvAnthropicAccount(id int64, priority int) Account { + return Account{ + ID: id, + Name: fmt.Sprintf("sched-inv-%d", id), + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Priority: priority, + Concurrency: 5, + Status: StatusActive, + Schedulable: true, + } +} + +// --------------------------------------------------------------------------- +// I-4.1 session hash 输入优先级 +// --------------------------------------------------------------------------- + +// TestSchedulingInvariant_SessionHashSourcePriority 固化 session hash 三级输入优先级: +// metadata.user_id 的 session_xxx > cache_control(ephemeral) 内容 > +// IP+UA+APIKeyID+system+messages 完整摘要兜底。 +func TestSchedulingInvariant_SessionHashSourcePriority(t *testing.T) { + svc := &GatewayService{} + sessionCtx := &SessionContext{ClientIP: "10.1.2.3", UserAgent: "claude-cli/2.0.0", APIKeyID: 77} + metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000" + cacheableSystem := []any{map[string]any{ + "type": "text", + "text": "long cacheable system prompt", + "cache_control": map[string]any{"type": "ephemeral"}, + }} + messages := []any{msg("user", "hello invariants")} + + // 1. 三个来源齐备 → metadata session_id 胜出 + full := mustParseSessionHashRequest(t, anthropicSessionBody(cacheableSystem, messages, metadata), sessionCtx) + require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", svc.GenerateSessionHash(full), + "metadata.user_id 的 session_id 必须有最高优先级") + + // 2. 去掉 metadata → cacheable content hash 胜出(与 SessionContext 无关) + cacheableOnly := mustParseSessionHashRequest(t, anthropicSessionBody(cacheableSystem, messages, ""), sessionCtx) + cacheableNoCtx := mustParseSessionHashRequest(t, anthropicSessionBody(cacheableSystem, messages, ""), nil) + h2 := svc.GenerateSessionHash(cacheableOnly) + require.NotEmpty(t, h2) + require.Equal(t, svc.GenerateSessionHash(cacheableNoCtx), h2, + "存在 cacheable content 时 SessionContext 不参与 hash(第 2 级优先于第 3 级)") + + // 3. 无 metadata 且无 cacheable content → 兜底摘要,SessionContext 参与区分 + fallbackA := mustParseSessionHashRequest(t, anthropicSessionBody("plain system", messages, ""), sessionCtx) + fallbackB := mustParseSessionHashRequest(t, anthropicSessionBody("plain system", messages, ""), + &SessionContext{ClientIP: "10.9.9.9", UserAgent: "other-agent/1.0", APIKeyID: 88}) + h3a := svc.GenerateSessionHash(fallbackA) + h3b := svc.GenerateSessionHash(fallbackB) + require.NotEmpty(t, h3a) + require.NotEqual(t, h3a, h3b, "兜底级 hash 必须混入 IP+UA+APIKeyID 区分因子") + require.NotEqual(t, h2, h3a, "第 2 级与第 3 级来源应产生不同 hash") +} + +// --------------------------------------------------------------------------- +// I-4.1 绑定后二次请求命中同账号 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_StickySession_SecondSelectionHitsSameAccount(t *testing.T) { + ctx := context.Background() + repo := schedInvAccountRepo( + schedInvAnthropicAccount(1, 1), + schedInvAnthropicAccount(2, 1), + schedInvAnthropicAccount(3, 1), + ) + cache := schedInvNewGatewayCache() + svc := schedInvLoadAwareService(repo, cache) + + const sessionHash = "sched-inv-sticky-hit" + + // 第一次选号:无绑定 → 任选一个账号并写入粘性绑定 + first, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-sonnet-4-5", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, first) + require.NotNil(t, first.Account) + require.True(t, first.Acquired) + boundID := cache.binding(0, sessionHash) + require.Equal(t, first.Account.ID, boundID, "首次选号成功后必须建立粘性绑定") + + // 第二次选号:同 sessionHash 必须命中同一账号(无论负载排序如何演化) + second, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-sonnet-4-5", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, second.Account) + require.Equal(t, first.Account.ID, second.Account.ID, "二次请求必须命中粘性绑定的同一账号") +} + +// --------------------------------------------------------------------------- +// I-4.2 粘性绑定 TTL = 1 小时 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_StickySession_TTLIsOneHour(t *testing.T) { + // 双断言①:常量引用(常量被改动时此处失败) + require.Equal(t, time.Hour, stickySessionTTL, "粘性会话 TTL 常量必须保持 1 小时") + + t.Run("BindStickySession传递TTL", func(t *testing.T) { + cache := schedInvNewGatewayCache() + svc := &GatewayService{cache: cache} + require.NoError(t, svc.BindStickySession(context.Background(), nil, "sess-ttl", 42)) + require.Len(t, cache.setCalls, 1) + // 双断言②:硬编码当前值 + 常量引用 + require.Equal(t, time.Hour, cache.setCalls[0].ttl) + require.Equal(t, stickySessionTTL, cache.setCalls[0].ttl) + }) + + t.Run("粘性命中刷新TTL为1小时", func(t *testing.T) { + ctx := context.Background() + repo := schedInvAccountRepo(schedInvAnthropicAccount(1, 1), schedInvAnthropicAccount(2, 1)) + cache := schedInvNewGatewayCache() + require.NoError(t, cache.SetSessionAccountID(ctx, 0, "sess-refresh", 1, stickySessionTTL)) + svc := schedInvLoadAwareService(repo, cache) + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-refresh", "claude-sonnet-4-5", nil, "", 0) + require.NoError(t, err) + require.Equal(t, int64(1), result.Account.ID, "应命中粘性账号") + require.NotEmpty(t, cache.refreshCalls, "粘性命中后必须刷新 TTL") + require.Equal(t, time.Hour, cache.refreshCalls[0].ttl) + }) + + t.Run("miniredis过期语义", func(t *testing.T) { + ctx := context.Background() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + svc := &GatewayService{cache: &schedInvRedisGatewayCache{rdb: rdb}} + + require.NoError(t, svc.BindStickySession(ctx, nil, "sess-expiry", 7)) + + got, err := svc.GetCachedSessionAccountID(ctx, nil, "sess-expiry") + require.NoError(t, err) + require.Equal(t, int64(7), got) + + // TTL 内(59 分钟后)仍然命中 + mr.FastForward(stickySessionTTL - time.Minute) + got, err = svc.GetCachedSessionAccountID(ctx, nil, "sess-expiry") + require.NoError(t, err) + require.Equal(t, int64(7), got) + + // 超过 TTL 后绑定消失(当前实现返回 0 + 底层 miss error,handler 忽略 error) + mr.FastForward(2 * time.Minute) + got, _ = svc.GetCachedSessionAccountID(ctx, nil, "sess-expiry") + require.Equal(t, int64(0), got, "TTL 过期后粘性绑定必须失效") + }) +} + +// --------------------------------------------------------------------------- +// I-4.3 sticky_escape:粘性账号不可用时逃逸重选 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_StickyEscape_ReselectsWhenBoundAccountUnavailable(t *testing.T) { + ctx := context.Background() + const model = "claude-sonnet-4-5" + + t.Run("绑定账号不可调度_逃逸并重绑", func(t *testing.T) { + // 账号 1 状态 error → 不在可调度列表;粘性绑定仍指向它 + broken := schedInvAnthropicAccount(1, 1) + broken.Status = StatusError + healthy := schedInvAnthropicAccount(2, 1) + repo := schedInvAccountRepo(broken, healthy) + cache := schedInvNewGatewayCache() + require.NoError(t, cache.SetSessionAccountID(ctx, 0, "sess-escape-a", 1, time.Hour)) + svc := schedInvLoadAwareService(repo, cache) + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-escape-a", model, nil, "", 0) + require.NoError(t, err, "粘性账号不可用时必须能逃逸重选,而不是失败") + require.NotNil(t, result.Account) + require.NotEqual(t, int64(1), result.Account.ID, "不得再选中不可用的粘性账号") + require.Equal(t, result.Account.ID, cache.binding(0, "sess-escape-a"), + "逃逸成功后粘性绑定更新为新账号(当前实现行为)") + }) + + t.Run("绑定账号模型限流_清除绑定并逃逸", func(t *testing.T) { + limited := schedInvAnthropicAccount(1, 1) + limited.Extra = map[string]any{ + "model_rate_limits": map[string]any{ + model: map[string]any{ + "rate_limit_reset_at": time.Now().Add(30 * time.Minute).Format(time.RFC3339), + }, + }, + } + healthy := schedInvAnthropicAccount(2, 1) + repo := schedInvAccountRepo(limited, healthy) + cache := schedInvNewGatewayCache() + require.NoError(t, cache.SetSessionAccountID(ctx, 0, "sess-escape-b", 1, time.Hour)) + svc := schedInvLoadAwareService(repo, cache) + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-escape-b", model, nil, "", 0) + require.NoError(t, err) + require.Equal(t, int64(2), result.Account.ID, "逃逸后应重选到可用账号") + require.NotEmpty(t, cache.deleteCalls, "模型限流的粘性绑定应被清除(当前实现行为)") + }) + + t.Run("绑定账号在failover排除集合_跳过粘性", func(t *testing.T) { + repo := schedInvAccountRepo(schedInvAnthropicAccount(1, 1), schedInvAnthropicAccount(2, 1)) + cache := schedInvNewGatewayCache() + require.NoError(t, cache.SetSessionAccountID(ctx, 0, "sess-escape-c", 1, time.Hour)) + svc := schedInvLoadAwareService(repo, cache) + + excluded := map[int64]struct{}{1: {}} + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-escape-c", model, excluded, "", 0) + require.NoError(t, err) + require.Equal(t, int64(2), result.Account.ID, "排除集合中的粘性账号必须被跳过") + }) +} + +// --------------------------------------------------------------------------- +// I-4.5 账号级模型映射白名单对选号的影响 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_AccountModelMappingFiltersSelection(t *testing.T) { + ctx := context.Background() + + t.Run("不支持请求模型的账号被排除", func(t *testing.T) { + // 账号 1 配了 model_mapping 白名单(只支持 opus),优先级更高; + // 账号 2 未配映射(支持所有模型)。请求 sonnet 必须落到账号 2。 + restricted := schedInvAnthropicAccount(1, 0) + restricted.Credentials = map[string]any{ + "model_mapping": map[string]any{"claude-opus-4-6": "claude-opus-4-6"}, + } + open := schedInvAnthropicAccount(2, 9) + repo := schedInvAccountRepo(restricted, open) + svc := schedInvLoadAwareService(repo, schedInvNewGatewayCache()) + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-sonnet-4-5", nil, "", 0) + require.NoError(t, err) + require.Equal(t, int64(2), result.Account.ID, + "配置了 model_mapping 白名单且不含请求模型的账号必须被排除(即使优先级更高)") + }) + + t.Run("支持模型的账号正常命中映射", func(t *testing.T) { + restricted := schedInvAnthropicAccount(1, 0) + restricted.Credentials = map[string]any{ + "model_mapping": map[string]any{"claude-opus-4-6": "claude-opus-4-6"}, + } + repo := schedInvAccountRepo(restricted) + svc := schedInvLoadAwareService(repo, schedInvNewGatewayCache()) + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-opus-4-6", nil, "", 0) + require.NoError(t, err) + require.Equal(t, int64(1), result.Account.ID) + }) + + t.Run("所有账号都不支持模型_返回无可用账号", func(t *testing.T) { + a := schedInvAnthropicAccount(1, 1) + a.Credentials = map[string]any{"model_mapping": map[string]any{"claude-opus-4-6": "claude-opus-4-6"}} + b := schedInvAnthropicAccount(2, 1) + b.Credentials = map[string]any{"model_mapping": map[string]any{"claude-haiku-4-5": "claude-haiku-4-5"}} + repo := schedInvAccountRepo(a, b) + svc := schedInvLoadAwareService(repo, schedInvNewGatewayCache()) + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-sonnet-4-5", nil, "", 0) + require.Error(t, err) + require.ErrorIs(t, err, ErrNoAvailableAccounts) + }) +} + +// --------------------------------------------------------------------------- +// I-4.5 渠道模型映射 + 定价限制对选号的影响 +// --------------------------------------------------------------------------- + +func TestSchedulingInvariant_ChannelMappingPricingRestrictionAffectsSelection(t *testing.T) { + const groupID = int64(10) + newSvc := func(pricingModels []string) (*GatewayService, *mockGroupRepoForGateway) { + ch := Channel{ + ID: 1, + Status: StatusActive, + GroupIDs: []int64{groupID}, + RestrictModels: true, + BillingModelSource: BillingModelSourceChannelMapped, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformAnthropic, Models: pricingModels}, + }, + ModelMapping: map[string]map[string]string{ + PlatformAnthropic: {"claude-sonnet-4-5": "claude-sonnet-4-6"}, + }, + } + channelSvc := newTestChannelService(makeStandardRepo(ch, map[int64]string{groupID: PlatformAnthropic})) + repo := schedInvAccountRepo(func() Account { + a := schedInvAnthropicAccount(1, 1) + a.AccountGroups = []AccountGroup{{AccountID: 1, GroupID: groupID}} + return a + }()) + groupRepo := &mockGroupRepoForGateway{ + groups: map[int64]*Group{ + groupID: {ID: groupID, Platform: PlatformAnthropic, Status: StatusActive, Hydrated: true}, + }, + } + svc := &GatewayService{ + accountRepo: repo, + groupRepo: groupRepo, + channelService: channelSvc, + cache: schedInvNewGatewayCache(), + cfg: testConfig(), + } + return svc, groupRepo + } + ctxWithGroup := func(groupRepo *mockGroupRepoForGateway) context.Context { + return context.WithValue(context.Background(), ctxkey.Group, groupRepo.groups[groupID]) + } + + t.Run("映射目标不在渠道定价_选号被拒绝", func(t *testing.T) { + // 渠道映射 sonnet-4-5 → sonnet-4-6,但定价列表只有 opus: + // 限制检查对象是"映射后模型",因此请求 sonnet-4-5 必须被拒绝。 + gid := groupID + svc, groupRepo := newSvc([]string{"claude-opus-4-6"}) + _, err := svc.SelectAccountWithLoadAwareness(ctxWithGroup(groupRepo), &gid, "", "claude-sonnet-4-5", nil, "", 0) + require.Error(t, err) + require.ErrorIs(t, err, ErrNoAvailableAccounts) + require.Contains(t, err.Error(), "channel pricing restriction", + "渠道定价限制必须以专属错误语义拒绝选号") + }) + + t.Run("映射目标在渠道定价_正常选号", func(t *testing.T) { + // 定价列表含映射后的 sonnet-4-6(即使不含原始请求模型 sonnet-4-5)→ 放行。 + // 这证明渠道映射结果(而非原始模型名)决定选号阶段的限制判定。 + gid := groupID + svc, groupRepo := newSvc([]string{"claude-sonnet-4-6"}) + result, err := svc.SelectAccountWithLoadAwareness(ctxWithGroup(groupRepo), &gid, "", "claude-sonnet-4-5", nil, "", 0) + require.NoError(t, err) + require.Equal(t, int64(1), result.Account.ID) + }) +} diff --git a/backend/scripts/bench-baseline.sh b/backend/scripts/bench-baseline.sh new file mode 100644 index 0000000000..5784507f94 --- /dev/null +++ b/backend/scripts/bench-baseline.sh @@ -0,0 +1,102 @@ +#!/usr/bin/env bash +# Phase-0 TASK-005 热路径基准采集与对比(插件化改造回归安全网)。 +# +# 用法: +# ./scripts/bench-baseline.sh collect # 采集基线 -> testdata/bench/baseline.txt +# ./scripts/bench-baseline.sh compare # 当前代码 vs 基线,违反阈值时退出码 1 +# +# 对比策略(每个基准按多次采样的均值比较): +# - allocs/op 严格:增加 > ALLOC_TOLERANCE(默认 0)即失败 +# - ns/op 宽松:劣化 > TIME_TOLERANCE_PCT%(默认 15,吸收 runner 抖动) +# 环境变量:BENCH_COUNT(采样次数,默认 6)、TIME_TOLERANCE_PCT、ALLOC_TOLERANCE +# +# 注意:基线与对比应在同类机器上进行;换机器/换 Go 版本后请重新 collect +# (环境信息记录在 testdata/bench/baseline.env.txt)。 + +set -euo pipefail +cd "$(dirname "$0")/.." + +PATTERN='BenchmarkGatewayForward|BenchmarkGatewayService_ParseSSEUsage|BenchmarkParseClaudeUsageFromResponseBody' +PKGS='./internal/service/' +BASELINE=testdata/bench/baseline.txt +COUNT="${BENCH_COUNT:-6}" +TIME_TOLERANCE_PCT="${TIME_TOLERANCE_PCT:-15}" +ALLOC_TOLERANCE="${ALLOC_TOLERANCE:-0}" + +run_bench() { + go test -tags=unit -run '^$' -bench "$PATTERN" -benchmem -count="$COUNT" "$PKGS" +} + +# 从基准输出中提取每个基准的均值:name mean_ns mean_allocs +# ns 与 allocs 独立计数;缺任一指标(如未带 -benchmem)的基准跳过并警告, +# 避免除零产生 inf/nan 静默污染对比结果。 +summarize() { + awk '/^Benchmark/ { + name=$1; sub(/-[0-9]+$/, "", name) + for (i=2; i<=NF; i++) { + if ($(i)=="ns/op") { ns[name]+=$(i-1); nns[name]++ } + if ($(i)=="allocs/op") { al[name]+=$(i-1); nal[name]++ } + } + } + END { + for (k in nns) { + if (nns[k]>0 && nal[k]>0) printf "%s %.2f %.2f\n", k, ns[k]/nns[k], al[k]/nal[k] + else print "warn: skip " k " (missing ns/op or allocs/op samples, run with -benchmem)" > "/dev/stderr" + } + }' "$1" | sort +} + +case "${1:-}" in + collect) + mkdir -p "$(dirname "$BASELINE")" + { + echo "go: $(go version)" + echo "cpu: $(grep -m1 'model name' /proc/cpuinfo 2>/dev/null | cut -d: -f2- | xargs || echo unknown)" + echo "date: $(date -u +%Y-%m-%dT%H:%M:%SZ)" + echo "count: $COUNT" + } > "$(dirname "$BASELINE")/baseline.env.txt" + run_bench | tee "$BASELINE" + echo "基线已写入 $BASELINE" + ;; + compare) + [ -f "$BASELINE" ] || { echo "缺少基线 $BASELINE,先运行: $0 collect" >&2; exit 2; } + new="$(mktemp)" + run_bench | tee "$new" + echo + if command -v benchstat >/dev/null 2>&1; then + benchstat "$BASELINE" "$new" || true + echo + fi + base_sum="$(mktemp)"; new_sum="$(mktemp)" + summarize "$BASELINE" > "$base_sum" + summarize "$new" > "$new_sum" + # 基准名集合对比(审计 F-1):join 只输出两侧共有的基准,增删/改名若不检查 + # 会被静默丢弃,造成"基准消失但 compare 仍绿"的假阴性。 + added="$(join -v2 "$base_sum" "$new_sum" | awk '{print $1}')" + missing="$(join -v1 "$base_sum" "$new_sum" | awk '{print $1}')" + table_status=0 + join "$base_sum" "$new_sum" | awk -v tp="$TIME_TOLERANCE_PCT" -v at="$ALLOC_TOLERANCE" ' + { + name=$1; bns=$2; bal=$3; nns=$4; nal=$5 + dns = (bns>0) ? (nns-bns)/bns*100 : 0 + dal = nal-bal + status="ok" + if (dal > at) { status="FAIL(allocs/op +" dal ")"; fail=1 } + else if (dns > tp) { status="FAIL(ns/op +" sprintf("%.1f", dns) "%)"; fail=1 } + printf "%-60s ns/op %10.0f -> %10.0f (%+6.1f%%) allocs/op %7.1f -> %7.1f %s\n", name, bns, nns, dns, bal, nal, status + } + END { exit fail ? 1 : 0 }' || table_status=$? + if [ -n "$added" ]; then + printf 'warn: 以下基准在基线中不存在,未参与对比(新增基准请重新 collect 基线):\n%s\n' "$added" >&2 + fi + if [ -n "$missing" ]; then + printf 'FAIL: 基线中的以下基准在本次运行缺失(被删除/改名?回归网已失效,须修复或重新 collect):\n%s\n' "$missing" >&2 + exit 1 + fi + exit "$table_status" + ;; + *) + echo "用法: $0 {collect|compare}" >&2 + exit 2 + ;; +esac diff --git a/backend/testdata/bench/baseline.env.txt b/backend/testdata/bench/baseline.env.txt new file mode 100644 index 0000000000..6ad24c63b6 --- /dev/null +++ b/backend/testdata/bench/baseline.env.txt @@ -0,0 +1,4 @@ +go: go version go1.26.4 linux/amd64 +cpu: 13th Gen Intel(R) Core(TM) i5-13600KF +date: 2026-06-11T06:26:18Z +count: 6 diff --git a/backend/testdata/bench/baseline.txt b/backend/testdata/bench/baseline.txt new file mode 100644 index 0000000000..2614c90256 --- /dev/null +++ b/backend/testdata/bench/baseline.txt @@ -0,0 +1,48 @@ +goos: linux +goarch: amd64 +pkg: github.com/Wei-Shaw/sub2api/internal/service +cpu: 13th Gen Intel(R) Core(TM) i5-13600KF +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 403580 2613 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 467905 2583 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 462894 2551 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 443625 2538 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 473317 2544 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageStart-20 465820 2481 ns/op 2224 B/op 40 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1746464 692.2 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1731408 677.6 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1786676 673.6 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1765081 675.2 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1762540 687.4 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageStart-20 1737454 679.1 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 470481 2281 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 512163 2243 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 540775 2273 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 509505 2292 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 510970 2300 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsage_MessageDelta-20 509253 2325 ns/op 1872 B/op 37 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1560142 769.5 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1574744 765.5 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1569315 767.5 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1561520 769.2 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1551661 801.5 ns/op 0 B/op 0 allocs/op +BenchmarkGatewayService_ParseSSEUsagePassthrough_MessageDelta-20 1555836 769.8 ns/op 0 B/op 0 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1430790 850.6 ns/op 320 B/op 2 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1431716 846.4 ns/op 320 B/op 2 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1427368 842.8 ns/op 320 B/op 2 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1414472 842.4 ns/op 320 B/op 2 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1410212 838.6 ns/op 320 B/op 2 allocs/op +BenchmarkParseClaudeUsageFromResponseBody-20 1424356 840.8 ns/op 320 B/op 2 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 120924 9941 ns/op 14667 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 117312 10316 ns/op 14667 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 113154 10088 ns/op 14666 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 116223 10061 ns/op 14667 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 119020 10349 ns/op 14667 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicNonStreamPassthrough-20 112460 10227 ns/op 14667 B/op 93 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 39451 29660 ns/op 19156 B/op 134 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 41238 29418 ns/op 19189 B/op 134 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 40820 29667 ns/op 19218 B/op 134 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 41295 29620 ns/op 19186 B/op 134 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 40642 29393 ns/op 19171 B/op 134 allocs/op +BenchmarkGatewayForward_AnthropicStreamPassthrough-20 40014 29780 ns/op 19176 B/op 134 allocs/op +PASS +ok github.com/Wei-Shaw/sub2api/internal/service 66.644s