Merge origin/main into codex/secure-protected-video-content-4498

Resolve conflicts with main:
- service/grok_media.go: keep both post-response blocks — #4497's empty
  image-output failover (main) runs first for image endpoints, then this
  branch's video-status content-URL rewrite; the endpoint conditions are
  mutually exclusive
- handler/openai_gateway_credential_failover_loop_test.go: mark the stub
  OAuth accounts media-eligible via the grok_media_eligible extra override,
  because this branch moved grok media failover coverage to the generation
  endpoint, which is now gated by #4540's paid-eligibility probe on main
This commit is contained in:
shaw
2026-07-18 21:31:28 +08:00
222 changed files with 11544 additions and 1414 deletions
+19 -5
View File
@@ -18,7 +18,9 @@ ARG NPM_CONFIG_REGISTRY=
# -----------------------------------------------------------------------------
# Stage 1: Frontend Builder
# -----------------------------------------------------------------------------
FROM ${NODE_IMAGE} AS frontend-builder
# --platform=$BUILDPLATFORM: the frontend output is JS (arch-neutral), so build
# it on the native host arch instead of under QEMU emulation for the target.
FROM --platform=${BUILDPLATFORM} ${NODE_IMAGE} AS frontend-builder
ARG NPM_CONFIG_REGISTRY
WORKDIR /app/frontend
@@ -44,7 +46,11 @@ RUN pnpm run build
# -----------------------------------------------------------------------------
# Stage 2: Backend Builder
# -----------------------------------------------------------------------------
FROM ${GOLANG_IMAGE} AS backend-builder
# --platform=$BUILDPLATFORM: run the Go toolchain on the native host arch and
# cross-compile to the target arch below. The binary is CGO_ENABLED=0, so this
# is a clean pure-Go cross-compile — no QEMU emulation of go mod download / go
# build (emulated networking here was dropping module fetches with EOF).
FROM --platform=${BUILDPLATFORM} ${GOLANG_IMAGE} AS backend-builder
# Build arguments for version info (set by CI)
ARG VERSION=
@@ -52,6 +58,9 @@ ARG COMMIT=docker
ARG DATE
ARG GOPROXY
ARG GOSUMDB
# Populated by buildx from the --platform target (e.g. linux/amd64).
ARG TARGETOS
ARG TARGETARCH
ENV GOPROXY=${GOPROXY}
ENV GOSUMDB=${GOSUMDB}
@@ -63,7 +72,10 @@ WORKDIR /app/backend
# Copy go mod files first (better caching)
COPY backend/go.mod backend/go.sum ./
RUN go mod download
# Cache mount keeps the module cache across builds so a transient CDN blip on
# retry resumes instead of re-fetching every zip from scratch.
RUN --mount=type=cache,id=sub2api-gomod,target=/go/pkg/mod \
go mod download
# Copy backend source first
COPY backend/ ./
@@ -73,10 +85,12 @@ COPY --from=frontend-builder /app/backend/internal/web/dist ./internal/web/dist
# Build the binary (BuildType=release for CI builds, embed frontend)
# Version precedence: build arg VERSION > exact git tag > cmd/server/VERSION
RUN VERSION_VALUE="${VERSION}" && \
RUN --mount=type=cache,id=sub2api-gomod,target=/go/pkg/mod \
--mount=type=cache,id=sub2api-gobuild,target=/root/.cache/go-build \
VERSION_VALUE="${VERSION}" && \
if [ -z "${VERSION_VALUE}" ]; then VERSION_VALUE="$(./scripts/resolve-version.sh)"; fi && \
DATE_VALUE="${DATE:-$(date -u +%Y-%m-%dT%H:%M:%SZ)}" && \
CGO_ENABLED=0 GOOS=linux go build \
CGO_ENABLED=0 GOOS=${TARGETOS:-linux} GOARCH=${TARGETARCH} go build \
-tags embed \
-ldflags="-s -w -X main.Version=${VERSION_VALUE} -X main.Commit=${COMMIT} -X main.Date=${DATE_VALUE} -X main.BuildType=release" \
-trimpath \
+16 -4
View File
@@ -639,8 +639,20 @@ override this limit.
The connection cap is coordinated through Redis using a 60-second lease that
is refreshed every 20 seconds. A process that cannot confirm a lease for a
full lease lifetime closes its local WebSocket rather than continuing outside
the global cap. Use `http_bridge` for client-WebSocket/upstream-HTTP operation
when rolling out or mitigating upstream WebSocket issues.
the global cap.
Enable the v2 mode router before selecting an account-level WS mode such as
`http_bridge`:
```yaml
gateway:
openai_ws:
mode_router_v2_enabled: true
```
Or set `GATEWAY_OPENAI_WS_MODE_ROUTER_V2_ENABLED=true` in the environment.
Use `http_bridge` for client-WebSocket/upstream-HTTP operation when rolling out
or mitigating upstream WebSocket issues.
#### ⚠️ Important: Creating the Admin Account
@@ -788,9 +800,9 @@ xAI quota is passive. Sub2API does not invent subscription quota values; it reco
`401` responses temporarily remove accounts with invalid credentials from scheduling. `403` responses are treated as access or entitlement failures instead of token-refresh loops. `429` responses use `Retry-After` or a short cooldown to temporarily remove the account from scheduling.
New Grok image and video generation requests use a media-specific eligibility check. An OAuth account is excluded from new media generation when its recorded weekly or monthly billing probe returns `403`; chat requests and video status lookups are not affected by this media-only quarantine. If no eligible account remains, the media endpoint returns HTTP `503` with error type `grok_media_no_eligible_account` instead of forwarding the request to a known-ineligible account.
New Grok image and video generation requests use a media-specific eligibility check. API-key accounts remain eligible. OAuth accounts require positive paid-entitlement evidence from the xAI billing probe; Free, forbidden, missing, malformed, and inconclusive billing observations are excluded from new media generation. Unobserved OAuth accounts are probed before the first media request is forwarded, and imports run the billing-first quota probe proactively. Chat requests and video status lookups are not affected by this media-only quarantine. If no eligible account remains, the media endpoint returns HTTP `503` with error type `grok_media_no_eligible_account`.
Administrators can override automatic media eligibility through the account create/update API by setting `extra.grok_media_eligible` to `false` (exclude) or `true` (force eligible). On update, set it to `null` to remove the override and return to automatic probe-based behavior; omitting the field preserves the current override. A missing billing observation does not block legacy routing, and a weekly allowance period by itself is not treated as evidence that the account is ineligible.
Administrators can override automatic media eligibility through the account create/update API by setting `extra.grok_media_eligible` to `false` (exclude) or `true` (force eligible). On update, set it to `null` to remove the override and return to automatic probe-based behavior; omitting the field preserves the current override. A weekly allowance period alone is not treated as a paid tier signal. Successful image responses must contain at least one actual image output; empty HTTP `200` responses trigger account failover instead of being counted and returned as successful generations.
---
@@ -0,0 +1,20 @@
# Ingress rejection log cleanup
This maintenance command removes historical admission rejections from
`ops_error_logs` without matching unrelated authentication or upstream errors.
It is a dry run unless `--execute` is supplied, and always requires an explicit
RFC3339 cutoff.
```sh
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z --execute
```
Run the execute form only after every application instance has been upgraded so
older instances cannot add new ingress rejection rows below the chosen cutoff.
The classifier intentionally retains invariant failures such as
`USER_NOT_FOUND`, database errors, quota/billing errors, and upstream failures.
After the rollout and cleanup are verified, run
`backend/scripts/finalize-ingress-reject-cleanup.sql` in a maintenance window to
remove the deprecated plaintext-key audit table and attribution columns.
@@ -0,0 +1,218 @@
package main
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"flag"
"fmt"
"log"
"sort"
"strings"
"time"
_ "github.com/Wei-Shaw/sub2api/ent/runtime"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/repository"
"github.com/lib/pq"
)
const classifierVersion = "ingress-reject-v1"
type candidate struct {
id int64
statusCode int
message string
body string
}
func main() {
beforeRaw := flag.String("before", "", "required RFC3339 cutoff; only older rows are considered")
execute := flag.Bool("execute", false, "delete matched rows (default is dry-run)")
batchSize := flag.Int("batch-size", 5000, "scan/delete batch size (1-5000)")
flag.Parse()
if *beforeRaw == "" {
log.Fatal("--before is required")
}
before, err := time.Parse(time.RFC3339, *beforeRaw)
if err != nil {
log.Fatalf("invalid --before: %v", err)
}
if *batchSize < 1 || *batchSize > 5000 {
log.Fatal("--batch-size must be between 1 and 5000")
}
cfg, err := config.LoadForBootstrap()
if err != nil {
log.Fatalf("load config: %v", err)
}
client, db, err := repository.InitEnt(cfg)
if err != nil {
log.Fatalf("initialize database: %v", err)
}
defer func() { _ = client.Close() }()
ctx := context.Background()
counts, scanned, matched, deleted, err := cleanup(ctx, db, before, *batchSize, *execute)
if err != nil {
log.Fatalf("cleanup failed: %v", err)
}
digest := sha256.Sum256([]byte(classifierVersion))
mode := "dry-run"
if *execute {
mode = "execute"
}
fmt.Printf("mode=%s before=%s classifier=%s scanned=%d matched=%d deleted=%d\n",
mode, before.UTC().Format(time.RFC3339), hex.EncodeToString(digest[:]), scanned, matched, deleted)
reasons := make([]string, 0, len(counts))
for reason := range counts {
reasons = append(reasons, reason)
}
sort.Strings(reasons)
for _, reason := range reasons {
fmt.Printf("reason=%s count=%d\n", reason, counts[reason])
}
if *execute && deleted > 0 {
fmt.Println("cleanup complete; schedule VACUUM (ANALYZE) ops_error_logs during normal maintenance")
}
}
func cleanup(ctx context.Context, db *sql.DB, before time.Time, batchSize int, execute bool) (map[string]int64, int64, int64, int64, error) {
counts := make(map[string]int64)
var cursor, scanned, matched, deleted int64
for {
rows, err := db.QueryContext(ctx, `
SELECT id, COALESCE(status_code, 0), COALESCE(error_message, ''), COALESCE(error_body, '')
FROM ops_error_logs
WHERE id > $1
AND created_at < $2
AND error_phase = 'auth'
AND account_id IS NULL
AND upstream_status_code IS NULL
AND COALESCE(upstream_error_message, '') = ''
AND COALESCE(upstream_error_detail, '') = ''
ORDER BY id ASC
LIMIT $3`, cursor, before, batchSize)
if err != nil {
return nil, scanned, matched, deleted, err
}
batch := make([]candidate, 0, batchSize)
for rows.Next() {
var item candidate
if err := rows.Scan(&item.id, &item.statusCode, &item.message, &item.body); err != nil {
_ = rows.Close()
return nil, scanned, matched, deleted, err
}
batch = append(batch, item)
cursor = item.id
}
if err := rows.Err(); err != nil {
_ = rows.Close()
return nil, scanned, matched, deleted, err
}
_ = rows.Close()
if len(batch) == 0 {
break
}
ids := make([]int64, 0, len(batch))
for _, item := range batch {
scanned++
if reason, ok := historicalIngressRejectReason(item); ok {
matched++
counts[reason]++
ids = append(ids, item.id)
}
}
if execute && len(ids) > 0 {
result, err := db.ExecContext(ctx,
`DELETE FROM ops_error_logs WHERE id = ANY($1) AND created_at < $2`, pq.Array(ids), before)
if err != nil {
return nil, scanned, matched, deleted, err
}
n, err := result.RowsAffected()
if err != nil {
return nil, scanned, matched, deleted, err
}
deleted += n
}
}
return counts, scanned, matched, deleted, nil
}
func historicalIngressRejectReason(item candidate) (string, bool) {
code, message := parseErrorIdentity(item.body, item.message)
switch code {
case "API_KEY_REQUIRED":
return "missing_key", true
case "INVALID_API_KEY":
return "invalid_key", true
case "API_KEY_DISABLED":
return "key_disabled", true
case "USER_INACTIVE":
return "user_inactive", true
case "GROUP_DELETED":
return "group_deleted", true
case "GROUP_DISABLED":
return "group_disabled", true
case "GROUP_NOT_ALLOWED":
return "group_forbidden", true
case "ACCESS_DENIED":
return "ip_acl_denied", true
case "api_key_in_query_deprecated":
return "query_key_deprecated", true
}
normalized := strings.TrimSpace(message)
switch {
case normalized == "API key is required":
return "missing_key", true
case normalized == "Invalid API key":
return "invalid_key", true
case normalized == "API key is disabled":
return "key_disabled", true
case normalized == "User account is not active":
return "user_inactive", true
case normalized == "API Key 所属分组已删除":
return "group_deleted", true
case normalized == "API Key 所属分组已停用":
return "group_disabled", true
case normalized == "API Key 所属专属分组不再允许当前用户使用":
return "group_forbidden", true
case normalized == "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.":
return "group_unassigned", true
case strings.HasPrefix(normalized, "Access denied. Your IP is "):
return "ip_acl_denied", true
case normalized == "Query parameter api_key is deprecated. Use Authorization header or key instead.":
return "query_key_deprecated", true
default:
return "", false
}
}
func parseErrorIdentity(body, fallbackMessage string) (string, string) {
var payload struct {
Code string `json:"code"`
Message string `json:"message"`
Error struct {
Code json.RawMessage `json:"code"`
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal([]byte(body), &payload); err != nil {
return "", fallbackMessage
}
message := payload.Message
if message == "" {
message = payload.Error.Message
}
if message == "" {
message = fallbackMessage
}
return strings.TrimSpace(payload.Code), message
}
@@ -0,0 +1,29 @@
package main
import "testing"
func TestHistoricalIngressRejectReason(t *testing.T) {
tests := []struct {
name string
item candidate
reason string
match bool
}{
{name: "standard invalid key", item: candidate{body: `{"code":"INVALID_API_KEY","message":"Invalid API key"}`}, reason: "invalid_key", match: true},
{name: "google missing key", item: candidate{body: `{"error":{"code":401,"message":"API key is required","status":"UNAUTHENTICATED"}}`}, reason: "missing_key", match: true},
{name: "google group deleted", item: candidate{body: `{"error":{"code":403,"message":"API Key 所属分组已删除","status":"PERMISSION_DENIED"}}`}, reason: "group_deleted", match: true},
{name: "ip acl", item: candidate{body: `{"code":"ACCESS_DENIED","message":"Access denied. Your IP is 192.0.2.1"}`}, reason: "ip_acl_denied", match: true},
{name: "user not found remains", item: candidate{body: `{"code":"USER_NOT_FOUND","message":"User associated with API key not found"}`}, match: false},
{name: "quota remains", item: candidate{body: `{"code":"API_KEY_QUOTA_EXHAUSTED","message":"quota"}`}, match: false},
{name: "database failure remains", item: candidate{statusCode: 500, message: "Failed to validate API key", body: `{"code":"INTERNAL_ERROR","message":"Failed to validate API key"}`}, match: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
reason, ok := historicalIngressRejectReason(tt.item)
if ok != tt.match || reason != tt.reason {
t.Fatalf("got (%q, %v), want (%q, %v)", reason, ok, tt.reason, tt.match)
}
})
}
}
+28
View File
@@ -81,6 +81,10 @@ func provideCleanup(
opsCleanup *service.OpsCleanupService,
opsScheduledReport *service.OpsScheduledReportService,
opsSystemLogSink *service.OpsSystemLogSink,
opsService *service.OpsService,
opsIngressReject *service.OpsIngressRejectAggregator,
apiKeyService *service.APIKeyService,
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
schedulerSnapshot *service.SchedulerSnapshotService,
tokenRefresh *service.TokenRefreshService,
accountExpiry *service.AccountExpiryService,
@@ -121,6 +125,30 @@ func provideCleanup(
// 应用层清理步骤可并行执行,基础设施资源(Redis/Ent)最后按顺序关闭。
parallelSteps := []cleanupStep{
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
}
return nil
}},
{"AuthCacheInvalidationWorker", func() error {
if authCacheInvalidationWorker != nil {
authCacheInvalidationWorker.Stop()
}
return nil
}},
{"AuthCacheInvalidationSubscriber", func() error {
if apiKeyService != nil {
apiKeyService.StopAuthCacheInvalidationSubscriber()
}
return nil
}},
{"OpsRuntimeSettingsRefresh", func() error {
if opsService != nil {
opsService.StopRuntimeSettingsRefresh()
}
return nil
}},
{"PromptAuditService", func() error {
if promptAudit != nil {
return promptAudit.Shutdown(ctx)
+34 -3
View File
@@ -157,7 +157,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
antigravityGatewayService := service.NewAntigravityGatewayService(accountRepository, gatewayCache, schedulerSnapshotService, antigravityTokenProvider, rateLimitService, httpUpstream, settingService, internal500CounterCache)
geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
authCacheInvalidationOutboxRepository := repository.NewAuthCacheInvalidationOutboxRepository(db)
authCacheInvalidationWorker := service.ProvideAuthCacheInvalidationWorker(authCacheInvalidationOutboxRepository, apiKeyCache, apiKeyService)
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService, authCacheInvalidationWorker, apiKeyService)
usageHandler := handler.NewUsageHandler(usageService, apiKeyService, opsService, settingService)
redeemHandler := handler.NewRedeemHandler(redeemService)
subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
@@ -270,7 +272,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
legacyEngine := securityaudit.NewLegacyModerationAdapter(contentModerationService)
coordinator := securityaudit.NewCoordinator(legacyEngine, promptService)
gatewayHandler := handler.ProvideGatewayHandler(gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService, coordinator)
openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig, coordinator)
openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, grokQuotaService, configConfig, coordinator)
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService)
totpHandler := handler.NewTotpHandler(totpService)
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService)
@@ -306,6 +308,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
opsAlertEvaluatorService := service.ProvideOpsAlertEvaluatorService(opsService, opsRepository, emailService, redisClient, configConfig, proxyRepository)
opsCleanupService := service.ProvideOpsCleanupService(opsRepository, db, redisClient, configConfig, channelMonitorService, settingRepository, opsService)
opsScheduledReportService := service.ProvideOpsScheduledReportService(opsService, userService, emailService, redisClient, configConfig)
opsIngressRejectAggregator := service.ProvideOpsIngressRejectAggregator(opsRepository, opsService)
accountExpiryService := service.ProvideAccountExpiryService(accountRepository)
proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository)
subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db)
@@ -314,7 +317,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db)
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
application := &Application{
Server: httpServer,
PromptAudit: promptService,
@@ -351,6 +354,10 @@ func provideCleanup(
opsCleanup *service.OpsCleanupService,
opsScheduledReport *service.OpsScheduledReportService,
opsSystemLogSink *service.OpsSystemLogSink,
opsService *service.OpsService,
opsIngressReject *service.OpsIngressRejectAggregator,
apiKeyService *service.APIKeyService,
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
schedulerSnapshot *service.SchedulerSnapshotService,
tokenRefresh *service.TokenRefreshService,
accountExpiry *service.AccountExpiryService,
@@ -390,6 +397,30 @@ func provideCleanup(
}
parallelSteps := []cleanupStep{
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
}
return nil
}},
{"AuthCacheInvalidationWorker", func() error {
if authCacheInvalidationWorker != nil {
authCacheInvalidationWorker.Stop()
}
return nil
}},
{"AuthCacheInvalidationSubscriber", func() error {
if apiKeyService != nil {
apiKeyService.StopAuthCacheInvalidationSubscriber()
}
return nil
}},
{"OpsRuntimeSettingsRefresh", func() error {
if opsService != nil {
opsService.StopRuntimeSettingsRefresh()
}
return nil
}},
{"PromptAuditService", func() error {
if promptAudit != nil {
return promptAudit.Shutdown(ctx)
+4
View File
@@ -58,6 +58,10 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
&service.OpsCleanupService{},
&service.OpsScheduledReportService{},
opsSystemLogSinkSvc,
nil, // opsService
nil, // opsIngressRejectAggregator
nil, // apiKeyService
nil, // authCacheInvalidationWorker
schedulerSnapshotSvc,
tokenRefreshSvc,
accountExpirySvc,
+75 -8
View File
@@ -645,6 +645,7 @@ type ServerConfig struct {
EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header
FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接
ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒)
MaxHeaderBytes int `mapstructure:"max_header_bytes"` // 请求头最大字节数(HTTP/2 映射为 header-list 上限)
IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒)
TrustedProxies []string `mapstructure:"trusted_proxies"` // 可信代理列表(CIDR/IP)
MaxRequestBodySize int64 `mapstructure:"max_request_body_size"` // 全局最大请求体限制
@@ -796,6 +797,8 @@ type GatewayConfig struct {
OpenAIHighEffortFirstOutputTimeoutSeconds int `mapstructure:"openai_high_effort_first_output_timeout_seconds"`
// 请求体最大字节数,用于网关请求体大小限制
MaxBodySize int64 `mapstructure:"max_body_size"`
// TextMaxBodySize limits endpoints that cannot carry inline image/video payloads.
TextMaxBodySize int64 `mapstructure:"text_max_body_size"`
// 非流式上游响应体读取上限(字节),用于防止无界读取导致内存放大
UpstreamResponseReadMaxBytes int64 `mapstructure:"upstream_response_read_max_bytes"`
// 代理探测响应体读取上限(字节)
@@ -1419,12 +1422,22 @@ type RateLimitConfig struct {
// APIKeyAuthCacheConfig API Key 认证缓存配置
type APIKeyAuthCacheConfig struct {
L1Size int `mapstructure:"l1_size"`
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
JitterPercent int `mapstructure:"jitter_percent"`
Singleflight bool `mapstructure:"singleflight"`
L1Size int `mapstructure:"l1_size"`
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
JitterPercent int `mapstructure:"jitter_percent"`
Singleflight bool `mapstructure:"singleflight"`
LookupConcurrency int `mapstructure:"lookup_concurrency"`
InvalidAbuse InvalidAuthAbuseConfig `mapstructure:"invalid_abuse"`
}
type InvalidAuthAbuseConfig struct {
Enabled bool `mapstructure:"enabled"`
Threshold int `mapstructure:"threshold"`
WindowSeconds int `mapstructure:"window_seconds"`
BlockSeconds int `mapstructure:"block_seconds"`
Capacity int `mapstructure:"capacity"`
}
// SubscriptionCacheConfig 订阅认证 L1 缓存配置
@@ -1698,8 +1711,9 @@ func setDefaults() {
viper.SetDefault("server.mode", "release")
viper.SetDefault("server.enable_server_timing", false)
viper.SetDefault("server.frontend_url", "")
viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
viper.SetDefault("server.read_header_timeout", 10) // 10秒读取请求头
viper.SetDefault("server.max_header_bytes", 64*1024)
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
viper.SetDefault("server.trusted_proxies", []string{})
viper.SetDefault("server.max_request_body_size", int64(256*1024*1024))
// H2C 默认配置
@@ -1983,6 +1997,12 @@ func setDefaults() {
viper.SetDefault("api_key_auth_cache.negative_ttl_seconds", 30)
viper.SetDefault("api_key_auth_cache.jitter_percent", 10)
viper.SetDefault("api_key_auth_cache.singleflight", true)
viper.SetDefault("api_key_auth_cache.lookup_concurrency", 64)
viper.SetDefault("api_key_auth_cache.invalid_abuse.enabled", true)
viper.SetDefault("api_key_auth_cache.invalid_abuse.threshold", 120)
viper.SetDefault("api_key_auth_cache.invalid_abuse.window_seconds", 60)
viper.SetDefault("api_key_auth_cache.invalid_abuse.block_seconds", 60)
viper.SetDefault("api_key_auth_cache.invalid_abuse.capacity", 16384)
// Subscription auth L1 cache
viper.SetDefault("subscription_cache.l1_size", 16384)
@@ -2111,6 +2131,7 @@ func setDefaults() {
viper.SetDefault("gateway.antigravity_fallback_cooldown_minutes", 1)
viper.SetDefault("gateway.antigravity_extra_retries", 10)
viper.SetDefault("gateway.max_body_size", int64(256*1024*1024))
viper.SetDefault("gateway.text_max_body_size", int64(32*1024*1024))
viper.SetDefault("gateway.upstream_response_read_max_bytes", DefaultUpstreamResponseReadMaxBytes)
viper.SetDefault("gateway.proxy_probe_response_read_max_bytes", int64(1024*1024))
viper.SetDefault("gateway.gemini_debug_response_headers", false)
@@ -2208,6 +2229,49 @@ func setDefaults() {
}
func (c *Config) Validate() error {
if c.Server.ReadHeaderTimeout < 1 || c.Server.ReadHeaderTimeout > 60 {
return fmt.Errorf("server.read_header_timeout must be between 1 and 60 seconds")
}
if c.Server.MaxHeaderBytes < 8*1024 || c.Server.MaxHeaderBytes > 1024*1024 {
return fmt.Errorf("server.max_header_bytes must be between 8192 and 1048576 bytes")
}
if c.Server.IdleTimeout <= 0 {
return fmt.Errorf("server.idle_timeout must be positive")
}
if c.Server.MaxRequestBodySize < 0 {
return fmt.Errorf("server.max_request_body_size must be non-negative")
}
if c.Server.H2C.Enabled {
if c.Server.H2C.MaxConcurrentStreams == 0 {
return fmt.Errorf("server.h2c.max_concurrent_streams must be positive")
}
if c.Server.H2C.IdleTimeout <= 0 {
return fmt.Errorf("server.h2c.idle_timeout must be positive")
}
if c.Server.H2C.MaxReadFrameSize < 16*1024 || c.Server.H2C.MaxReadFrameSize > 16*1024*1024-1 {
return fmt.Errorf("server.h2c.max_read_frame_size must be between 16384 and 16777215 bytes")
}
if c.Server.H2C.MaxUploadBufferPerConnection < 65535 {
return fmt.Errorf("server.h2c.max_upload_buffer_per_connection must be at least 65535 bytes")
}
if c.Server.H2C.MaxUploadBufferPerStream <= 0 {
return fmt.Errorf("server.h2c.max_upload_buffer_per_stream must be positive")
}
}
if c.APIKeyAuth.InvalidAbuse.Enabled {
if c.APIKeyAuth.InvalidAbuse.Threshold < 10 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.threshold must be at least 10")
}
if c.APIKeyAuth.InvalidAbuse.WindowSeconds < 1 || c.APIKeyAuth.InvalidAbuse.WindowSeconds > 3600 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.window_seconds must be between 1 and 3600")
}
if c.APIKeyAuth.InvalidAbuse.BlockSeconds < 1 || c.APIKeyAuth.InvalidAbuse.BlockSeconds > 3600 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.block_seconds must be between 1 and 3600")
}
if c.APIKeyAuth.InvalidAbuse.Capacity < 256 || c.APIKeyAuth.InvalidAbuse.Capacity > 1_000_000 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.capacity must be between 256 and 1000000")
}
}
jwtSecret := strings.TrimSpace(c.JWT.Secret)
if jwtSecret == "" {
return fmt.Errorf("jwt.secret is required")
@@ -2733,6 +2797,9 @@ func (c *Config) Validate() error {
if c.Gateway.MaxBodySize <= 0 {
return fmt.Errorf("gateway.max_body_size must be positive")
}
if c.Gateway.TextMaxBodySize <= 0 || c.Gateway.TextMaxBodySize > c.Gateway.MaxBodySize {
return fmt.Errorf("gateway.text_max_body_size must be positive and no greater than gateway.max_body_size")
}
if c.Gateway.UpstreamResponseReadMaxBytes <= 0 {
return fmt.Errorf("gateway.upstream_response_read_max_bytes must be positive")
}
+62
View File
@@ -35,6 +35,18 @@ func TestLoadServerTimingConfig(t *testing.T) {
})
}
func TestLoadHTTPIngressSafetyDefaults(t *testing.T) {
resetViperWithJWTSecret(t)
cfg, err := Load()
require.NoError(t, err)
require.Equal(t, 10, cfg.Server.ReadHeaderTimeout)
require.Equal(t, 64*1024, cfg.Server.MaxHeaderBytes)
require.Equal(t, int64(32*1024*1024), cfg.Gateway.TextMaxBodySize)
require.True(t, cfg.APIKeyAuth.InvalidAbuse.Enabled)
require.Equal(t, 120, cfg.APIKeyAuth.InvalidAbuse.Threshold)
require.Equal(t, 16384, cfg.APIKeyAuth.InvalidAbuse.Capacity)
}
func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) {
viper.Reset()
t.Setenv("JWT_SECRET", "")
@@ -1185,6 +1197,51 @@ func TestValidateConfigErrors(t *testing.T) {
mutate func(*Config)
wantErr string
}{
{
name: "server read header timeout",
mutate: func(c *Config) { c.Server.ReadHeaderTimeout = 0 },
wantErr: "server.read_header_timeout",
},
{
name: "server max header bytes too small",
mutate: func(c *Config) { c.Server.MaxHeaderBytes = 4096 },
wantErr: "server.max_header_bytes",
},
{
name: "server max request body size",
mutate: func(c *Config) { c.Server.MaxRequestBodySize = -1 },
wantErr: "server.max_request_body_size",
},
{
name: "h2c zero concurrent streams",
mutate: func(c *Config) {
c.Server.H2C.Enabled = true
c.Server.H2C.MaxConcurrentStreams = 0
},
wantErr: "server.h2c.max_concurrent_streams",
},
{
name: "h2c oversized read frame",
mutate: func(c *Config) {
c.Server.H2C.Enabled = true
c.Server.H2C.MaxReadFrameSize = 16 * 1024 * 1024
},
wantErr: "server.h2c.max_read_frame_size",
},
{
name: "invalid auth abuse threshold too small",
mutate: func(c *Config) {
c.APIKeyAuth.InvalidAbuse.Threshold = 9
},
wantErr: "api_key_auth_cache.invalid_abuse.threshold",
},
{
name: "invalid auth abuse capacity too small",
mutate: func(c *Config) {
c.APIKeyAuth.InvalidAbuse.Capacity = 255
},
wantErr: "api_key_auth_cache.invalid_abuse.capacity",
},
{
name: "jwt secret required",
mutate: func(c *Config) { c.JWT.Secret = "" },
@@ -1386,6 +1443,11 @@ func TestValidateConfigErrors(t *testing.T) {
mutate: func(c *Config) { c.Gateway.MaxBodySize = 0 },
wantErr: "gateway.max_body_size",
},
{
name: "gateway text body exceeds media body",
mutate: func(c *Config) { c.Gateway.TextMaxBodySize = c.Gateway.MaxBodySize + 1 },
wantErr: "gateway.text_max_body_size",
},
{
name: "gateway response header timeout",
mutate: func(c *Config) { c.Gateway.ResponseHeaderTimeout = -1 },
@@ -61,7 +61,7 @@ type AccountHandler struct {
sessionLimitCache service.SessionLimitCache
rpmCache service.RPMCache
tokenCacheInvalidator service.TokenCacheInvalidator
grokImportProber grokUsageProber
grokImportProber grokImportProber
upstreamBillingProbe *service.UpstreamBillingProbeService
}
@@ -121,6 +121,7 @@ type CreateAccountRequest struct {
GroupIDs []int64 `json:"group_ids"`
ExpiresAt *int64 `json:"expires_at"`
AutoPauseOnExpired *bool `json:"auto_pause_on_expired"`
ProbeEnabled *bool `json:"upstream_billing_probe_enabled"`
ConfirmMixedChannelRisk *bool `json:"confirm_mixed_channel_risk"` // 用户确认混合渠道风险
}
@@ -159,6 +160,7 @@ type BulkUpdateAccountsRequest struct {
GroupIDs *[]int64 `json:"group_ids"`
Credentials map[string]any `json:"credentials"`
Extra map[string]any `json:"extra"`
ProbeEnabled *bool `json:"upstream_billing_probe_enabled"`
ConfirmMixedChannelRisk *bool `json:"confirm_mixed_channel_risk"` // 用户确认混合渠道风险
}
@@ -828,6 +830,7 @@ func (h *AccountHandler) Create(c *gin.Context) {
GroupIDs: req.GroupIDs,
ExpiresAt: req.ExpiresAt,
AutoPauseOnExpired: req.AutoPauseOnExpired,
ProbeEnabled: req.ProbeEnabled,
SkipMixedChannelCheck: skipCheck,
})
if execErr != nil {
@@ -1907,7 +1910,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) {
req.Schedulable != nil ||
req.GroupIDs != nil ||
len(req.Credentials) > 0 ||
len(req.Extra) > 0
len(req.Extra) > 0 ||
req.ProbeEnabled != nil
if !hasUpdates {
response.BadRequest(c, "No updates provided")
@@ -1928,6 +1932,7 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) {
GroupIDs: req.GroupIDs,
Credentials: req.Credentials,
Extra: req.Extra,
ProbeEnabled: req.ProbeEnabled,
SkipMixedChannelCheck: skipCheck,
})
if err != nil {
@@ -222,3 +222,22 @@ func TestBulkUpdateAcceptsFilterTargetRequest(t *testing.T) {
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, float64(0), resp["code"])
}
func TestBulkUpdateAcceptsDedicatedUpstreamBillingProbeSetting(t *testing.T) {
adminSvc := newStubAdminService()
router := setupAccountMixedChannelRouter(adminSvc)
body, _ := json.Marshal(map[string]any{
"account_ids": []int64{1, 2},
"upstream_billing_probe_enabled": false,
})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.NotNil(t, adminSvc.lastBulkUpdateAccountInput)
require.NotNil(t, adminSvc.lastBulkUpdateAccountInput.ProbeEnabled)
require.False(t, *adminSvc.lastBulkUpdateAccountInput.ProbeEnabled)
}
@@ -33,6 +33,7 @@ type stubAdminService struct {
createSparkShadowErr error
updateAccountErr error
bulkUpdateAccountErr error
lastBulkUpdateAccountInput *service.BulkUpdateAccountsInput
getAccountResult *service.Account
updateAccountCalls int
updateAccountExtraCalls int
@@ -478,6 +479,7 @@ func (s *stubAdminService) SetAccountSchedulable(ctx context.Context, id int64,
}
func (s *stubAdminService) BulkUpdateAccounts(ctx context.Context, input *service.BulkUpdateAccountsInput) (*service.BulkUpdateAccountsResult, error) {
s.lastBulkUpdateAccountInput = input
if s.bulkUpdateAccountErr != nil {
return nil, s.bulkUpdateAccountErr
}
@@ -15,12 +15,12 @@ const (
grokImportProbeTimeout = 25 * time.Second
)
type grokUsageProber interface {
ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error)
type grokImportProber interface {
QueryQuota(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error)
}
type grokImportProbeTask struct {
prober grokUsageProber
prober grokImportProber
accountID int64
}
@@ -51,7 +51,7 @@ func newGrokImportProbeScheduler(concurrency int, timeout time.Duration) *grokIm
}
}
func (s *grokImportProbeScheduler) schedule(prober grokUsageProber, account *service.Account) {
func (s *grokImportProbeScheduler) schedule(prober grokImportProber, account *service.Account) {
if s == nil || prober == nil || account == nil || account.ID <= 0 {
return
}
@@ -97,7 +97,7 @@ func (s *grokImportProbeScheduler) nextTask() (grokImportProbeTask, bool) {
return task, true
}
func (s *grokImportProbeScheduler) run(prober grokUsageProber, accountID int64) {
func (s *grokImportProbeScheduler) run(prober grokImportProber, accountID int64) {
defer func() {
if recovered := recover(); recovered != nil {
slog.Error(
@@ -112,7 +112,7 @@ func (s *grokImportProbeScheduler) run(prober grokUsageProber, accountID int64)
// while this timeout only bounds the actual upstream probe execution.
ctx, cancel := context.WithTimeout(context.Background(), s.timeout)
defer cancel()
result, err := prober.ProbeUsage(ctx, accountID)
result, err := prober.QueryQuota(ctx, accountID)
if err != nil {
slog.Warn(
"grok_import_active_probe_failed",
@@ -36,7 +36,7 @@ func newGrokImportProbeStub(buffer int) *grokImportProbeStub {
}
}
func (s *grokImportProbeStub) ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) {
func (s *grokImportProbeStub) QueryQuota(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) {
_, deadlineSeen := ctx.Deadline()
s.mu.Lock()
s.calls[accountID]++
@@ -69,7 +69,7 @@ func (s *grokImportProbeStub) ProbeUsage(ctx context.Context, accountID int64) (
return nil, failure
}
return &service.GrokQuotaProbeResult{
Source: "active_probe",
Source: "hybrid_probe",
Model: "grok-4.5",
StatusCode: 200,
ResetSupported: false,
@@ -23,7 +23,7 @@ type GrokOAuthHandler struct {
grokOAuthService *service.GrokOAuthService
adminService service.AdminService
quotaService *service.GrokQuotaService
importProber grokUsageProber
importProber grokImportProber
reconciler service.GrokOAuthReconciler
}
@@ -0,0 +1,21 @@
package admin
import (
"net/http"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/gin-gonic/gin"
)
// GetAuthCacheInvalidationHealth exposes durable outbox lag and subscriber health.
func (h *OpsHandler) GetAuthCacheInvalidationHealth(c *gin.Context) {
if h.opsService == nil {
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
return
}
if err := h.opsService.RequireMonitoringEnabled(c.Request.Context()); err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, h.opsService.GetAuthCacheInvalidationHealth(c.Request.Context()))
}
@@ -0,0 +1,131 @@
package admin
import (
"net/http"
"net/netip"
"strconv"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
var ingressRejectReasons = map[string]struct{}{
"query_api_key_deprecated": {}, "api_key_required": {}, "invalid_api_key": {},
"invalid_auth_rate_limited": {},
"api_key_auth_overloaded": {},
"api_key_disabled": {}, "ip_restricted": {}, "user_inactive": {}, "group_deleted": {},
"group_disabled": {}, "group_not_allowed": {}, "group_unassigned": {}, "other": {},
}
var ingressRejectRouteFamilies = map[string]struct{}{
"antigravity": {}, "gemini": {}, "codex": {}, "messages": {}, "responses": {},
"chat_completions": {}, "images": {}, "videos": {}, "embeddings": {}, "models": {}, "other": {},
}
var ingressRejectProtocols = map[string]struct{}{
"google": {}, "anthropic": {}, "openai": {}, "gateway": {}, "other": {},
}
// ListIngressRejects returns bounded security aggregates, never raw credentials or request bodies.
func (h *OpsHandler) ListIngressRejects(c *gin.Context) {
if h.opsService == nil {
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
return
}
page, pageSize := response.ParsePagination(c)
if pageSize > 200 {
pageSize = 200
}
startTime, endTime, err := parseOpsTimeRange(c, "1h")
if err != nil {
response.BadRequest(c, err.Error())
return
}
filter := &service.OpsIngressRejectFilter{Page: page, PageSize: pageSize}
if !startTime.IsZero() {
filter.StartTime = &startTime
}
if !endTime.IsZero() {
filter.EndTime = &endTime
}
if filter.RejectReason, err = parseIngressRejectEnum(c, "reason", ingressRejectReasons); err != nil {
response.BadRequest(c, err.Error())
return
}
if filter.RouteFamily, err = parseIngressRejectEnum(c, "route_family", ingressRejectRouteFamilies); err != nil {
response.BadRequest(c, err.Error())
return
}
if filter.Protocol, err = parseIngressRejectEnum(c, "protocol", ingressRejectProtocols); err != nil {
response.BadRequest(c, err.Error())
return
}
if raw := strings.TrimSpace(c.Query("client_ip")); raw != "" {
addr, parseErr := netip.ParseAddr(raw)
if parseErr != nil {
response.BadRequest(c, "Invalid client_ip")
return
}
addr = addr.Unmap()
if addr.Is6() {
addr = netip.PrefixFrom(addr, 64).Masked().Addr()
}
filter.ClientIP = addr.String()
}
if filter.UserID, err = parseOptionalPositiveID(c, "user_id"); err != nil {
response.BadRequest(c, err.Error())
return
}
if filter.APIKeyID, err = parseOptionalPositiveID(c, "api_key_id"); err != nil {
response.BadRequest(c, err.Error())
return
}
result, err := h.opsService.ListIngressRejects(c.Request.Context(), filter)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, result)
}
func (h *OpsHandler) GetIngressRejectHealth(c *gin.Context) {
if h.opsService == nil {
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
return
}
if err := h.opsService.RequireMonitoringEnabled(c.Request.Context()); err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, h.opsService.GetIngressRejectHealth())
}
func parseIngressRejectEnum(c *gin.Context, name string, allowed map[string]struct{}) (string, error) {
value := strings.TrimSpace(c.Query(name))
if value == "" {
return "", nil
}
if _, ok := allowed[value]; !ok {
return "", &ingressRejectQueryError{message: "Invalid " + name}
}
return value, nil
}
func parseOptionalPositiveID(c *gin.Context, name string) (*int64, error) {
raw := strings.TrimSpace(c.Query(name))
if raw == "" {
return nil, nil
}
value, err := strconv.ParseInt(raw, 10, 64)
if err != nil || value <= 0 {
return nil, &ingressRejectQueryError{message: "Invalid " + name}
}
return &value, nil
}
type ingressRejectQueryError struct{ message string }
func (e *ingressRejectQueryError) Error() string { return e.message }
@@ -0,0 +1,41 @@
package admin
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func newIngressRejectHandlerForTest() *OpsHandler {
return NewOpsHandler(service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil))
}
func TestListIngressRejectsValidatesFilters(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, path := range []string{
"/api/v1/admin/ops/ingress-rejections?reason=not-valid",
"/api/v1/admin/ops/ingress-rejections?client_ip=not-an-ip",
"/api/v1/admin/ops/ingress-rejections?user_id=0",
"/api/v1/admin/ops/ingress-rejections?api_key_id=-1",
} {
t.Run(path, func(t *testing.T) {
context, _ := gin.CreateTestContext(httptest.NewRecorder())
context.Request = httptest.NewRequest(http.MethodGet, path, nil)
newIngressRejectHandlerForTest().ListIngressRejects(context)
require.Equal(t, http.StatusBadRequest, context.Writer.Status())
})
}
}
func TestGetIngressRejectHealthAvailable(t *testing.T) {
recorder := httptest.NewRecorder()
context, _ := gin.CreateTestContext(recorder)
context.Request = httptest.NewRequest(http.MethodGet, "/api/v1/admin/ops/ingress-rejections/health", nil)
newIngressRejectHandlerForTest().GetIngressRejectHealth(context)
require.Equal(t, http.StatusOK, recorder.Code)
require.Contains(t, recorder.Body.String(), `"capacity":8192`)
}
@@ -1688,6 +1688,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
response.ErrorFrom(c, err)
return
}
if h.opsService != nil {
h.opsService.SetMonitoringEnabled(settings.OpsMonitoringEnabled)
}
// Update OpenAI fast policy (stored under dedicated key, only when provided).
if req.OpenAIFastPolicySettings != nil {
@@ -166,6 +166,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
body = parsedReq.Body.Bytes()
reqModel := parsedReq.Model
reqStream := parsedReq.Stream
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
@@ -1882,6 +1883,7 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
body = parsedReq.Body.Bytes()
// count_tokens 走 messages 严格校验时,复用已解析请求,避免二次反序列化。
SetClaudeCodeClientContext(c, body, parsedReq)
reqLog = reqLog.With(zap.String("model", parsedReq.Model), zap.Bool("stream", parsedReq.Stream))
+36 -1
View File
@@ -181,6 +181,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
sameAccountRetryCount := make(map[int64]int)
var lastFailoverErr *service.UpstreamFailoverError
var oauth429FailoverState service.OpenAIOAuth429FailoverState
mediaEligibilityRejected := false
switchCount := 0
maxAccountSwitches := h.maxAccountSwitches
if maxAccountSwitches <= 0 {
@@ -216,7 +217,8 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
zap.Error(err),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if endpoint.IsGenerationRequest() && len(failedAccountIDs) == 0 && errors.Is(err, service.ErrNoAvailableAccounts) {
if endpoint.IsGenerationRequest() && errors.Is(err, service.ErrNoAvailableAccounts) &&
(len(failedAccountIDs) == 0 || (mediaEligibilityRejected && lastFailoverErr == nil)) {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts")
return
@@ -268,6 +270,25 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
)
account := selection.Account
if endpoint.IsGenerationRequest() {
eligible, eligibilityReason, eligibilityErr := h.ensureGrokMediaAccountEligibility(requestCtx, account)
if !eligible {
mediaEligibilityRejected = true
failedAccountIDs[account.ID] = struct{}{}
reqLog.Warn("grok_media.account_eligibility_rejected",
zap.Int64("account_id", account.ID),
zap.String("reason", eligibilityReason),
zap.Bool("probe_failed", eligibilityErr != nil),
)
if switchCount >= maxAccountSwitches {
markOpsRoutingCapacityLimited(c)
h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts")
return
}
switchCount++
continue
}
}
sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account)
setOpsSelectedAccount(c, account.ID, account.Platform)
@@ -393,6 +414,20 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
}
}
func (h *OpenAIGatewayHandler) ensureGrokMediaAccountEligibility(ctx context.Context, account *service.Account) (bool, string, error) {
if account == nil {
return false, "missing_account", errors.New("grok media account is required")
}
eligible, reason := account.GrokMediaGenerationEligibility()
if eligible || reason != "billing_unobserved" {
return eligible, reason, nil
}
if h == nil || h.grokMediaEligibilityProber == nil {
return false, "billing_probe_unavailable", errors.New("grok media eligibility probe is not configured")
}
return h.grokMediaEligibilityProber.ProbeMediaEligibility(ctx, account.ID)
}
func grokMediaRequiredCapability(endpoint service.GrokMediaEndpoint) service.OpenAIEndpointCapability {
if endpoint.IsGenerationRequest() {
return service.OpenAIEndpointCapabilityGrokMediaGeneration
@@ -1,12 +1,26 @@
package handler
import (
"context"
"errors"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
type grokMediaEligibilityProberStub struct {
eligible bool
reason string
err error
calls int
}
func (s *grokMediaEligibilityProberStub) ProbeMediaEligibility(context.Context, int64) (bool, string, error) {
s.calls++
return s.eligible, s.reason, s.err
}
func TestShouldRecordGrokMediaUsage(t *testing.T) {
tests := []struct {
name string
@@ -80,3 +94,55 @@ func TestGrokMediaRequiredCapability(t *testing.T) {
})
}
}
func TestEnsureGrokMediaAccountEligibility(t *testing.T) {
t.Run("non oauth account does not probe", func(t *testing.T) {
prober := &grokMediaEligibilityProberStub{}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeAPIKey}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.NoError(t, err)
require.True(t, eligible)
require.Equal(t, "non_oauth", reason)
require.Zero(t, prober.calls)
})
t.Run("unobserved oauth is probed before forwarding", func(t *testing.T) {
prober := &grokMediaEligibilityProberStub{eligible: true, reason: "eligible"}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{ID: 7, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.NoError(t, err)
require.True(t, eligible)
require.Equal(t, "eligible", reason)
require.Equal(t, 1, prober.calls)
})
t.Run("missing prober fails closed", func(t *testing.T) {
h := &OpenAIGatewayHandler{}
account := &service.Account{ID: 8, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.Error(t, err)
require.False(t, eligible)
require.Equal(t, "billing_probe_unavailable", reason)
})
t.Run("probe failure fails closed", func(t *testing.T) {
probeErr := errors.New("probe failed")
prober := &grokMediaEligibilityProberStub{reason: "billing_unobserved", err: probeErr}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{ID: 9, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.ErrorIs(t, err, probeErr)
require.False(t, eligible)
require.Equal(t, "billing_unobserved", reason)
})
}
+6 -5
View File
@@ -38,11 +38,12 @@ type noAccountErrorClassification struct {
// The classifier intentionally does not consume the original error: the
// selection layer never tells us *why* the pool came up empty (rate-limited
// vs. unsupported model are both wrapped as ErrNoAvailableAccounts). Instead
// we re-check pool composition through DiagnoseModelAvailabilityForPlatform,
// which only inspects model_mapping configuration and ignores transient
// state. That guarantees a 404 is only returned when no operator action
// short of editing the account's model_mapping could make this request
// succeed.
// we re-check pool composition through DiagnoseModelAvailabilityForPlatform.
// Its dedicated database query considers only persistent eligibility
// (active status + schedulable setting) and model_mapping, bypassing scheduler
// snapshots and transient filters. That guarantees a 404 is only returned
// when persistent account/group/model configuration must change before the
// request can succeed.
//
// routingModel is the model name that account selection actually compared
// against (i.e. after group-level dispatch mapping). displayModel is the
@@ -154,6 +154,20 @@ func TestClassifyNoAccountError_HasModelSupport_KeepsRoutingMessageGenerationToC
require.False(t, cls.ModelNotFound)
}
func TestClassifyNoAccountError_ModelSupportedOnlyByRateLimitedAccount_Returns503(t *testing.T) {
c := newTestGinContextWithRequest()
// The diagnoser's configured-state lookup still sees the model-supporting
// account even though normal scheduling has excluded it during cooldown.
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: true}}
apiKey := &service.APIKey{GroupID: ptrInt64(7)}
cls := classifyNoAccountErrorFromGin(c, fd, apiKey, "claude-opus-4-8", "claude-opus-4-8", service.PlatformAnthropic)
require.Equal(t, http.StatusServiceUnavailable, cls.Status)
require.Equal(t, "api_error", cls.ErrType)
require.False(t, cls.ModelNotFound, "temporary account cooldown must remain retryable")
}
func TestClassifyNoAccountError_NoAccountsInPool_Stays503(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: false, HasModelSupport: false}}
@@ -305,7 +305,10 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
upstreamErrorAlreadyCommunicated := openAIForwardErrorAlreadyCommunicated(c, writerSizeBeforeForward, err)
wroteFallback := false
if !upstreamErrorAlreadyCommunicated {
wroteFallback = h.ensureForwardErrorResponse(c, streamStarted)
wroteFallback = h.ensureOpenAIStreamReadErrorResponse(c, err, streamStarted)
if !wroteFallback {
wroteFallback = h.ensureForwardErrorResponse(c, streamStarted)
}
}
reqLog.Warn("openai_chat_completions.forward_failed",
zap.Int64("account_id", account.ID),
@@ -67,6 +67,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
body = parsedReq.Body.Bytes()
if parsedReq.Model == "" {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
@@ -824,6 +824,7 @@ func newGrokCredentialFailoverHandler(t *testing.T, mode string) (*OpenAIGateway
"access_token": "expired", "refresh_token": "revoked-refresh",
"expires_at": time.Now().Add(-time.Minute).UTC().Format(time.RFC3339),
},
Extra: map[string]any{service.GrokMediaEligibleExtraKey: true},
},
{
ID: 802, Name: "healthy", Platform: service.PlatformGrok, Type: service.AccountTypeOAuth,
@@ -832,6 +833,7 @@ func newGrokCredentialFailoverHandler(t *testing.T, mode string) (*OpenAIGateway
"access_token": "healthy-access", "refresh_token": "healthy-refresh",
"expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339),
},
Extra: map[string]any{service.GrokMediaEligibleExtraKey: true},
},
}
if mode == "postmap_cancel" || mode == "first_429" || mode == "all_429" || mode == "mixed_429_500" || mode == "mixed_500_429" || mode == "oauth_429_apikey_500" {
@@ -845,6 +847,7 @@ func newGrokCredentialFailoverHandler(t *testing.T, mode string) (*OpenAIGateway
"access_token": "untried-healthy-access", "refresh_token": "untried-healthy-refresh",
"expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339),
},
Extra: map[string]any{service.GrokMediaEligibleExtraKey: true},
})
}
if mode == "oauth_429_apikey_500" {
@@ -206,6 +206,10 @@ func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T)
require.NotContains(t, job.entry.ErrorMessage, "earlier inference failure")
require.NotNil(t, job.entry.UpstreamStatusCode)
require.Zero(t, *job.entry.UpstreamStatusCode)
require.Len(t, job.entry.UpstreamErrors, 2)
require.Equal(t, http.StatusForbidden, job.entry.UpstreamErrors[0].UpstreamStatusCode)
require.Nil(t, job.entry.UpstreamErrors)
require.NotNil(t, job.entry.UpstreamErrorsJSON)
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
require.NoError(t, err)
require.Len(t, events, 2)
require.Equal(t, http.StatusForbidden, events[0].UpstreamStatusCode)
}
@@ -28,18 +28,23 @@ import (
// OpenAIGatewayHandler handles OpenAI API gateway requests
type OpenAIGatewayHandler struct {
gatewayService *service.OpenAIGatewayService
billingCacheService *service.BillingCacheService
apiKeyService *service.APIKeyService
usageRecordWorkerPool *service.UsageRecordWorkerPool
errorPassthroughService *service.ErrorPassthroughService
contentModerationService *service.ContentModerationService
securityAuditCoordinator *securityaudit.Coordinator
opsService *service.OpsService
concurrencyHelper *ConcurrencyHelper
imageLimiter *imageConcurrencyLimiter
maxAccountSwitches int
cfg *config.Config
gatewayService *service.OpenAIGatewayService
billingCacheService *service.BillingCacheService
apiKeyService *service.APIKeyService
usageRecordWorkerPool *service.UsageRecordWorkerPool
errorPassthroughService *service.ErrorPassthroughService
contentModerationService *service.ContentModerationService
securityAuditCoordinator *securityaudit.Coordinator
grokMediaEligibilityProber grokMediaEligibilityProber
opsService *service.OpsService
concurrencyHelper *ConcurrencyHelper
imageLimiter *imageConcurrencyLimiter
maxAccountSwitches int
cfg *config.Config
}
type grokMediaEligibilityProber interface {
ProbeMediaEligibility(ctx context.Context, accountID int64) (bool, string, error)
}
const maxOpenAIFirstOutputTimeoutSwitches = 1
@@ -2231,12 +2236,29 @@ func (h *OpenAIGatewayHandler) mapUpstreamError(statusCode int) (int, string, st
// handleStreamingAwareError handles errors that may occur after streaming has started
func (h *OpenAIGatewayHandler) handleStreamingAwareError(c *gin.Context, status int, errType, message string, streamStarted bool) {
h.handleStreamingAwareErrorWithCode(c, status, errType, "", message, streamStarted, false)
}
func (h *OpenAIGatewayHandler) handleStreamingAwareErrorWithCode(
c *gin.Context,
status int,
errType string,
code string,
message string,
streamStarted bool,
countTowardsSLA bool,
) {
// body-signal compact 心跳可能已把响应头提交为 200:先停心跳(建立
// happens-before,接管 ResponseWriter),并升级为流内错误处理。
if service.StopOpenAICompactSSEKeepaliveCommitted(c) {
streamStarted = true
}
if streamStarted {
if countTowardsSLA {
service.MarkOpsStreamFailure(c, errType, code, message, status)
} else {
service.MarkOpsStreamError(c, errType, message, status)
}
// /v1/responses 的严格 SDK(Codex CLI)要求终止事件必须属于
// response.completed/failed/incomplete/cancelled 集合。
// 通用 `event: error` 帧不被识别为终止事件,会导致
@@ -2249,8 +2271,15 @@ func (h *OpenAIGatewayHandler) handleStreamingAwareError(c *gin.Context, status
// Stream already started, send error as SSE event then close
flusher, ok := c.Writer.(http.Flusher)
if ok {
// SSE 错误事件固定 schema,使用 Quote 直拼可避免额外 Marshal 分配。
errorEvent := "event: error\ndata: " + `{"error":{"type":` + strconv.Quote(errType) + `,"message":` + strconv.Quote(message) + `}}` + "\n\n"
errorObject := gin.H{"type": errType, "message": message}
if code != "" {
errorObject["code"] = code
}
payload, err := json.Marshal(gin.H{"error": errorObject})
if err != nil {
payload = []byte(`{"error":{"type":"upstream_error","message":"Upstream request failed"}}`)
}
errorEvent := "event: error\ndata: " + string(payload) + "\n\n"
if _, err := fmt.Fprint(c.Writer, errorEvent); err != nil {
_ = c.Error(err)
}
@@ -2260,7 +2289,27 @@ func (h *OpenAIGatewayHandler) handleStreamingAwareError(c *gin.Context, status
}
// Normal case: return JSON response with proper status code
h.errorResponse(c, status, errType, message)
if code == "" {
h.errorResponse(c, status, errType, message)
return
}
c.JSON(status, gin.H{"error": gin.H{
"type": errType, "code": code, "message": message,
}})
}
func (h *OpenAIGatewayHandler) ensureOpenAIStreamReadErrorResponse(c *gin.Context, err error, streamStarted bool) bool {
code, message, ok := service.OpenAIUpstreamStreamReadErrorDetails(err)
if !ok || c == nil || c.Writer == nil || service.IsResponseCommitted(c) {
return false
}
if c.Writer.Written() {
streamStarted = true
}
h.handleStreamingAwareErrorWithCode(
c, http.StatusBadGateway, "upstream_error", code, message, streamStarted, true,
)
return true
}
// ensureForwardErrorResponse 在 Forward 返回错误但尚未写响应时补写统一错误响应。
@@ -96,6 +96,35 @@ func TestOpenAIHandleStreamingAwareError_JSONEscaping(t *testing.T) {
}
}
func TestOpenAIHandleStreamingAwareErrorWithCode_EmitsStableClassification(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
h := &OpenAIGatewayHandler{}
h.handleStreamingAwareErrorWithCode(
c,
http.StatusBadGateway,
"upstream_error",
service.OpenAIUpstreamHTTP2StreamErrorCode,
"Upstream HTTP/2 stream failed",
true,
true,
)
body := w.Body.String()
require.Contains(t, body, "event: error\n")
require.Equal(t, "upstream_error", gjson.Get(body[strings.Index(body, "{"):], "error.type").String())
require.Equal(t, service.OpenAIUpstreamHTTP2StreamErrorCode, gjson.Get(body[strings.Index(body, "{"):], "error.code").String())
require.NotContains(t, body, "stream ID")
streamErr, ok := service.GetOpsStreamError(c)
require.True(t, ok)
require.True(t, streamErr.CountTowardsSLA)
require.Equal(t, http.StatusBadGateway, streamErr.IntendedStatus)
}
func TestOpenAIForwardSucceededForScheduling(t *testing.T) {
require.True(t, openAIForwardSucceededForScheduling(nil))
require.True(t, openAIForwardSucceededForScheduling(&service.OpenAIForwardResult{}))
@@ -1875,6 +1904,212 @@ func TestOpenAIResponsesWebSocket_FailoverOnUpstreamUsageLimitEvent(t *testing.T
require.Equal(t, []int64{int64(9902)}, accountRepo.rateLimitedIDs)
}
func TestOpenAIResponsesWebSocket_FirstOutputTimeoutWithoutDownstreamReusesClientForOneFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
firstHitCh := make(chan []byte, 1)
secondHitCh := make(chan []byte, 1)
var firstConnections atomic.Int32
var secondConnections atomic.Int32
firstUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
firstConnections.Add(1)
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, payload, readErr := conn.Read(readCtx)
cancelRead()
if readErr == nil {
firstHitCh <- payload
}
select {
case <-r.Context().Done():
case <-time.After(3 * time.Second):
}
}))
defer firstUpstream.Close()
secondUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
secondConnections.Add(1)
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, payload, readErr := conn.Read(readCtx)
cancelRead()
if readErr == nil {
secondHitCh <- payload
}
for _, event := range []string{
`{"type":"response.created","response":{"id":"resp_ws_timeout_b","model":"gpt-5.1"}}`,
`{"type":"response.output_text.delta","response_id":"resp_ws_timeout_b","delta":"recovered"}`,
`{"type":"response.completed","response":{"id":"resp_ws_timeout_b","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`,
} {
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
writeErr := conn.Write(writeCtx, coderws.MessageText, []byte(event))
cancelWrite()
if writeErr != nil {
return
}
}
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
_, _, _ = conn.Read(readCtx)
cancelRead()
}))
defer secondUpstream.Close()
groupID := int64(4212)
accounts := []service.Account{
{
ID: 9912,
Name: "openai-ws-first-semantic-timeout",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 1,
Credentials: map[string]any{"api_key": "sk-first", "base_url": firstUpstream.URL},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
},
},
{
ID: 9913,
Name: "openai-ws-failover-healthy",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 2,
Credentials: map[string]any{"api_key": "sk-second", "base_url": secondUpstream.URL},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
},
},
}
cfg := &config.Config{}
cfg.RunMode = config.RunModeSimple
cfg.Default.RateMultiplier = 1
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIFirstOutputTimeoutSeconds = 1
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 3
cfg.Gateway.MaxAccountSwitches = 3
accountRepo := &openAIWSFailoverHandlerAccountRepoStub{accounts: accounts}
rateLimitSvc := service.NewRateLimitService(accountRepo, nil, cfg, nil, nil)
billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
gatewaySvc := service.NewOpenAIGatewayService(
accountRepo, nil, nil, nil, nil, nil, nil, cfg, nil, nil,
service.NewBillingService(cfg, nil), rateLimitSvc, billingCacheSvc,
nil, &service.DeferredService{}, nil, nil, nil, nil, nil, nil, nil,
)
cache := &concurrencyCacheMock{
acquireUserSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil },
acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) {
return true, nil
},
}
h := &OpenAIGatewayHandler{
gatewayService: gatewaySvc,
billingCacheService: billingCacheSvc,
apiKeyService: &service.APIKeyService{},
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second),
maxAccountSwitches: 3,
}
apiKey := &service.APIKey{
ID: 1812,
GroupID: &groupID,
User: &service.User{ID: 1712, Status: service.StatusActive},
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, Status: service.StatusActive},
}
handlerDone := make(chan struct{})
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware.ContextKeyAPIKey), apiKey)
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1})
c.Next()
})
router.GET("/openai/v1/responses", func(c *gin.Context) {
h.ResponsesWebSocket(c)
close(handlerDone)
})
handlerServer := httptest.NewServer(router)
defer handlerServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(
dialCtx,
"ws"+strings.TrimPrefix(handlerServer.URL, "http")+"/openai/v1/responses",
&coderws.DialOptions{CompressionMode: coderws.CompressionContextTakeover},
)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false}`))
cancelWrite()
require.NoError(t, err)
var eventTypes []string
readCtx, cancelRead := context.WithTimeout(context.Background(), 6*time.Second)
for {
_, event, readErr := clientConn.Read(readCtx)
require.NoError(t, readErr)
eventType := gjson.GetBytes(event, "type").String()
eventTypes = append(eventTypes, eventType)
if eventType == "response.completed" {
require.Equal(t, "resp_ws_timeout_b", gjson.GetBytes(event, "response.id").String())
break
}
}
cancelRead()
require.Contains(t, eventTypes, "response.output_text.delta")
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
select {
case <-handlerDone:
case <-time.After(3 * time.Second):
t.Fatal("websocket handler did not finish after healthy failover turn")
}
select {
case <-firstHitCh:
case <-time.After(3 * time.Second):
t.Fatal("first upstream did not receive replayable request")
}
select {
case <-secondHitCh:
case <-time.After(3 * time.Second):
t.Fatal("second upstream did not receive replayed request")
}
require.Equal(t, int32(1), firstConnections.Load())
require.Equal(t, int32(1), secondConnections.Load())
require.NotContains(t, accountRepo.rateLimitedIDs, int64(9913), "healthy failover account must not be penalized")
}
func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSUsageLogCase) openAIResponsesWSUsageLogResult {
t.Helper()
gin.SetMode(gin.TestMode)
+146 -79
View File
@@ -72,24 +72,10 @@ const (
opsErrorLogMinQueueSize = 256
opsErrorLogMaxQueueSize = 8192
opsErrorLogBatchSize = 32
opsErrorLogMaxQueueBytes = 32 * 1024 * 1024
opsErrorLogMaxUserAgentBytes = 512
)
// looksLikeSystemKey 粗筛"形似本系统 key"的输入:长度 16-128 且仅含 [a-zA-Z0-9_-]。
// 不用前缀匹配(APIKeyPrefix 可配置)。用于反查审计表前挡掉随机扫描的乱码输入。
func looksLikeSystemKey(key string) bool {
if len(key) < 16 || len(key) > 128 {
return false
}
for _, c := range key {
allowed := (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
(c >= '0' && c <= '9') || c == '_' || c == '-'
if !allowed {
return false
}
}
return true
}
// keyPrefix 返回脱敏前缀(前 n 个字符);不足 n 则原样返回。
func keyPrefix(key string, n int) string {
if len(key) <= n {
@@ -98,43 +84,26 @@ func keyPrefix(key string, n int) string {
return key[:n]
}
// extractAttemptedKey 按认证中间件同样的顺序从请求头提取提交的 key 明文。
// 与 api_key_auth.go:43-59 一致:Authorization 仅取 Bearer scheme,非 Bearer 则忽略并继续 x-api-key → x-goog-api-key。
func extractAttemptedKey(c *gin.Context) string {
if h := c.GetHeader("Authorization"); h != "" {
parts := strings.SplitN(h, " ", 2)
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
return strings.TrimSpace(parts[1])
}
// 非 Bearer:与中间件一致,忽略 Authorization,继续尝试其它 header(不在此 return)。
}
if k := c.GetHeader("x-api-key"); k != "" {
return strings.TrimSpace(k)
}
if k := c.GetHeader("x-goog-api-key"); k != "" {
return strings.TrimSpace(k)
}
return ""
}
type opsErrorLogJob struct {
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
queuedBytes int64
}
var (
opsErrorLogOnce sync.Once
opsErrorLogQueue chan opsErrorLogJob
opsErrorLogStopOnce sync.Once
opsErrorLogWorkersWg sync.WaitGroup
opsErrorLogMu sync.RWMutex
opsErrorLogStopping bool
opsErrorLogQueueLen atomic.Int64
opsErrorLogEnqueued atomic.Int64
opsErrorLogDropped atomic.Int64
opsErrorLogProcessed atomic.Int64
opsErrorLogSanitized atomic.Int64
opsErrorLogStopOnce sync.Once
opsErrorLogWorkersWg sync.WaitGroup
opsErrorLogMu sync.RWMutex
opsErrorLogStopping bool
opsErrorLogQueueLen atomic.Int64
opsErrorLogQueueBytes atomic.Int64
opsErrorLogEnqueued atomic.Int64
opsErrorLogDropped atomic.Int64
opsErrorLogProcessed atomic.Int64
opsErrorLogSanitized atomic.Int64
opsErrorLogLastDropLogAt atomic.Int64
@@ -154,6 +123,7 @@ func startOpsErrorLogWorkers() {
workerCount, queueSize := opsErrorLogConfig()
opsErrorLogQueue = make(chan opsErrorLogJob, queueSize)
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
opsErrorLogWorkersWg.Add(workerCount)
for i := 0; i < workerCount; i++ {
@@ -165,6 +135,7 @@ func startOpsErrorLogWorkers() {
return
}
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-job.queuedBytes)
batch := make([]opsErrorLogJob, 0, opsErrorLogBatchSize)
batch = append(batch, job)
@@ -184,6 +155,7 @@ func startOpsErrorLogWorkers() {
return
}
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-nextJob.queuedBytes)
batch = append(batch, nextJob)
case <-timer.C:
break batchLoop
@@ -239,6 +211,20 @@ func enqueueOpsErrorLog(ops *service.OpsService, entry *service.OpsInsertErrorLo
if ops == nil || entry == nil {
return
}
entry.UserAgent = normalizeOpsPersistentUserAgent(entry.UserAgent)
if entry.ErrorBody != "" {
originalBody := entry.ErrorBody
body, truncated := service.SanitizeOpsErrorBodyForQueue(originalBody)
entry.ErrorBody = body
if truncated || body != originalBody {
opsErrorLogSanitized.Add(1)
}
}
if err := service.SanitizeOpsUpstreamErrorsForQueue(entry); err != nil {
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
return
}
select {
case <-opsErrorLogShutdownCh:
return
@@ -259,18 +245,29 @@ func enqueueOpsErrorLog(ops *service.OpsService, entry *service.OpsInsertErrorLo
if opsErrorLogStopping || opsErrorLogQueue == nil {
return
}
queuedBytes := estimateOpsErrorLogJobBytes(entry)
if !reserveOpsErrorLogQueueBytes(queuedBytes) {
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
return
}
select {
case opsErrorLogQueue <- opsErrorLogJob{ops: ops, entry: entry}:
opsErrorLogQueueLen.Add(1)
case opsErrorLogQueue <- opsErrorLogJob{ops: ops, entry: entry, queuedBytes: queuedBytes}:
opsErrorLogEnqueued.Add(1)
default:
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-queuedBytes)
// Queue is full; drop to avoid blocking request handling.
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
}
}
func normalizeOpsPersistentUserAgent(value string) string {
return truncateString(strings.TrimSpace(strings.ToValidUTF8(value, "")), opsErrorLogMaxUserAgentBytes)
}
func StopOpsErrorLogWorkers() bool {
opsErrorLogStopOnce.Do(func() {
opsErrorLogShutdownOnce.Do(func() {
@@ -293,6 +290,7 @@ func stopOpsErrorLogWorkers() bool {
if ch == nil {
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
return true
}
@@ -305,6 +303,7 @@ func stopOpsErrorLogWorkers() bool {
select {
case <-done:
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
return true
case <-time.After(opsErrorLogDrainTimeout):
return false
@@ -315,6 +314,14 @@ func OpsErrorLogQueueLength() int64 {
return opsErrorLogQueueLen.Load()
}
func OpsErrorLogQueueBytes() int64 {
return opsErrorLogQueueBytes.Load()
}
func OpsErrorLogQueueBytesCapacity() int64 {
return opsErrorLogMaxQueueBytes
}
func OpsErrorLogQueueCapacity() int {
opsErrorLogMu.RLock()
ch := opsErrorLogQueue
@@ -355,12 +362,15 @@ func maybeLogOpsErrorLogDrop() {
}
queued := opsErrorLogQueueLen.Load()
queuedBytes := opsErrorLogQueueBytes.Load()
queueCap := OpsErrorLogQueueCapacity()
log.Printf(
"[OpsErrorLogger] queue is full; dropping logs (queued=%d cap=%d enqueued_total=%d dropped_total=%d processed_total=%d sanitized_total=%d)",
"[OpsErrorLogger] queue is full; dropping logs (queued=%d cap=%d queued_bytes=%d bytes_cap=%d enqueued_total=%d dropped_total=%d processed_total=%d sanitized_total=%d)",
queued,
queueCap,
queuedBytes,
opsErrorLogMaxQueueBytes,
opsErrorLogEnqueued.Load(),
opsErrorLogDropped.Load(),
opsErrorLogProcessed.Load(),
@@ -368,6 +378,46 @@ func maybeLogOpsErrorLogDrop() {
)
}
func reserveOpsErrorLogQueueBytes(size int64) bool {
if size < 1 {
size = 1
}
for {
current := opsErrorLogQueueBytes.Load()
if current > opsErrorLogMaxQueueBytes-size {
return false
}
if opsErrorLogQueueBytes.CompareAndSwap(current, current+size) {
opsErrorLogQueueLen.Add(1)
return true
}
}
}
func estimateOpsErrorLogJobBytes(entry *service.OpsInsertErrorLogInput) int64 {
if entry == nil {
return 1
}
const fixedOverhead = 512
size := fixedOverhead + len(entry.RequestID) + len(entry.ClientRequestID) +
len(entry.Platform) + len(entry.Model) + len(entry.RequestPath) +
len(entry.InboundEndpoint) + len(entry.UpstreamEndpoint) +
len(entry.RequestedModel) + len(entry.UpstreamModel) + len(entry.UserAgent) +
len(entry.ErrorPhase) + len(entry.ErrorType) + len(entry.Severity) +
len(entry.ErrorMessage) + len(entry.ErrorBody) + len(entry.ErrorSource) +
len(entry.ErrorOwner) + len(entry.APIKeyPrefix)
if entry.UpstreamErrorMessage != nil {
size += len(*entry.UpstreamErrorMessage)
}
if entry.UpstreamErrorDetail != nil {
size += len(*entry.UpstreamErrorDetail)
}
if entry.UpstreamErrorsJSON != nil {
size += len(*entry.UpstreamErrorsJSON)
}
return int64(size)
}
func opsErrorLogConfig() (workerCount int, queueSize int) {
workerCount = runtime.GOMAXPROCS(0) * 2
if workerCount < opsErrorLogMinWorkerCount {
@@ -470,9 +520,12 @@ type opsCaptureWriter struct {
gin.ResponseWriter
limit int
buf bytes.Buffer
ctx *gin.Context
}
const opsCaptureWriterLimit = 64 * 1024
const opsCaptureWriterLimit = service.OpsErrorLogQueueBodyMaxBytes
const opsCaptureWriterPoolMaxRetainedCapacity = service.OpsErrorLogQueueBodyMaxBytes
var opsCaptureWriterPool = sync.Pool{
New: func() any {
@@ -496,11 +549,19 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) {
return
}
w.ResponseWriter = nil
w.ctx = nil
w.limit = opsCaptureWriterLimit
if !shouldPoolOpsCaptureWriter(w) {
return
}
w.buf.Reset()
opsCaptureWriterPool.Put(w)
}
func shouldPoolOpsCaptureWriter(w *opsCaptureWriter) bool {
return w != nil && w.buf.Cap() <= opsCaptureWriterPoolMaxRetainedCapacity
}
func (w *opsCaptureWriter) Status() int {
if w.ResponseWriter == nil {
return 0
@@ -577,7 +638,7 @@ func (w *opsCaptureWriter) Write(b []byte) (int, error) {
if w.ResponseWriter == nil {
return 0, nil
}
if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
remaining := w.limit - w.buf.Len()
if len(b) > remaining {
_, _ = w.buf.Write(b[:remaining])
@@ -592,7 +653,7 @@ func (w *opsCaptureWriter) WriteString(s string) (int, error) {
if w.ResponseWriter == nil {
return 0, nil
}
if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
remaining := w.limit - w.buf.Len()
if len(s) > remaining {
_, _ = w.buf.WriteString(s[:remaining])
@@ -603,6 +664,14 @@ func (w *opsCaptureWriter) WriteString(s string) (int, error) {
return w.ResponseWriter.WriteString(s)
}
func (w *opsCaptureWriter) shouldCapture() bool {
if w.ctx == nil {
return true
}
_, rejected := middleware2.GetIngressRejectReason(w.ctx)
return !rejected
}
// OpsErrorLoggerMiddleware records error responses (status >= 400) into ops_error_logs.
//
// Notes:
@@ -612,6 +681,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return func(c *gin.Context) {
originalWriter := c.Writer
w := acquireOpsCaptureWriter(originalWriter)
w.ctx = c
defer func() {
// Restore the original writer before returning so outer middlewares
// don't observe a pooled wrapper that has been released.
@@ -623,6 +693,10 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
c.Writer = w
c.Next()
if _, rejected := middleware2.GetIngressRejectReason(c); rejected {
return
}
if ops == nil {
return
}
@@ -989,7 +1063,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
IsCountTokens: isCountTokensRequest(c),
ErrorMessage: parsed.Message,
// Keep the full captured error body (capture is already capped at 64KB) so the
// Keep the captured error body (already capped at the queue-safe limit) so the
// service layer can sanitize JSON before truncating for storage.
ErrorBody: string(body),
ErrorSource: errorSource,
@@ -1002,7 +1076,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
// 有效(未删除)key 报错时快照前缀,key 之后被删也保留;与 INVALID_API_KEY 的 attempted_key_prefix 互斥。
// 有效 key 报错时快照前缀,key 之后被删也保留。
entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
@@ -1022,22 +1096,6 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
entry.ClientIP = &clientIP
}
// 已删除 key 归因:仅 INVALID_API_KEY 才尝试。响应已写出,此处不阻塞客户端。
if parsed.Code == opsCodeInvalidAPIKey {
if attemptedKey := extractAttemptedKey(c); attemptedKey != "" {
entry.AttemptedKeyPrefix = keyPrefix(attemptedKey, 8)
if looksLikeSystemKey(attemptedKey) {
if res, lookupErr := ops.LookupDeletedKeyAudit(c.Request.Context(), attemptedKey); lookupErr != nil {
log.Printf("[OpsErrorLogger] LookupDeletedKeyAudit failed: %v", lookupErr)
} else if res != nil {
owner := res.UserID
entry.DeletedKeyOwnerUserID = &owner
entry.DeletedKeyName = res.KeyName
}
}
}
}
enqueueOpsErrorLog(ops, entry)
}
}
@@ -1072,8 +1130,20 @@ func logOpsStreamError(c *gin.Context, ops *service.OpsService, wireStatus int)
if classifyStatus <= 0 {
classifyStatus = wireStatus
}
normalizedType := normalizeOpsErrorType(streamErr.ErrType, "")
phase, isBusinessLimited, errorOwner, errorSource := classifyOpsErrorLog(c, normalizedType, streamErr.Message, "", classifyStatus)
normalizedType := normalizeOpsErrorType(streamErr.ErrType, streamErr.Code)
phase, isBusinessLimited, errorOwner, errorSource := classifyOpsErrorLog(c, normalizedType, streamErr.Message, streamErr.Code, classifyStatus)
recordedStatus := wireStatus
if streamErr.CountTowardsSLA && streamErr.IntendedStatus >= 400 {
recordedStatus = streamErr.IntendedStatus
}
errorBody := ""
if streamErr.Code != "" {
if payload, err := json.Marshal(gin.H{"error": gin.H{
"type": normalizedType, "code": streamErr.Code, "message": streamErr.Message,
}}); err == nil {
errorBody = string(payload)
}
}
apiKey := getOpsAPIKey(c)
clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
@@ -1140,12 +1210,12 @@ func logOpsStreamError(c *gin.Context, ops *service.OpsService, wireStatus int)
ErrorPhase: phase,
ErrorType: normalizedType,
Severity: classifyOpsSeverity(normalizedType, classifyStatus),
StatusCode: wireStatus,
StatusCode: recordedStatus,
IsBusinessLimited: isBusinessLimited,
IsCountTokens: isCountTokensRequest(c),
ErrorMessage: streamErr.Message,
ErrorBody: "",
ErrorBody: errorBody,
ErrorSource: errorSource,
ErrorOwner: errorOwner,
@@ -1694,11 +1764,8 @@ func shouldSkipOpsErrorLog(ctx context.Context, ops *service.OpsService, message
}
// Get advanced settings to check filter configuration
settings, err := ops.GetOpsAdvancedSettings(ctx)
if err != nil || settings == nil {
// If we can't get settings, don't skip (fail open)
return false
}
_ = ctx
settings := ops.OpsAdvancedSettingsSnapshot()
msgLower := strings.ToLower(message)
bodyLower := strings.ToLower(body)
@@ -1,39 +1,9 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
)
func TestLooksLikeSystemKey(t *testing.T) {
cases := []struct {
in string
want bool
}{
{"sk-abcdef0123456789", true},
{"ABCdef_-0123456789", true},
{"short", false},
{"with space xxxxxxxxxx", false},
{"汉字key1234567890", false},
{"", false},
}
for _, c := range cases {
if got := looksLikeSystemKey(c.in); got != c.want {
t.Errorf("looksLikeSystemKey(%q)=%v want %v", c.in, got, c.want)
}
}
long := make([]byte, 129)
for i := range long {
long[i] = 'a'
}
if looksLikeSystemKey(string(long)) {
t.Errorf("129-char key should be rejected")
}
}
func TestKeyPrefix(t *testing.T) {
if got := keyPrefix("sk-3f2a9c7e", 8); got != "sk-3f2a9" {
t.Errorf("keyPrefix=%q want %q", got, "sk-3f2a9")
@@ -42,77 +12,3 @@ func TestKeyPrefix(t *testing.T) {
t.Errorf("short key should be returned as-is, got %q", got)
}
}
func TestExtractAttemptedKey(t *testing.T) {
gin.SetMode(gin.TestMode)
cases := []struct {
name string
headers map[string]string
want string
}{
{
name: "Bearer in Authorization",
headers: map[string]string{"Authorization": "Bearer sk-testkey0123456789"},
want: "sk-testkey0123456789",
},
{
name: "Bearer case-insensitive",
headers: map[string]string{"Authorization": "BEARER sk-testkey0123456789"},
want: "sk-testkey0123456789",
},
{
name: "x-api-key header",
headers: map[string]string{"x-api-key": "sk-xapikey0123456789"},
want: "sk-xapikey0123456789",
},
{
name: "x-goog-api-key header",
headers: map[string]string{"x-goog-api-key": "sk-goog0123456789"},
want: "sk-goog0123456789",
},
{
name: "Authorization takes priority over x-api-key",
headers: map[string]string{"Authorization": "Bearer sk-auth0123456789", "x-api-key": "sk-xapi0123456789"},
want: "sk-auth0123456789",
},
{
name: "x-api-key takes priority over x-goog-api-key",
headers: map[string]string{"x-api-key": "sk-xapi0123456789", "x-goog-api-key": "sk-goog0123456789"},
want: "sk-xapi0123456789",
},
{
name: "no key headers",
headers: map[string]string{},
want: "",
},
{
name: "Bearer with leading/trailing spaces trimmed",
headers: map[string]string{"Authorization": "Bearer sk-trimmed0123456789 "},
want: "sk-trimmed0123456789",
},
{
// 非 Bearer Authorization 应被忽略,继续 fall-through 到 x-api-key(与认证中间件一致)
name: "non-Bearer Authorization falls through to x-api-key",
headers: map[string]string{"Authorization": "junk-not-bearer", "x-api-key": "sk-realkey1234567"},
want: "sk-realkey1234567",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
for k, v := range tc.headers {
req.Header.Set(k, v)
}
c.Request = req
got := extractAttemptedKey(c)
if got != tc.want {
t.Errorf("extractAttemptedKey(%v) = %q, want %q", tc.headers, got, tc.want)
}
})
}
}
@@ -1,10 +1,13 @@
package handler
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"unicode/utf8"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
@@ -12,6 +15,82 @@ import (
"github.com/stretchr/testify/require"
)
type ingressRejectSettingRepo struct {
service.SettingRepository
getValueCalls int
}
func (r *ingressRejectSettingRepo) GetValue(context.Context, string) (string, error) {
r.getValueCalls++
return "", service.ErrSettingNotFound
}
func (r *ingressRejectSettingRepo) GetMultiple(context.Context, []string) (map[string]string, error) {
r.getValueCalls++
return map[string]string{}, nil
}
func (r *ingressRejectSettingRepo) Set(context.Context, string, string) error {
return nil
}
type ingressRejectOpsRepo struct {
service.OpsRepository
insertCalls int
}
func (r *ingressRejectOpsRepo) InsertErrorLog(context.Context, *service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
return 0, nil
}
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(context.Context, []*service.OpsInsertErrorLogInput) (int64, error) {
r.insertCalls++
return 0, nil
}
func TestOpsErrorLogQueueByteBudget(t *testing.T) {
previousBytes := opsErrorLogQueueBytes.Load()
previousLen := opsErrorLogQueueLen.Load()
opsErrorLogQueueBytes.Store(0)
opsErrorLogQueueLen.Store(0)
t.Cleanup(func() {
opsErrorLogQueueBytes.Store(previousBytes)
opsErrorLogQueueLen.Store(previousLen)
})
if !reserveOpsErrorLogQueueBytes(opsErrorLogMaxQueueBytes - 1) {
t.Fatal("first reservation within byte budget should succeed")
}
if reserveOpsErrorLogQueueBytes(2) {
t.Fatal("reservation beyond byte budget should be rejected")
}
if got := OpsErrorLogQueueBytes(); got != opsErrorLogMaxQueueBytes-1 {
t.Fatalf("queued bytes = %d, want %d", got, opsErrorLogMaxQueueBytes-1)
}
if got := OpsErrorLogQueueLength(); got != 1 {
t.Fatalf("queue length = %d, want 1", got)
}
}
func TestEstimateOpsErrorLogJobBytesIncludesVariablePayloads(t *testing.T) {
base := estimateOpsErrorLogJobBytes(&service.OpsInsertErrorLogInput{})
message := "upstream message"
detail := "upstream detail"
events := `[{"error":"x"}]`
entry := &service.OpsInsertErrorLogInput{
ErrorBody: strings.Repeat("x", 1024),
ErrorMessage: "client error",
UserAgent: "test-agent",
UpstreamErrorMessage: &message,
UpstreamErrorDetail: &detail,
UpstreamErrorsJSON: &events,
}
if got := estimateOpsErrorLogJobBytes(entry); got <= base+1024 {
t.Fatalf("estimated bytes = %d, expected variable payloads above %d", got, base+1024)
}
}
func resetOpsErrorLoggerStateForTest(t *testing.T) {
t.Helper()
@@ -119,6 +198,29 @@ func TestOpsCaptureWriterPool_ResetOnRelease(t *testing.T) {
require.Zero(t, reused.buf.Len(), "writer should be reset before reuse")
}
func TestOpsCaptureWriterPool_DropsLargeBuffers(t *testing.T) {
w := &opsCaptureWriter{}
w.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
require.False(t, shouldPoolOpsCaptureWriter(w))
}
func TestEnqueueOpsErrorLog_SanitizesAndBoundsBodyBeforeQueue(t *testing.T) {
setupOpsErrorLogTestQueue(t, 1)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
secret := strings.Repeat("s", service.OpsErrorLogQueueBodyMaxBytes)
entry := &service.OpsInsertErrorLogInput{
ErrorPhase: "request",
ErrorType: "api_error",
ErrorBody: `{"authorization":"Bearer ` + secret + `","message":"failed"}`,
}
enqueueOpsErrorLog(ops, entry)
job := <-opsErrorLogQueue
require.LessOrEqual(t, len(job.entry.ErrorBody), service.OpsErrorLogQueueBodyMaxBytes)
require.NotContains(t, job.entry.ErrorBody, secret)
require.Equal(t, int64(1), OpsErrorLogSanitizedTotal())
}
func TestOpsErrorLoggerMiddleware_DoesNotBreakOuterMiddlewares(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -149,6 +251,42 @@ func setupOpsErrorLogTestQueue(t *testing.T, size int) {
opsErrorLogMu.Unlock()
}
func TestOpsErrorLoggerMiddleware_HardSkipsIngressRejection(t *testing.T) {
setupOpsErrorLogTestQueue(t, 4)
gin.SetMode(gin.TestMode)
settings := &ingressRejectSettingRepo{}
repo := &ingressRejectOpsRepo{}
ops := service.NewOpsService(repo, settings, nil, nil, nil, nil, nil, nil, nil, nil, nil)
// Construction may read unrelated runtime settings; only request-path reads matter here.
settings.getValueCalls = 0
router := gin.New()
router.Use(OpsErrorLoggerMiddleware(ops))
router.GET("/v1/messages", func(c *gin.Context) {
middleware2.MarkIngressRejected(c, middleware2.IngressRejectInvalidAPIKey)
c.JSON(http.StatusUnauthorized, gin.H{"code": "INVALID_API_KEY", "message": "Invalid API key"})
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
router.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.JSONEq(t, `{"code":"INVALID_API_KEY","message":"Invalid API key"}`, w.Body.String())
require.Zero(t, settings.getValueCalls, "ingress rejection must bypass monitoring settings reads")
require.Zero(t, repo.insertCalls, "ingress rejection must bypass inserts")
require.Zero(t, OpsErrorLogEnqueuedTotal(), "ingress rejection must not enter the error queue")
}
func TestNormalizeOpsPersistentUserAgentBoundsAndPreservesUTF8(t *testing.T) {
value := strings.Repeat("a", opsErrorLogMaxUserAgentBytes-1) + "你" + strings.Repeat("b", 32)
got := normalizeOpsPersistentUserAgent(" " + value + " ")
require.LessOrEqual(t, len(got), opsErrorLogMaxUserAgentBytes)
require.True(t, utf8.ValidString(got))
require.NotContains(t, got, "b")
}
// 就地(in-band) SSE 错误挂在已固化的 HTTP 200 流上:wire 状态码为 200,
// 常规 status>=400 采集路径不会触发。logOpsStreamError 必须据 MarkOpsStreamError
// 补记一条错误日志,且用 IntendedStatus(429) 分级、StatusCode 仍记 wire 的 200。
@@ -182,6 +320,36 @@ func TestLogOpsStreamError_RecordsInBandConcurrencyLimit(t *testing.T) {
require.Equal(t, "Concurrency limit exceeded for account, please retry later", job.entry.ErrorMessage)
}
func TestLogOpsStreamError_UpstreamFailureCountsTowardsSLA(t *testing.T) {
setupOpsErrorLogTestQueue(t, 4)
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
c.Set(opsModelKey, "gpt-5.6-sol")
service.MarkOpsStreamFailure(
c,
"upstream_error",
service.OpenAIUpstreamHTTP2StreamErrorCode,
"Upstream HTTP/2 stream failed",
http.StatusBadGateway,
)
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
logOpsStreamError(c, ops, http.StatusOK)
job := <-opsErrorLogQueue
require.NotNil(t, job.entry)
require.Equal(t, http.StatusBadGateway, job.entry.StatusCode)
require.Equal(t, "upstream_error", job.entry.ErrorType)
require.Equal(t, "upstream", job.entry.ErrorPhase)
require.Equal(t, "provider", job.entry.ErrorOwner)
require.False(t, job.entry.IsBusinessLimited)
require.Contains(t, job.entry.ErrorBody, service.OpenAIUpstreamHTTP2StreamErrorCode)
}
// 未标记流内错误时 logOpsStreamError 必须是 no-op(不误记正常的 200 流)。
func TestLogOpsStreamError_NoopWhenNotMarked(t *testing.T) {
setupOpsErrorLogTestQueue(t, 4)
@@ -0,0 +1,25 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestOpsCaptureWriterDoesNotCopyIngressRejectBody(t *testing.T) {
gin.SetMode(gin.TestMode)
context, _ := gin.CreateTestContext(httptest.NewRecorder())
writer := acquireOpsCaptureWriter(context.Writer)
defer releaseOpsCaptureWriter(writer)
writer.ctx = context
context.Writer = writer
middleware2.MarkIngressRejected(context, middleware2.IngressRejectInvalidAPIKey)
context.Status(http.StatusUnauthorized)
_, err := context.Writer.WriteString(`{"code":"INVALID_API_KEY","message":"Invalid API key"}`)
require.NoError(t, err)
require.Zero(t, writer.buf.Len())
}
+2
View File
@@ -120,12 +120,14 @@ func ProvideOpenAIGatewayHandler(
errorPassthroughService *service.ErrorPassthroughService,
contentModerationService *service.ContentModerationService,
opsService *service.OpsService,
grokQuotaService *service.GrokQuotaService,
cfg *config.Config,
coordinator *securityaudit.Coordinator,
) *OpenAIGatewayHandler {
h := NewOpenAIGatewayHandler(gatewayService, concurrencyService, billingCacheService, apiKeyService,
usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, cfg)
h.securityAuditCoordinator = coordinator
h.grokMediaEligibilityProber = grokQuotaService
return h
}
@@ -152,11 +152,25 @@ type AnthropicEventToResponsesState struct {
// For message output: accumulate text parts
ContentIndex int
// TextAccum accumulates the current text part so that output_text.done and
// content_part.done can carry the full text (deltas carry increments only).
TextAccum string
// For function_call: track per-output info
CurrentCallID string
CurrentName string
// Content of the currently open item, folded into Outputs when it closes.
CurrentContent []ResponsesContentPart // message
CurrentArgs string // function_call
CurrentSummary string // reasoning
// Outputs accumulates every closed output item so that response.completed
// can carry the full output list. The OpenAI SDK's get_final_response()
// parses the terminal event's response directly; without this, clients see
// an empty output_text.
Outputs []ResponsesOutput
// Usage from message_start / message_delta. InputTokens here follows
// Anthropic semantics (excludes cached tokens); they are added back when
// emitting the OpenAI Responses usage.
@@ -293,6 +307,22 @@ func anthToResHandleContentBlockStart(evt *AnthropicStreamEvent, state *Anthropi
}))
}
// response.content_part.added must precede the output_text.delta events
// for that part. The message item is added with content: [], and the
// OpenAI SDK's accumulating stream helper (client.responses.stream) only
// appends a content part when it sees content_part.added. Without it the
// following output_text.delta indexes output.content[content_index] and
// raises IndexError. Raw event iteration
// (responses.create(stream=True)) does not accumulate, which is why this
// went unnoticed.
events = append(events, makeResponsesEvent(state, "response.content_part.added", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
ContentIndex: state.ContentIndex,
ItemID: state.CurrentItemID,
Part: &ResponsesContentPart{Type: "output_text", Text: ""},
}))
state.TextAccum = ""
case "tool_use":
// Close previous item if any
events = append(events, closeCurrentResponsesItem(state)...)
@@ -327,6 +357,7 @@ func anthToResHandleContentBlockDelta(evt *AnthropicStreamEvent, state *Anthropi
if evt.Delta.Text == "" {
return nil
}
state.TextAccum += evt.Delta.Text
return []ResponsesStreamEvent{makeResponsesEvent(state, "response.output_text.delta", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
ContentIndex: state.ContentIndex,
@@ -338,6 +369,7 @@ func anthToResHandleContentBlockDelta(evt *AnthropicStreamEvent, state *Anthropi
if evt.Delta.Thinking == "" {
return nil
}
state.CurrentSummary += evt.Delta.Thinking
return []ResponsesStreamEvent{makeResponsesEvent(state, "response.reasoning_summary_text.delta", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
SummaryIndex: 0,
@@ -349,6 +381,7 @@ func anthToResHandleContentBlockDelta(evt *AnthropicStreamEvent, state *Anthropi
if evt.Delta.PartialJSON == "" {
return nil
}
state.CurrentArgs += evt.Delta.PartialJSON
return []ResponsesStreamEvent{makeResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
Delta: evt.Delta.PartialJSON,
@@ -393,12 +426,24 @@ func anthToResHandleContentBlockStop(evt *AnthropicStreamEvent, state *Anthropic
return events
case "message":
// Emit output_text.done (text block is done, but message item stays open for potential more blocks)
// Text block is done: emit output_text.done then content_part.done (the
// order OpenAI uses), both carrying the part's full text. The message
// item itself stays open since more blocks may follow.
text := state.TextAccum
state.TextAccum = ""
state.CurrentContent = append(state.CurrentContent, ResponsesContentPart{Type: "output_text", Text: text})
return []ResponsesStreamEvent{
makeResponsesEvent(state, "response.output_text.done", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
ContentIndex: state.ContentIndex,
ItemID: state.CurrentItemID,
Text: text,
}),
makeResponsesEvent(state, "response.content_part.done", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex,
ContentIndex: state.ContentIndex,
ItemID: state.CurrentItemID,
Part: &ResponsesContentPart{Type: "output_text", Text: text},
}),
}
}
@@ -454,24 +499,48 @@ func closeCurrentResponsesItem(state *AnthropicEventToResponsesState) []Response
return nil
}
itemType := state.CurrentItemType
itemID := state.CurrentItemID
// Assemble the full item: both output_item.done and response.completed must
// carry its content. Emitting only {type,id,status} makes SDK-side
// accumulation produce an empty output.
item := ResponsesOutput{
Type: state.CurrentItemType,
ID: state.CurrentItemID,
Status: "completed",
}
switch state.CurrentItemType {
case "message":
item.Role = "assistant"
item.Content = state.CurrentContent
case "function_call":
item.CallID = state.CurrentCallID
item.Name = state.CurrentName
args := state.CurrentArgs
if args == "" {
args = "{}"
}
item.Arguments = args
case "reasoning":
if state.CurrentSummary != "" {
item.Summary = []ResponsesSummary{{Type: "summary_text", Text: state.CurrentSummary}}
}
}
state.Outputs = append(state.Outputs, item)
// Reset
state.CurrentItemType = ""
state.CurrentItemID = ""
state.CurrentCallID = ""
state.CurrentName = ""
state.CurrentContent = nil
state.CurrentArgs = ""
state.CurrentSummary = ""
state.TextAccum = ""
state.OutputIndex++
state.ContentIndex = 0
return []ResponsesStreamEvent{makeResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
OutputIndex: state.OutputIndex - 1, // Use the index before increment
Item: &ResponsesOutput{
Type: itemType,
ID: itemID,
Status: "completed",
},
Item: &item,
})}
}
@@ -519,6 +588,14 @@ func makeResponsesCompletedEvent(
eventType = "response.incomplete"
}
// Carry the output items accumulated over the stream. The SDK's
// get_final_response() reads them straight from the terminal event, so an
// empty list leaves clients with an empty result.
outputs := state.Outputs
if outputs == nil {
outputs = []ResponsesOutput{}
}
return ResponsesStreamEvent{
Type: eventType,
SequenceNumber: seq,
@@ -527,7 +604,7 @@ func makeResponsesCompletedEvent(
Object: "response",
Model: state.Model,
Status: status,
Output: []ResponsesOutput{},
Output: outputs,
Usage: usage,
IncompleteDetails: incompleteDetails,
},
@@ -0,0 +1,187 @@
package apicompat
import "testing"
// TestAnthropicEventToResponses_TextEmitsContentPart pins that a message text
// stream emits response.content_part.added, and that it precedes the first
// output_text.delta for that part.
//
// Why: the OpenAI SDK's accumulating stream helper (client.responses.stream)
// only appends a content part to the message item when it sees
// content_part.added. The item is added with content: [], so a missing event
// makes the following output_text.delta index output.content[content_index] and
// raise IndexError. Raw event iteration does not accumulate, so a regression
// here is easy to miss.
func TestAnthropicEventToResponses_TextEmitsContentPart(t *testing.T) {
state := NewAnthropicEventToResponsesState()
state.Model = "claude-sonnet-4-5"
var types []string
feed := func(evt *AnthropicStreamEvent) {
for _, out := range AnthropicEventToResponsesEvents(evt, state) {
types = append(types, out.Type)
}
}
idx := 0
feed(&AnthropicStreamEvent{Type: "message_start", Message: &AnthropicResponse{ID: "msg_1", Model: "claude-sonnet-4-5"}})
feed(&AnthropicStreamEvent{Type: "content_block_start", Index: &idx, ContentBlock: &AnthropicContentBlock{Type: "text"}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{Type: "text_delta", Text: "Hel"}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{Type: "text_delta", Text: "lo"}})
feed(&AnthropicStreamEvent{Type: "content_block_stop", Index: &idx})
feed(&AnthropicStreamEvent{Type: "message_stop"})
posOf := func(target string) int {
for i, ty := range types {
if ty == target {
return i
}
}
return -1
}
partAdded := posOf("response.content_part.added")
firstDelta := posOf("response.output_text.delta")
if partAdded < 0 {
t.Fatalf("response.content_part.added was not emitted; got %v", types)
}
if firstDelta < 0 {
t.Fatalf("response.output_text.delta was not emitted; got %v", types)
}
if partAdded > firstDelta {
t.Errorf("content_part.added must precede the first output_text.delta; got %v", types)
}
if posOf("response.content_part.done") < 0 {
t.Errorf("response.content_part.done was not emitted; got %v", types)
}
}
// TestAnthropicEventToResponses_DoneEventsCarryFullText pins that done events
// carry the part's full text (deltas carry increments only).
func TestAnthropicEventToResponses_DoneEventsCarryFullText(t *testing.T) {
state := NewAnthropicEventToResponsesState()
state.Model = "claude-sonnet-4-5"
var events []ResponsesStreamEvent
feed := func(evt *AnthropicStreamEvent) {
events = append(events, AnthropicEventToResponsesEvents(evt, state)...)
}
idx := 0
feed(&AnthropicStreamEvent{Type: "message_start", Message: &AnthropicResponse{ID: "msg_1"}})
feed(&AnthropicStreamEvent{Type: "content_block_start", Index: &idx, ContentBlock: &AnthropicContentBlock{Type: "text"}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{Type: "text_delta", Text: "Hello "}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{Type: "text_delta", Text: "world"}})
feed(&AnthropicStreamEvent{Type: "content_block_stop", Index: &idx})
const want = "Hello world"
var sawTextDone, sawPartDone bool
for _, e := range events {
switch e.Type {
case "response.output_text.done":
sawTextDone = true
if e.Text != want {
t.Errorf("output_text.done text = %q, want %q", e.Text, want)
}
case "response.content_part.done":
sawPartDone = true
if e.Part == nil || e.Part.Text != want {
t.Errorf("content_part.done part = %+v, want text %q", e.Part, want)
}
}
}
if !sawTextDone || !sawPartDone {
t.Errorf("missing done events: output_text.done=%v content_part.done=%v", sawTextDone, sawPartDone)
}
}
// TestAnthropicEventToResponses_CompletedCarriesOutput pins that
// response.completed carries the full output list. The SDK's
// get_final_response() and tracing integrations parse the terminal event's
// response directly; an empty output leaves them with nothing (the text still
// renders from deltas, which is why this is invisible when only watching the
// stream).
func TestAnthropicEventToResponses_CompletedCarriesOutput(t *testing.T) {
state := NewAnthropicEventToResponsesState()
state.Model = "claude-sonnet-4-5"
var events []ResponsesStreamEvent
feed := func(evt *AnthropicStreamEvent) {
events = append(events, AnthropicEventToResponsesEvents(evt, state)...)
}
idx := 0
feed(&AnthropicStreamEvent{Type: "message_start", Message: &AnthropicResponse{ID: "msg_1"}})
feed(&AnthropicStreamEvent{Type: "content_block_start", Index: &idx, ContentBlock: &AnthropicContentBlock{Type: "text"}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{Type: "text_delta", Text: "4826"}})
feed(&AnthropicStreamEvent{Type: "content_block_stop", Index: &idx})
feed(&AnthropicStreamEvent{Type: "message_stop"})
var completed *ResponsesStreamEvent
for i := range events {
if events[i].Type == "response.completed" {
completed = &events[i]
}
}
if completed == nil || completed.Response == nil {
t.Fatalf("response.completed was not emitted")
}
if len(completed.Response.Output) == 0 {
t.Fatalf("response.completed carries an empty output; clients would see no result")
}
msg := completed.Response.Output[0]
if msg.Type != "message" || len(msg.Content) == 0 {
t.Fatalf("output[0] = %+v, want a message with content", msg)
}
if msg.Content[0].Text != "4826" {
t.Errorf("output[0].content[0].text = %q, want %q", msg.Content[0].Text, "4826")
}
}
// TestAnthropicEventToResponses_ToolCallCompletedCarriesArguments pins that a
// function call's accumulated arguments survive into output_item.done and
// response.completed.
func TestAnthropicEventToResponses_ToolCallCompletedCarriesArguments(t *testing.T) {
state := NewAnthropicEventToResponsesState()
state.Model = "claude-sonnet-4-5"
var events []ResponsesStreamEvent
feed := func(evt *AnthropicStreamEvent) {
events = append(events, AnthropicEventToResponsesEvents(evt, state)...)
}
idx := 0
feed(&AnthropicStreamEvent{Type: "message_start", Message: &AnthropicResponse{ID: "msg_1"}})
feed(&AnthropicStreamEvent{Type: "content_block_start", Index: &idx, ContentBlock: &AnthropicContentBlock{
Type: "tool_use", ID: "toolu_1", Name: "get_weather",
}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{
Type: "input_json_delta", PartialJSON: `{"city":`,
}})
feed(&AnthropicStreamEvent{Type: "content_block_delta", Index: &idx, Delta: &AnthropicDelta{
Type: "input_json_delta", PartialJSON: `"SH"}`,
}})
feed(&AnthropicStreamEvent{Type: "content_block_stop", Index: &idx})
feed(&AnthropicStreamEvent{Type: "message_stop"})
var completed *ResponsesStreamEvent
for i := range events {
if events[i].Type == "response.completed" {
completed = &events[i]
}
}
if completed == nil || completed.Response == nil || len(completed.Response.Output) == 0 {
t.Fatalf("response.completed carries no output")
}
fc := completed.Response.Output[0]
if fc.Type != "function_call" {
t.Fatalf("output[0].type = %q, want function_call", fc.Type)
}
if fc.Arguments != `{"city":"SH"}` {
t.Errorf("arguments = %q, want %q", fc.Arguments, `{"city":"SH"}`)
}
if fc.Name != "get_weather" {
t.Errorf("name = %q, want get_weather", fc.Name)
}
}
+7 -75
View File
@@ -8,40 +8,10 @@ import (
"github.com/gin-gonic/gin"
)
// GetClientIP 从 Gin Context 中提取客户端真实 IP 地址。
// 按以下优先级检查 Header:
// 1. CF-Connecting-IP (Cloudflare)
// 2. X-Real-IP (Nginx)
// 3. X-Forwarded-For (取第一个非私有 IP)
// 4. c.ClientIP() (Gin 内置方法)
// GetClientIP resolves a client address only through Gin's configured trusted
// proxy chain. Forwarding headers from a direct or untrusted peer are ignored.
func GetClientIP(c *gin.Context) string {
// 1. Cloudflare
if ip := c.GetHeader("CF-Connecting-IP"); ip != "" {
return normalizeIP(ip)
}
// 2. Nginx X-Real-IP
if ip := c.GetHeader("X-Real-IP"); ip != "" {
return normalizeIP(ip)
}
// 3. X-Forwarded-For (多个 IP 时取第一个公网 IP)
if xff := c.GetHeader("X-Forwarded-For"); xff != "" {
ips := strings.Split(xff, ",")
for _, ip := range ips {
ip = strings.TrimSpace(ip)
if ip != "" && !isPrivateIP(ip) {
return normalizeIP(ip)
}
}
// 如果都是私有 IP,返回第一个
if len(ips) > 0 {
return normalizeIP(strings.TrimSpace(ips[0]))
}
}
// 4. Gin 内置方法
return normalizeIP(c.ClientIP())
return GetTrustedClientIP(c)
}
// GetTrustedClientIP 从 Gin 的可信代理解析链提取客户端 IP。
@@ -54,14 +24,10 @@ func GetTrustedClientIP(c *gin.Context) string {
return normalizeIP(c.ClientIP())
}
// GetSecurityClientIP 返回安全敏感场景(API Key IP 限制、审计日志、会话 IP/UA 绑定)
// 使用的客户端 IP。trustForwarded 对应系统设置「信任反代传递的客户端 IP」:
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
// 关闭时走 Gin trusted_proxies 解析链。
func GetSecurityClientIP(c *gin.Context, trustForwarded bool) string {
if trustForwarded {
return GetClientIP(c)
}
// GetSecurityClientIP returns the address resolved through Gin's configured
// trusted-proxy chain. The legacy toggle is retained for configuration/API
// compatibility, but never makes raw forwarding headers trustworthy by itself.
func GetSecurityClientIP(c *gin.Context, _ bool) string {
return GetTrustedClientIP(c)
}
@@ -75,9 +41,6 @@ func normalizeIP(ip string) string {
return ip
}
// privateNets 预编译私有 IP CIDR 块,避免每次调用 isPrivateIP 时重复解析
var privateNets []*net.IPNet
// CompiledIPRules 表示预编译的 IP 匹配规则。
// PatternCount 记录原始规则数量,用于保留“规则存在但全无效”时的行为语义。
type CompiledIPRules struct {
@@ -86,23 +49,6 @@ type CompiledIPRules struct {
PatternCount int
}
func init() {
for _, cidr := range []string{
"10.0.0.0/8",
"172.16.0.0/12",
"192.168.0.0/16",
"127.0.0.0/8",
"::1/128",
"fc00::/7",
} {
_, block, err := net.ParseCIDR(cidr)
if err != nil {
panic("invalid CIDR: " + cidr)
}
privateNets = append(privateNets, block)
}
}
// CompileIPRules 将 IP/CIDR 字符串规则预编译为可复用结构。
// 非法规则会被忽略,但 PatternCount 会保留原始规则条数。
func CompileIPRules(patterns []string) *CompiledIPRules {
@@ -150,20 +96,6 @@ func matchesCompiledRules(parsedIP net.IP, rules *CompiledIPRules) bool {
return false
}
// isPrivateIP 检查 IP 是否为私有地址。
func isPrivateIP(ipStr string) bool {
ip := net.ParseIP(ipStr)
if ip == nil {
return false
}
for _, block := range privateNets {
if block.Contains(ip) {
return true
}
}
return false
}
// MatchesPattern 检查 IP 是否匹配指定的模式(支持单个 IP 或 CIDR)。
// pattern 可以是:
// - 单个 IP: "192.168.1.100"
+18 -45
View File
@@ -10,48 +10,6 @@ import (
"github.com/stretchr/testify/require"
)
func TestIsPrivateIP(t *testing.T) {
tests := []struct {
name string
ip string
expected bool
}{
// 私有 IPv4
{"10.x 私有地址", "10.0.0.1", true},
{"10.x 私有地址段末", "10.255.255.255", true},
{"172.16.x 私有地址", "172.16.0.1", true},
{"172.31.x 私有地址", "172.31.255.255", true},
{"192.168.x 私有地址", "192.168.1.1", true},
{"127.0.0.1 本地回环", "127.0.0.1", true},
{"127.x 回环段", "127.255.255.255", true},
// 公网 IPv4
{"8.8.8.8 公网 DNS", "8.8.8.8", false},
{"1.1.1.1 公网", "1.1.1.1", false},
{"172.15.255.255 非私有", "172.15.255.255", false},
{"172.32.0.0 非私有", "172.32.0.0", false},
{"11.0.0.1 公网", "11.0.0.1", false},
// IPv6
{"::1 IPv6 回环", "::1", true},
{"fc00:: IPv6 私有", "fc00::1", true},
{"fd00:: IPv6 私有", "fd00::1", true},
{"2001:db8::1 IPv6 公网", "2001:db8::1", false},
// 无效输入
{"空字符串", "", false},
{"非法字符串", "not-an-ip", false},
{"不完整 IP", "192.168", false},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := isPrivateIP(tc.ip)
require.Equal(t, tc.expected, got, "isPrivateIP(%q)", tc.ip)
})
}
}
func TestGetTrustedClientIPUsesGinClientIP(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -95,7 +53,7 @@ func TestCheckIPRestrictionWithCompiledRules_InvalidWhitelistStillDenies(t *test
require.Equal(t, "access denied", reason)
}
func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
func TestGetSecurityClientIPNeverTrustsHeadersFromUntrustedPeer(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, tc := range []struct {
@@ -103,8 +61,8 @@ func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
trustForwarded bool
want string
}{
{name: "trust disabled uses trusted proxy chain", trustForwarded: false, want: "9.9.9.9"},
{name: "trust enabled uses forwarded header", trustForwarded: true, want: "1.2.3.4"},
{name: "legacy toggle disabled", trustForwarded: false, want: "9.9.9.9"},
{name: "legacy toggle enabled", trustForwarded: true, want: "9.9.9.9"},
} {
t.Run(tc.name, func(t *testing.T) {
r := gin.New()
@@ -124,3 +82,18 @@ func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
})
}
}
func TestGetSecurityClientIPUsesConfiguredTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
r := gin.New()
require.NoError(t, r.SetTrustedProxies([]string{"9.9.9.9"}))
r.GET("/t", func(c *gin.Context) { c.String(200, GetSecurityClientIP(c, true)) })
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/t", nil)
req.RemoteAddr = "9.9.9.9:12345"
req.Header.Set("X-Forwarded-For", "1.2.3.4")
r.ServeHTTP(w, req)
require.Equal(t, "1.2.3.4", w.Body.String())
}
+4
View File
@@ -26,6 +26,10 @@ const (
LevelWarn = zapcore.WarnLevel
LevelError = zapcore.ErrorLevel
LevelFatal = zapcore.FatalLevel
// OpsSystemLogSkipField keeps an event in the standard logger while
// preventing the database-backed Ops system-log sink from indexing it.
OpsSystemLogSkipField = "ops_system_log_skip"
)
type Sink interface {
+129 -12
View File
@@ -893,6 +893,20 @@ func (r *accountRepository) ListOpsAccountsForStats(ctx context.Context, platfor
func accountListOrder(params pagination.PaginationParams) []func(*entsql.Selector) {
sortBy := strings.ToLower(strings.TrimSpace(params.SortBy))
sortOrder := params.NormalizedSortOrder(pagination.SortOrderAsc)
if sortBy == "upstream_billing_rate" {
direction := "ASC"
tieOrder := entsql.Asc
if sortOrder == pagination.SortOrderDesc {
direction = "DESC"
tieOrder = entsql.Desc
}
return []func(*entsql.Selector){func(s *entsql.Selector) {
extra := s.C(dbaccount.FieldExtra)
expression := upstreamBillingRateSortExpression(extra)
s.OrderExpr(entsql.Expr(expression + " " + direction + " NULLS LAST"))
s.OrderBy(tieOrder(s.C(dbaccount.FieldID)))
}}
}
field := dbaccount.FieldName
defaultOrder := true
@@ -934,6 +948,40 @@ func accountListOrder(params pagination.PaginationParams) []func(*entsql.Selecto
return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(dbaccount.FieldID)}
}
func upstreamBillingRateSortExpression(extra string) string {
status := extra + " #>> '{upstream_billing_probe,status}'"
effectiveJSON := extra + " #> '{upstream_billing_probe,data,effective_rate_multiplier}'"
effective := extra + " #>> '{upstream_billing_probe,data,effective_rate_multiplier}'"
resolvedJSON := extra + " #> '{upstream_billing_probe,data,resolved_rate_multiplier}'"
resolved := extra + " #>> '{upstream_billing_probe,data,resolved_rate_multiplier}'"
peakEnabledJSON := extra + " #> '{upstream_billing_probe,data,peak_rate_enabled}'"
peakEnabled := extra + " #>> '{upstream_billing_probe,data,peak_rate_enabled}'"
peakStart := extra + " #>> '{upstream_billing_probe,data,peak_start}'"
peakEnd := extra + " #>> '{upstream_billing_probe,data,peak_end}'"
peakMultiplierJSON := extra + " #> '{upstream_billing_probe,data,peak_rate_multiplier}'"
peakMultiplier := extra + " #>> '{upstream_billing_probe,data,peak_rate_multiplier}'"
peakMultiplierValue := "(CASE WHEN jsonb_typeof(" + peakMultiplierJSON + ") = 'number' THEN (" + peakMultiplier + ")::numeric END)"
billingScope := extra + " #>> '{upstream_billing_probe,data,billing_scope}'"
timezone := extra + " #>> '{upstream_billing_probe,data,timezone}'"
validClock := "'^([01][0-9]|2[0-3]):[0-5][0-9]$'"
startMinute := "(CASE WHEN " + peakStart + " ~ " + validClock + " THEN split_part(" + peakStart + ", ':', 1)::numeric * 60 + split_part(" + peakStart + ", ':', 2)::numeric END)"
endMinute := "(CASE WHEN " + peakEnd + " ~ " + validClock + " THEN split_part(" + peakEnd + ", ':', 1)::numeric * 60 + split_part(" + peakEnd + ", ':', 2)::numeric END)"
localMinute := "(EXTRACT(HOUR FROM (CURRENT_TIMESTAMP AT TIME ZONE (" + timezone + "))) * 60 + EXTRACT(MINUTE FROM (CURRENT_TIMESTAMP AT TIME ZONE (" + timezone + "))))"
validPeakWindow := peakStart + " ~ " + validClock + " AND " +
peakEnd + " ~ " + validClock + " AND " +
startMinute + " < " + endMinute
validPeakConfig := validPeakWindow + " AND " + peakMultiplierValue + " >= 0 AND " +
"EXISTS (SELECT 1 FROM pg_timezone_names WHERE name = " + timezone + ")"
dynamicRate := "CASE WHEN " + peakEnabled + " = 'false' THEN (" + resolved + ")::numeric WHEN " + peakEnabled + " = 'true' AND " + validPeakConfig +
" THEN (" + resolved + ")::numeric * CASE WHEN " + localMinute + " >= " + startMinute + " AND " + localMinute + " < " + endMinute +
" THEN " + peakMultiplierValue + " ELSE 1 END ELSE NULL END"
legacySnapshot := "jsonb_typeof(" + resolvedJSON + ") IS NULL AND jsonb_typeof(" + peakEnabledJSON + ") IS NULL"
return "CASE WHEN " + status + " IN ('ok', 'failed') AND (jsonb_typeof(" + resolvedJSON + ") = 'number' OR jsonb_typeof(" + effectiveJSON + ") = 'number') THEN CASE WHEN jsonb_typeof(" +
resolvedJSON + ") = 'number' AND jsonb_typeof(" + peakEnabledJSON + ") = 'boolean' THEN CASE WHEN " + billingScope + " = 'token' THEN " + dynamicRate + " ELSE NULL END WHEN " + legacySnapshot +
" AND jsonb_typeof(" + effectiveJSON + ") = 'number' THEN (" + effective + ")::numeric END END"
}
func (r *accountRepository) ListByGroup(ctx context.Context, groupID int64) ([]service.Account, error) {
accounts, err := r.queryAccountsByGroup(ctx, groupID, accountGroupQueryOptions{
status: service.StatusActive,
@@ -1877,6 +1925,46 @@ func (r *accountRepository) ListSchedulableByGroupIDAndPlatforms(ctx context.Con
})
}
// ListModelAvailabilityCandidates returns the persistently configured account
// pool used to decide whether a model is supported. Unlike scheduling queries,
// it intentionally ignores transient runtime state (rate limits, overload,
// temporary unschedulability, and expiry windows).
func (r *accountRepository) ListModelAvailabilityCandidates(
ctx context.Context,
groupID *int64,
platforms []string,
includeGrouped bool,
) ([]service.Account, error) {
if len(platforms) == 0 {
return []service.Account{}, nil
}
if groupID != nil {
return r.queryAccountsByGroup(ctx, *groupID, accountGroupQueryOptions{
status: service.StatusActive,
schedulable: true,
ignoreTransientState: true,
platforms: platforms,
})
}
preds := []dbpredicate.Account{
dbaccount.StatusEQ(service.StatusActive),
dbaccount.SchedulableEQ(true),
dbaccount.PlatformIn(platforms...),
}
if !includeGrouped {
preds = append(preds, dbaccount.Not(dbaccount.HasAccountGroups()))
}
accounts, err := r.client.Account.Query().
Where(preds...).
Order(dbent.Asc(dbaccount.FieldPriority)).
All(ctx)
if err != nil {
return nil, err
}
return r.accountsToService(ctx, accounts)
}
func (r *accountRepository) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
now := time.Now()
_, err := r.client.Account.Update().
@@ -2575,6 +2663,12 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
args = append(args, *updates.Schedulable)
idx++
}
if updates.ProbeEnabled != nil {
if updates.Extra == nil {
updates.Extra = make(map[string]any)
}
updates.Extra[service.UpstreamBillingProbeEnabledExtraKey] = *updates.ProbeEnabled
}
// JSONB 需要合并而非覆盖,使用 raw SQL 保持旧行为。
if len(updates.Credentials) > 0 {
payload, err := json.Marshal(updates.Credentials)
@@ -2605,8 +2699,14 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
setClauses = append(setClauses, "updated_at = NOW()")
query := "UPDATE accounts SET " + joinClauses(setClauses, ", ") + " WHERE id = ANY($" + itoa(idx) + ") AND deleted_at IS NULL"
whereClause := " WHERE id = ANY($" + itoa(idx) + ") AND deleted_at IS NULL"
args = append(args, pq.Array(ids))
idx++
if updates.ProbeEnabled != nil {
whereClause += " AND platform = $" + itoa(idx) + " AND type = $" + itoa(idx+1)
args = append(args, service.PlatformOpenAI, service.AccountTypeAPIKey)
}
query := "UPDATE accounts SET " + joinClauses(setClauses, ", ") + whereClause
baseCtx := ctx
contextTx := dbent.TxFromContext(ctx)
@@ -2635,6 +2735,20 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
if err != nil {
return 0, err
}
if updates.ProbeEnabled != nil {
expectedRows := int64(0)
seenIDs := make(map[int64]struct{}, len(ids))
for _, id := range ids {
if _, seen := seenIDs[id]; seen {
continue
}
seenIDs[id] = struct{}{}
expectedRows++
}
if rows != expectedRows {
return 0, service.ErrUpstreamBillingProbeAccountInvalid
}
}
if rows > 0 {
payload := map[string]any{"account_ids": ids}
if err := enqueueSchedulerOutbox(ctx, exec, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil {
@@ -2662,9 +2776,10 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
}
type accountGroupQueryOptions struct {
status string
schedulable bool
platforms []string // 允许的多个平台,空切片表示不进行平台过滤
status string
schedulable bool
ignoreTransientState bool
platforms []string // 允许的多个平台,空切片表示不进行平台过滤
}
func (r *accountRepository) queryAccountsByGroup(ctx context.Context, groupID int64, opts accountGroupQueryOptions) ([]service.Account, error) {
@@ -2681,14 +2796,16 @@ func (r *accountRepository) queryAccountsByGroup(ctx context.Context, groupID in
preds = append(preds, dbaccount.PlatformIn(opts.platforms...))
}
if opts.schedulable {
now := time.Now()
preds = append(preds,
dbaccount.SchedulableEQ(true),
tempUnschedulablePredicate(),
notExpiredPredicate(now),
dbaccount.Or(dbaccount.OverloadUntilIsNil(), dbaccount.OverloadUntilLTE(now)),
dbaccount.Or(dbaccount.RateLimitResetAtIsNil(), dbaccount.RateLimitResetAtLTE(now)),
)
preds = append(preds, dbaccount.SchedulableEQ(true))
if !opts.ignoreTransientState {
now := time.Now()
preds = append(preds,
tempUnschedulablePredicate(),
notExpiredPredicate(now),
dbaccount.Or(dbaccount.OverloadUntilIsNil(), dbaccount.OverloadUntilLTE(now)),
dbaccount.Or(dbaccount.RateLimitResetAtIsNil(), dbaccount.RateLimitResetAtLTE(now)),
)
}
}
if len(preds) > 0 {
@@ -0,0 +1,58 @@
package repository
import (
"context"
"strings"
"testing"
"github.com/DATA-DOG/go-sqlmock"
dbent "github.com/Wei-Shaw/sub2api/ent"
_ "github.com/Wei-Shaw/sub2api/ent/runtime"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
"entgo.io/ent/dialect"
entsql "entgo.io/ent/dialect/sql"
)
func TestListModelAvailabilityCandidates_GroupQueryIgnoresTransientState(t *testing.T) {
var capturedSQL string
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(captureEntQueryMatcher{actual: &capturedSQL}))
require.NoError(t, err)
t.Cleanup(func() { _ = db.Close() })
driver := entsql.OpenDB(dialect.Postgres, db)
client := dbent.NewClient(dbent.Driver(driver))
t.Cleanup(func() { _ = client.Close() })
repo := newAccountRepositoryWithSQL(client, db, nil)
mock.ExpectQuery("model availability candidates").
WillReturnRows(sqlmock.NewRows([]string{"id"}))
groupID := int64(42)
accounts, err := repo.ListModelAvailabilityCandidates(
context.Background(),
&groupID,
[]string{service.PlatformAnthropic},
false,
)
require.NoError(t, err)
require.Empty(t, accounts)
require.NoError(t, mock.ExpectationsWereMet())
normalized := normalizeSQLWhitespace(capturedSQL)
_, whereClause, found := strings.Cut(normalized, " WHERE ")
require.True(t, found, "expected WHERE clause in query: %s", normalized)
whereClause, _, _ = strings.Cut(whereClause, " ORDER BY ")
for _, configuredPredicate := range []string{"group_id", "status", "schedulable", "platform"} {
require.Contains(t, whereClause, configuredPredicate)
}
for _, transientPredicate := range []string{
"rate_limit_reset_at",
"overload_until",
"temp_unschedulable_until",
"expires_at",
"auto_pause_on_expired",
} {
require.NotContains(t, whereClause, transientPredicate, "configured-state diagnosis must not filter transient predicate %q", transientPredicate)
}
}
@@ -3,6 +3,9 @@
package repository
import (
"fmt"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
"github.com/Wei-Shaw/sub2api/internal/service"
)
@@ -33,3 +36,105 @@ func (s *AccountRepoSuite) TestListWithFilters_SortByPriorityDesc() {
s.Require().Equal("high-priority", accounts[0].Name)
s.Require().Equal("low-priority", accounts[1].Name)
}
func (s *AccountRepoSuite) TestListWithFilters_SortByUpstreamBillingRateWithMissingLast() {
makeAccount := func(name, status string, rate any) {
extra := map[string]any{}
if rate != nil {
extra[service.UpstreamBillingProbeExtraKey] = map[string]any{
"status": status,
"data": map[string]any{"effective_rate_multiplier": rate},
}
}
mustCreateAccount(s.T(), s.client, &service.Account{Name: name, Extra: extra})
}
makeAccount("high-rate", service.UpstreamBillingProbeStatusOK, 0.8)
makeAccount("low-rate", service.UpstreamBillingProbeStatusOK, 0.03)
makeAccount("missing-rate", "", nil)
makeAccount("unsupported-with-retained-rate", service.UpstreamBillingProbeStatusUnsupported, 0.01)
for _, tc := range []struct {
order string
want []string
}{
{order: "asc", want: []string{"low-rate", "high-rate", "missing-rate", "unsupported-with-retained-rate"}},
{order: "desc", want: []string{"high-rate", "low-rate", "unsupported-with-retained-rate", "missing-rate"}},
} {
accounts, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{
Page: 1, PageSize: 10, SortBy: "upstream_billing_rate", SortOrder: tc.order,
}, "", "", "", "", 0, "")
s.Require().NoError(err)
s.Require().Len(accounts, 4)
for i, name := range tc.want {
s.Require().Equal(name, accounts[i].Name)
}
}
}
func (s *AccountRepoSuite) TestListWithFilters_SortByCurrentUpstreamBillingRateDuringPeak() {
now := time.Now()
locations := []string{"UTC", "Asia/Shanghai", "America/New_York", "Europe/London"}
var timezone string
var minute int
for _, name := range locations {
location, err := time.LoadLocation(name)
s.Require().NoError(err)
local := now.In(location)
candidate := local.Hour()*60 + local.Minute()
if candidate >= 2 && candidate <= 1436 {
timezone = name
minute = candidate
break
}
}
s.Require().NotEmpty(timezone)
peakStart := fmt.Sprintf("%02d:%02d", (minute-2)/60, (minute-2)%60)
peakEnd := fmt.Sprintf("%02d:%02d", (minute+3)/60, (minute+3)%60)
mustCreateAccount(s.T(), s.client, &service.Account{
Name: "current-peak-rate",
Extra: map[string]any{
service.UpstreamBillingProbeExtraKey: map[string]any{
"status": service.UpstreamBillingProbeStatusOK,
"data": map[string]any{
"billing_scope": "token",
"resolved_rate_multiplier": 1.0,
"effective_rate_multiplier": 1.0,
"peak_rate_enabled": true,
"peak_start": peakStart,
"peak_end": peakEnd,
"peak_rate_multiplier": 10.0,
"timezone": timezone,
},
},
},
})
mustCreateAccount(s.T(), s.client, &service.Account{
Name: "current-off-peak-rate",
Extra: map[string]any{
service.UpstreamBillingProbeExtraKey: map[string]any{
"status": service.UpstreamBillingProbeStatusOK,
"data": map[string]any{
"effective_rate_multiplier": 5.0,
},
},
},
})
for _, tc := range []struct {
order string
want []string
}{
{order: "asc", want: []string{"current-off-peak-rate", "current-peak-rate"}},
{order: "desc", want: []string{"current-peak-rate", "current-off-peak-rate"}},
} {
accounts, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{
Page: 1, PageSize: 10, SortBy: "upstream_billing_rate", SortOrder: tc.order,
}, "", "", "", "", 0, "")
s.Require().NoError(err)
s.Require().Len(accounts, 2)
for i, name := range tc.want {
s.Require().Equal(name, accounts[i].Name)
}
}
}
@@ -163,6 +163,46 @@ func TestBulkUpdateNilProbeRemovesKeyInsteadOfWritingJSONNull(t *testing.T) {
require.Contains(t, normalizeSQLWhitespace(exec.execQueries[0]), "- 'upstream_billing_probe'")
}
func TestBulkUpdateDisablingProbeRemovesSnapshot(t *testing.T) {
exec := &recordingSQLExecutor{result: rowsAffectedResult(1)}
repo := newAccountRepositoryWithSQL(nil, exec, nil)
_, err := repo.BulkUpdate(context.Background(), []int64{27}, service.AccountBulkUpdate{
Extra: map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false},
})
require.NoError(t, err)
require.NotEmpty(t, exec.execQueries)
require.Contains(t, normalizeSQLWhitespace(exec.execQueries[0]), "- 'upstream_billing_probe'")
payload, ok := exec.execArgs[0][0].([]byte)
require.True(t, ok)
require.Equal(t, `{"upstream_billing_probe_enabled":false}`, string(payload))
}
func TestBulkUpdateProbeEligibilityMismatchRollsBack(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
t.Cleanup(func() { _ = db.Close() })
client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db)))
t.Cleanup(func() { _ = client.Close() })
enabled := true
mock.ExpectBegin()
mock.ExpectExec(`(?s)UPDATE accounts SET extra = .* WHERE id = ANY\(\$2\) AND deleted_at IS NULL AND platform = \$3 AND type = \$4`).
WithArgs(sqlmock.AnyArg(), `{27,28}`, service.PlatformOpenAI, service.AccountTypeAPIKey).
WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectRollback()
repo := newAccountRepositoryWithSQL(client, db, nil)
rows, err := repo.BulkUpdate(context.Background(), []int64{27, 28}, service.AccountBulkUpdate{
ProbeEnabled: &enabled,
})
require.ErrorIs(t, err, service.ErrUpstreamBillingProbeAccountInvalid)
require.Zero(t, rows)
require.NoError(t, mock.ExpectationsWereMet())
}
func TestUpdateCredentialsAtomicallyClearsProbeForOpenAIAPIKeyIdentityChange(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
+18 -21
View File
@@ -110,28 +110,25 @@ func (c *apiKeyCache) SubscribeAuthCacheInvalidation(ctx context.Context, handle
return fmt.Errorf("subscribe to auth cache invalidation: %w", err)
}
go func() {
defer func() {
if err := pubsub.Close(); err != nil {
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
}
}()
ch := pubsub.Channel()
for {
select {
case <-ctx.Done():
return
case msg, ok := <-ch:
if !ok {
return
}
if msg != nil {
handler(msg.Payload)
}
}
defer func() {
if err := pubsub.Close(); err != nil {
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
}
}()
service.NotifyAuthCacheSubscriptionReady(ctx)
return nil
ch := pubsub.Channel()
for {
select {
case <-ctx.Done():
return ctx.Err()
case msg, ok := <-ch:
if !ok {
return errors.New("auth cache invalidation pubsub channel closed")
}
if msg != nil {
handler(msg.Payload)
}
}
}
}
@@ -0,0 +1,49 @@
package repository
import (
"context"
"errors"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func TestAPIKeyCacheSubscriber_BlocksUntilContextCancellation(t *testing.T) {
server := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
defer func() { _ = client.Close() }()
cache := NewAPIKeyCache(client)
ctx, cancel := context.WithCancel(context.Background())
received := make(chan string, 1)
returned := make(chan error, 1)
go func() {
returned <- cache.SubscribeAuthCacheInvalidation(ctx, func(value string) { received <- value })
}()
var value string
require.Eventually(t, func() bool {
require.NoError(t, client.Publish(context.Background(), authCacheInvalidateChannel, "hash").Err())
select {
case value = <-received:
return true
default:
return false
}
}, time.Second, 10*time.Millisecond)
require.Equal(t, "hash", value)
select {
case err := <-returned:
t.Fatalf("subscriber returned while connection was active: %v", err)
default:
}
cancel()
select {
case err := <-returned:
require.True(t, errors.Is(err, context.Canceled))
case <-time.After(time.Second):
t.Fatal("subscriber did not stop after context cancellation")
}
}
+6 -18
View File
@@ -326,16 +326,14 @@ func (r *apiKeyRepository) Delete(ctx context.Context, id int64) error {
return nil
}
// DeleteWithAudit 在同一事务内:
// 1. 把(明文 key、所有者、key 名称)写入 deleted_api_key_audits;
// 2. 软删除该 key(tombstone 覆盖 key 列以释放唯一约束)。
//
// 保证"被删除的 key 一定能反查到所有者"。事务模式与 group_repo.DeleteCascade 一致。
// DeleteWithAudit keeps the legacy method name for rolling-upgrade compatibility.
// It atomically tombstones and soft-deletes the key without retaining credential
// material. Tombstoning releases the unique key value for safe reuse.
func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error {
tombstoneKey := fmt.Sprintf("__deleted__%d__%d", id, time.Now().UnixNano())
if existingTx := dbent.TxFromContext(ctx); existingTx != nil {
return r.deleteWithAudit(ctx, existingTx.Client(), id, tombstoneKey)
return r.deleteWithTombstone(ctx, existingTx.Client(), id, tombstoneKey)
}
tx, err := r.client.Tx(ctx)
@@ -348,7 +346,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
exec = tx.Client()
}
if err := r.deleteWithAudit(ctx, exec, id, tombstoneKey); err != nil {
if err := r.deleteWithTombstone(ctx, exec, id, tombstoneKey); err != nil {
return err
}
@@ -358,17 +356,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
return nil
}
func (r *apiKeyRepository) deleteWithAudit(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
// 1. 审计:数据源即 api_keys 当前行;WHERE deleted_at IS NULL 保证只对未删除行写一次。
if _, err := exec.ExecContext(ctx, `
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
SELECT key, id, user_id, name, NOW()
FROM api_keys
WHERE id = $1 AND deleted_at IS NULL`, id); err != nil {
return err
}
// 2. 软删除(tombstone 覆盖 key)。
func (r *apiKeyRepository) deleteWithTombstone(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
res, err := exec.ExecContext(ctx, `
UPDATE api_keys
SET key = $1, deleted_at = NOW(), updated_at = NOW()
@@ -556,7 +556,7 @@ func TestIncrementQuotaUsed_Concurrent(t *testing.T) {
"并发递增后总和应为 %v,实际为 %v", float64(goroutines)*increment, got.QuotaUsed)
}
func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
func (s *APIKeyRepoSuite) TestDeleteWithAudit_TombstonesWithoutRetainingCredential() {
user := s.mustCreateUser("delwithaudit@test.com")
key := &service.APIKey{
UserID: user.ID,
@@ -571,18 +571,24 @@ func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
_, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().Error(err)
rows, qErr := s.client.QueryContext(s.ctx,
`SELECT key, key_name, user_id, api_key_id FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
s.Require().NoError(qErr)
defer rows.Close()
s.Require().True(rows.Next(), "expected one audit row")
var auditKey, auditName string
var auditUserID, auditAPIKeyID int64
s.Require().NoError(rows.Scan(&auditKey, &auditName, &auditUserID, &auditAPIKeyID))
s.Require().Equal("sk-del-audit-1", auditKey)
s.Require().Equal("Audit Me", auditName)
s.Require().Equal(user.ID, auditUserID)
s.Require().Equal(key.ID, auditAPIKeyID)
var tombstone string
var deletedAt time.Time
rows, err := s.repo.sql.QueryContext(s.ctx, `SELECT key, deleted_at FROM api_keys WHERE id = $1`, key.ID)
s.Require().NoError(err)
s.Require().True(rows.Next())
s.Require().NoError(rows.Scan(&tombstone, &deletedAt))
s.Require().NoError(rows.Close())
s.Require().NotEqual("sk-del-audit-1", tombstone)
s.Require().Contains(tombstone, "__deleted__")
var auditCount int
auditRows, err := s.repo.sql.QueryContext(s.ctx,
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
s.Require().NoError(err)
s.Require().True(auditRows.Next())
s.Require().NoError(auditRows.Scan(&auditCount))
s.Require().NoError(auditRows.Close())
s.Require().Zero(auditCount, "deleted credentials must not be retained")
}
func (s *APIKeyRepoSuite) TestDeleteWithAudit_RepeatIsIdempotent() {
@@ -0,0 +1,120 @@
//go:build integration
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func TestAuthCacheInvalidationTriggers_CoverSecurityMutationsOnly(t *testing.T) {
ctx := context.Background()
suffix := time.Now().UnixNano()
group := mustCreateGroup(t, integrationEntClient, &service.Group{
Name: fmt.Sprintf("auth-outbox-group-%d", suffix), RateMultiplier: 1, IsExclusive: true,
})
user := mustCreateUser(t, integrationEntClient, &service.User{
Email: fmt.Sprintf("auth-outbox-%d@example.com", suffix), Concurrency: 5,
})
groupID := group.ID
keyValue := fmt.Sprintf("sk-auth-outbox-%d", suffix)
apiKeyRepo := NewAPIKeyRepository(integrationEntClient, integrationDB)
key := &service.APIKey{UserID: user.ID, GroupID: &groupID, Key: keyValue, Name: "outbox", Status: service.StatusActive}
require.NoError(t, apiKeyRepo.Create(ctx, key))
sum := sha256.Sum256([]byte(keyValue))
cacheKey := hex.EncodeToString(sum[:])
clear := func() {
_, err := integrationDB.ExecContext(ctx, "DELETE FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey)
require.NoError(t, err)
}
count := func() int {
var value int
require.NoError(t, integrationDB.QueryRowContext(ctx,
"SELECT COUNT(*) FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey).Scan(&value))
return value
}
clear()
t.Cleanup(clear)
t.Cleanup(func() {
// Keep the shared integration database isolated for suites that assert
// platform-wide group counts. The final clear cleanup runs after this one
// and removes invalidations emitted by these hard deletes.
_, err := integrationDB.ExecContext(ctx, "DELETE FROM user_allowed_groups WHERE user_id = $1 OR group_id = $2", user.ID, group.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM api_keys WHERE id = $1", key.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", user.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM groups WHERE id = $1", group.ID)
require.NoError(t, err)
})
_, err := integrationDB.ExecContext(ctx, `
UPDATE api_keys
SET quota_used = quota_used + 1,
usage_5h = usage_5h + 1,
last_used_at = NOW()
WHERE id = $1`, key.ID)
require.NoError(t, err)
require.Zero(t, count(), "usage-only key updates must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'disabled' WHERE id = $1", key.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "key disable must enqueue")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'active' WHERE id = $1", key.ID)
require.NoError(t, err)
clear()
userRepo := NewUserRepository(integrationEntClient, integrationDB)
loadedUser, err := userRepo.GetByID(ctx, user.ID)
require.NoError(t, err)
loadedUser.Balance += 10
require.NoError(t, userRepo.Update(ctx, loadedUser))
require.Zero(t, count(), "balance update with unchanged allowed groups must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'disabled' WHERE id = $1", user.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "user disable must enqueue all active keys")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'active' WHERE id = $1", user.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET name = name || '-cosmetic' WHERE id = $1", group.ID)
require.NoError(t, err)
require.Zero(t, count(), "cosmetic group update must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'disabled' WHERE id = $1", group.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "group disable must enqueue bound keys")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'active' WHERE id = $1", group.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx,
"INSERT INTO user_allowed_groups (user_id, group_id) VALUES ($1, $2)", user.ID, group.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx,
"DELETE FROM user_allowed_groups WHERE user_id = $1 AND group_id = $2", user.ID, group.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "exclusive-group revocation must enqueue")
clear()
require.NoError(t, apiKeyRepo.DeleteWithAudit(ctx, key.ID))
require.Equal(t, 1, count(), "tombstone delete must hash OLD.key exactly once")
var stored string
require.NoError(t, integrationDB.QueryRowContext(ctx,
"SELECT cache_key FROM auth_cache_invalidation_outbox WHERE cache_key = $1 LIMIT 1", cacheKey).Scan(&stored))
require.Equal(t, cacheKey, stored)
require.NotContains(t, stored, keyValue)
}
@@ -0,0 +1,159 @@
package repository
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
)
type authCacheInvalidationOutboxRepository struct {
db *sql.DB
}
func NewAuthCacheInvalidationOutboxRepository(db *sql.DB) service.AuthCacheInvalidationOutboxRepository {
return &authCacheInvalidationOutboxRepository{db: db}
}
func (r *authCacheInvalidationOutboxRepository) Claim(ctx context.Context, workerID string, limit int, lease time.Duration) ([]service.AuthCacheInvalidationEvent, error) {
if r == nil || r.db == nil {
return nil, errors.New("nil auth cache invalidation outbox database")
}
if limit <= 0 {
limit = 100
}
leaseSeconds := int64(lease / time.Second)
if leaseSeconds < 1 {
leaseSeconds = 30
}
rows, err := r.db.QueryContext(ctx, `
WITH candidates AS (
SELECT id
FROM auth_cache_invalidation_outbox
WHERE available_at <= NOW()
AND (claimed_at IS NULL OR claimed_at < NOW() - ($3 * INTERVAL '1 second'))
ORDER BY id ASC
LIMIT $2
FOR UPDATE SKIP LOCKED
)
UPDATE auth_cache_invalidation_outbox AS o
SET claimed_at = NOW(), claimed_by = $1
FROM candidates AS c
WHERE o.id = c.id
RETURNING o.id, o.cache_key, o.attempts, o.delivery_stage, o.created_at
`, workerID, limit, leaseSeconds)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
events := make([]service.AuthCacheInvalidationEvent, 0, limit)
for rows.Next() {
var event service.AuthCacheInvalidationEvent
if err := rows.Scan(&event.ID, &event.CacheKey, &event.Attempts, &event.Stage, &event.CreatedAt); err != nil {
return nil, err
}
event.CacheKey = strings.TrimSpace(event.CacheKey)
events = append(events, event)
}
if err := rows.Err(); err != nil {
return nil, err
}
return events, nil
}
func (r *authCacheInvalidationOutboxRepository) ScheduleSecondPass(ctx context.Context, id int64, workerID string, availableAt time.Time) error {
result, err := r.db.ExecContext(ctx, `
UPDATE auth_cache_invalidation_outbox
SET delivery_stage = 1,
available_at = $3,
last_error = NULL,
claimed_at = NULL,
claimed_by = NULL
WHERE id = $1 AND claimed_by = $2 AND delivery_stage = 0
`, id, workerID, availableAt)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d cannot schedule second pass", id)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) DeleteClaimed(ctx context.Context, id int64, workerID string) error {
result, err := r.db.ExecContext(ctx, `
DELETE FROM auth_cache_invalidation_outbox
WHERE id = $1 AND claimed_by = $2
`, id, workerID)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) RetryClaimed(ctx context.Context, id int64, workerID string, availableAt time.Time, lastError string) error {
result, err := r.db.ExecContext(ctx, `
UPDATE auth_cache_invalidation_outbox
SET attempts = attempts + 1,
available_at = $3,
last_error = $4,
claimed_at = NULL,
claimed_by = NULL
WHERE id = $1 AND claimed_by = $2
`, id, workerID, availableAt, lastError)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) Stats(ctx context.Context) (service.AuthCacheInvalidationOutboxStats, error) {
var (
stats service.AuthCacheInvalidationOutboxStats
oldest sql.NullTime
lastError sql.NullString
)
err := r.db.QueryRowContext(ctx, `
SELECT COUNT(*), MIN(created_at), COALESCE(MAX(attempts), 0),
(SELECT last_error
FROM auth_cache_invalidation_outbox
WHERE last_error IS NOT NULL
ORDER BY available_at DESC, id DESC
LIMIT 1)
FROM auth_cache_invalidation_outbox
`).Scan(&stats.Pending, &oldest, &stats.MaxAttempts, &lastError)
if err != nil {
return stats, err
}
if oldest.Valid {
value := oldest.Time
stats.OldestCreatedAt = &value
}
if lastError.Valid {
stats.LastError = lastError.String
}
return stats, nil
}
@@ -0,0 +1,123 @@
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"strings"
"testing"
"time"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/Wei-Shaw/sub2api/migrations"
"github.com/stretchr/testify/require"
)
func TestAuthCacheInvalidationOutboxRepository_ClaimUsesLeaseAndSkipLocked(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
created := time.Now().UTC()
mock.ExpectQuery("(?s)claimed_at < NOW\\(\\) - .*FOR UPDATE SKIP LOCKED.*RETURNING").
WithArgs("worker-a", 100, int64(30)).
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}).
AddRow(int64(4), strings.Repeat("a", 64), 2, 1, created))
repo := NewAuthCacheInvalidationOutboxRepository(db)
events, err := repo.Claim(context.Background(), "worker-a", 100, 30*time.Second)
require.NoError(t, err)
require.Len(t, events, 1)
require.Equal(t, int64(4), events[0].ID)
require.Equal(t, 1, events[0].Stage)
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_ClaimIsBoundedByDefault(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
mock.ExpectQuery("(?s)FROM auth_cache_invalidation_outbox.*LIMIT \\$2.*SKIP LOCKED").
WithArgs("worker", 100, int64(30)).
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}))
repo := NewAuthCacheInvalidationOutboxRepository(db)
_, err = repo.Claim(context.Background(), "worker", 0, 0)
require.NoError(t, err)
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_ClaimOwnershipTransitions(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
repo := NewAuthCacheInvalidationOutboxRepository(db)
next := time.Now().UTC().Add(time.Minute)
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
WithArgs(int64(1), "worker", next).
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.ScheduleSecondPass(context.Background(), 1, "worker", next))
retryAt := next.Add(time.Minute)
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
WithArgs(int64(2), "worker", retryAt, "publish failed").
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.RetryClaimed(context.Background(), 2, "worker", retryAt, "publish failed"))
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
WithArgs(int64(3), "worker").
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.DeleteClaimed(context.Background(), 3, "worker"))
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_RejectsLostClaim(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
WithArgs(int64(3), "old-worker").
WillReturnResult(sqlmock.NewResult(0, 0))
repo := NewAuthCacheInvalidationOutboxRepository(db)
err = repo.DeleteClaimed(context.Background(), 3, "old-worker")
require.ErrorContains(t, err, "no longer owned")
}
func TestAuthCacheInvalidationOutboxRepository_StatsExposeDurableLagAndFailures(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
oldest := time.Now().UTC().Add(-time.Minute)
mock.ExpectQuery("(?s)SELECT COUNT\\(\\*\\), MIN\\(created_at\\), COALESCE\\(MAX\\(attempts\\), 0\\)").
WillReturnRows(sqlmock.NewRows([]string{"count", "min", "max", "last_error"}).AddRow(5, oldest, 7, "redis down"))
repo := NewAuthCacheInvalidationOutboxRepository(db)
stats, err := repo.Stats(context.Background())
require.NoError(t, err)
require.Equal(t, int64(5), stats.Pending)
require.Equal(t, 7, stats.MaxAttempts)
require.Equal(t, "redis down", stats.LastError)
require.NotNil(t, stats.OldestCreatedAt)
}
func TestAuthCacheInvalidationMigration_SecurityCoverageAndNoPlaintextPayload(t *testing.T) {
content, err := migrations.FS.ReadFile("184_auth_cache_invalidation_outbox.sql")
require.NoError(t, err)
sqlText := string(content)
for _, required := range []string{
"encode(sha256(convert_to(raw_key, 'UTF8')), 'hex')",
"OLD.key", "OLD.status", "OLD.deleted_at", "OLD.user_id", "OLD.group_id",
"OLD.ip_whitelist", "OLD.ip_blacklist", "OLD.expires_at",
"trg_users_auth_cache_invalidation", "trg_groups_auth_cache_invalidation",
"trg_user_allowed_groups_auth_cache_invalidation", "FOR EACH ROW",
"delivery_stage", "claimed_at", "available_at",
} {
require.Contains(t, sqlText, required)
}
require.NotContains(t, sqlText, "quota_used IS DISTINCT")
require.NotContains(t, sqlText, "last_used_at IS DISTINCT")
plaintext := "sk-plaintext-must-not-be-stored"
sum := sha256.Sum256([]byte(plaintext))
require.Len(t, hex.EncodeToString(sum[:]), 64)
require.NotContains(t, sqlText, plaintext)
}
+18 -1
View File
@@ -7,6 +7,7 @@ import (
"compress/gzip"
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"io"
@@ -323,7 +324,7 @@ func (t *grokAccessDeniedFallbackTransport) RoundTrip(req *http.Request) (*http.
}
body, ok := bufferSmallResponseBody(resp, grokFallbackBodyLimit)
if !ok || !bytes.Contains(bytes.ToLower(body), []byte("access denied")) {
if !ok || !isGrokCLICompatibilityAccessDenied(body) {
return resp, nil
}
@@ -350,6 +351,22 @@ func (t *grokAccessDeniedFallbackTransport) RoundTrip(req *http.Request) (*http.
return fallbackResp, nil
}
func isGrokCLICompatibilityAccessDenied(body []byte) bool {
lower := bytes.ToLower(body)
if bytes.Contains(lower, []byte("access denied")) {
return true
}
var payload struct {
Code string `json:"code"`
Error string `json:"error"`
}
if err := json.Unmarshal(body, &payload); err != nil || !strings.EqualFold(strings.TrimSpace(payload.Code), "permission_denied") {
return false
}
const chatEndpointDeniedPrefix = "access to the chat endpoint is denied. please ensure you're using the correct credentials. if you believe this is a mistake, please"
return strings.HasPrefix(strings.ToLower(strings.TrimSpace(payload.Error)), chatEndpointDeniedPrefix)
}
func isGrokCLIAccessDeniedFallbackCandidate(req *http.Request, resp *http.Response) bool {
return req != nil && req.URL != nil && req.GetBody != nil && resp != nil &&
resp.StatusCode == http.StatusForbidden &&
@@ -311,6 +311,119 @@ func TestHTTPUpstreamDoFallsBackToOfficialGrokAPIOnCLIAccessDenied(t *testing.T)
require.Empty(t, fallbackHeaders.Get("User-Agent"))
}
func TestGrokAccessDeniedFallbackRecognizesChatEndpointPermissionDenied(t *testing.T) {
var hosts []string
transport := &grokAccessDeniedFallbackTransport{
base: roundTripFunc(func(req *http.Request) (*http.Response, error) {
hosts = append(hosts, req.URL.Hostname())
if req.URL.Hostname() == grokCLIProxyHost {
return &http.Response{
StatusCode: http.StatusForbidden,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(
`{"code":"permission_denied","error":"Access to the chat endpoint is denied. Please ensure you're using the correct credentials. If you believe this is a mistake, please contact support."}`,
)),
Request: req,
}, nil
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(`{"id":"response-ok"}`)),
Request: req,
}, nil
}),
}
req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", strings.NewReader(`{"model":"grok-4.5"}`))
require.NoError(t, err)
req.Header.Set("Authorization", "Bearer oauth-token")
req.Header.Set("X-XAI-Token-Auth", "xai-grok-cli")
resp, err := transport.RoundTrip(req)
require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
require.NoError(t, resp.Body.Close())
require.Equal(t, []string{grokCLIProxyHost, grokOfficialAPIHost}, hosts)
}
func TestIsGrokCLICompatibilityAccessDenied(t *testing.T) {
tests := []struct {
name string
body string
want bool
}{
{name: "legacy compatibility wording", body: `{"error":"Access denied"}`, want: true},
{
name: "observed chat endpoint permission denial",
body: `{"code":"permission_denied","error":"Access to the chat endpoint is denied. Please ensure you're using the correct credentials. If you believe this is a mistake, please contact support."}`,
want: true,
},
{
name: "entitlement denial using the same broad terms",
body: `{"code":"permission_denied","error":"Access to the chat endpoint is denied because a subscription is required"}`,
want: false,
},
{
name: "different permission denied endpoint",
body: `{"code":"permission_denied","error":"Access to the billing endpoint is denied."}`,
want: false,
},
{
name: "wrong structured error code",
body: `{"code":"subscription_required","error":"Access to the chat endpoint is denied. Please ensure you're using the correct credentials. If you believe this is a mistake, please contact support."}`,
want: false,
},
{name: "malformed response", body: `permission_denied: chat endpoint denied`, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, isGrokCLICompatibilityAccessDenied([]byte(tt.body)))
})
}
}
func TestIsGrokCLIAccessDeniedFallbackCandidateRequiresAuthenticatedReplayableCLI403(t *testing.T) {
newRequest := func() *http.Request {
req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", strings.NewReader(`{"model":"grok-4.5"}`))
require.NoError(t, err)
req.Header.Set("Authorization", "Bearer oauth-token")
req.Header.Set("X-XAI-Token-Auth", "xai-grok-cli")
return req
}
newResponse := func() *http.Response { return &http.Response{StatusCode: http.StatusForbidden} }
t.Run("valid candidate", func(t *testing.T) {
require.True(t, isGrokCLIAccessDeniedFallbackCandidate(newRequest(), newResponse()))
})
t.Run("non CLI host", func(t *testing.T) {
req := newRequest()
req.URL.Host = "api.x.ai"
require.False(t, isGrokCLIAccessDeniedFallbackCandidate(req, newResponse()))
})
t.Run("missing CLI identity", func(t *testing.T) {
req := newRequest()
req.Header.Del("X-XAI-Token-Auth")
require.False(t, isGrokCLIAccessDeniedFallbackCandidate(req, newResponse()))
})
t.Run("missing bearer authentication", func(t *testing.T) {
req := newRequest()
req.Header.Del("Authorization")
require.False(t, isGrokCLIAccessDeniedFallbackCandidate(req, newResponse()))
})
t.Run("non forbidden response", func(t *testing.T) {
resp := newResponse()
resp.StatusCode = http.StatusUnauthorized
require.False(t, isGrokCLIAccessDeniedFallbackCandidate(newRequest(), resp))
})
t.Run("non replayable request", func(t *testing.T) {
req := newRequest()
req.GetBody = nil
require.False(t, isGrokCLIAccessDeniedFallbackCandidate(req, newResponse()))
})
}
func TestHTTPUpstreamDoDoesNotFallbackForGrokEntitlementDenial(t *testing.T) {
transport := &grokAccessDeniedFallbackTransport{
base: roundTripFunc(func(req *http.Request) (*http.Response, error) {
@@ -128,23 +128,31 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
// 获取分布式锁,确保多实例部署时只有一个实例执行迁移。
// 这是 PostgreSQL 特有的 Advisory Lock 机制。
if err := pgAdvisoryLock(ctx, db); err != nil {
lockConn, err := db.Conn(ctx)
if err != nil {
return fmt.Errorf("acquire migrations lock connection: %w", err)
}
defer func() { _ = lockConn.Close() }()
if err := pgAdvisoryLock(ctx, lockConn); err != nil {
return err
}
defer func() {
// 无论迁移是否成功,都要释放锁。
// 使用 context.Background() 确保即使原 ctx 已取消也能释放锁。
_ = pgAdvisoryUnlock(context.Background(), db)
// 独立超时确保原 ctx 取消后仍会尝试释放,但数据库链路异常不会
// 无限阻塞进程退出。
unlockCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = pgAdvisoryUnlock(unlockCtx, lockConn)
}()
// 创建迁移记录表(如果不存在)。
// 该表记录所有已应用的迁移及其校验和。
if _, err := db.ExecContext(ctx, schemaMigrationsTableDDL); err != nil {
if _, err := lockConn.ExecContext(ctx, schemaMigrationsTableDDL); err != nil {
return fmt.Errorf("create schema_migrations: %w", err)
}
// 自动对齐 Atlas 基线(如果检测到 legacy schema_migrations 且缺失 atlas_schema_revisions)。
if err := ensureAtlasBaselineAligned(ctx, db, fsys); err != nil {
if err := ensureAtlasBaselineAligned(ctx, lockConn, fsys); err != nil {
return err
}
@@ -175,7 +183,7 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
// 检查该迁移是否已经应用
var existing string
rowErr := db.QueryRowContext(ctx, "SELECT checksum FROM schema_migrations WHERE filename = $1", name).Scan(&existing)
rowErr := lockConn.QueryRowContext(ctx, "SELECT checksum FROM schema_migrations WHERE filename = $1", name).Scan(&existing)
if rowErr == nil {
// 迁移已应用,验证校验和是否匹配
if existing != checksum {
@@ -207,7 +215,7 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
}
if nonTx {
if err := prepareNonTransactionalMigration(ctx, db, name); err != nil {
if err := prepareNonTransactionalMigration(ctx, lockConn, name); err != nil {
return fmt.Errorf("prepare migration %s: %w", name, err)
}
@@ -222,18 +230,18 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
if stripSQLLineComment(trimmed) == "" {
continue
}
if _, err := db.ExecContext(ctx, trimmed); err != nil {
if _, err := lockConn.ExecContext(ctx, trimmed); err != nil {
return fmt.Errorf("apply migration %s (non-tx statement %d): %w", name, i+1, err)
}
}
if _, err := db.ExecContext(ctx, "INSERT INTO schema_migrations (filename, checksum) VALUES ($1, $2)", name, checksum); err != nil {
if _, err := lockConn.ExecContext(ctx, "INSERT INTO schema_migrations (filename, checksum) VALUES ($1, $2)", name, checksum); err != nil {
return fmt.Errorf("record migration %s (non-tx): %w", name, err)
}
continue
}
// 默认迁移在事务中执行,确保原子性:要么完全成功,要么完全回滚。
tx, err := db.BeginTx(ctx, nil)
tx, err := lockConn.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin migration %s: %w", name, err)
}
@@ -260,7 +268,14 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
return nil
}
func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name string) error {
type migrationConnection interface {
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error)
}
func prepareNonTransactionalMigration(ctx context.Context, db migrationConnection, name string) error {
switch name {
case paymentOrdersOutTradeNoUniqueMigration:
return preparePaymentOrdersOutTradeNoUniqueMigration(ctx, db)
@@ -273,7 +288,7 @@ func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name stri
}
}
func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db *sql.DB) error {
func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db migrationConnection) error {
duplicates, err := findDuplicatePaymentOrderOutTradeNos(ctx, db)
if err != nil {
return fmt.Errorf("precheck duplicate out_trade_no: %w", err)
@@ -289,7 +304,7 @@ func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db *sql.
return dropInvalidIndexIfPresent(ctx, db, paymentOrdersOutTradeNoUniqueIndex)
}
func dropInvalidIndexIfPresent(ctx context.Context, db *sql.DB, indexName string) error {
func dropInvalidIndexIfPresent(ctx context.Context, db migrationConnection, indexName string) error {
invalid, err := indexIsInvalid(ctx, db, indexName)
if err != nil {
return fmt.Errorf("check invalid index %s: %w", indexName, err)
@@ -304,7 +319,7 @@ func dropInvalidIndexIfPresent(ctx context.Context, db *sql.DB, indexName string
return nil
}
func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db *sql.DB) ([]string, error) {
func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db migrationConnection) ([]string, error) {
rows, err := db.QueryContext(ctx, `
SELECT out_trade_no, COUNT(*) AS duplicate_count
FROM payment_orders
@@ -336,7 +351,7 @@ func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db *sql.DB) ([]st
return duplicates, nil
}
func indexIsInvalid(ctx context.Context, db *sql.DB, indexName string) (bool, error) {
func indexIsInvalid(ctx context.Context, db migrationConnection, indexName string) (bool, error) {
var invalid bool
err := db.QueryRowContext(ctx, `
SELECT EXISTS (
@@ -352,7 +367,7 @@ func indexIsInvalid(ctx context.Context, db *sql.DB, indexName string) (bool, er
return invalid, err
}
func ensureAtlasBaselineAligned(ctx context.Context, db *sql.DB, fsys fs.FS) error {
func ensureAtlasBaselineAligned(ctx context.Context, db migrationConnection, fsys fs.FS) error {
hasLegacy, err := tableExists(ctx, db, "schema_migrations")
if err != nil {
return fmt.Errorf("check schema_migrations: %w", err)
@@ -393,7 +408,7 @@ func ensureAtlasBaselineAligned(ctx context.Context, db *sql.DB, fsys fs.FS) err
return nil
}
func tableExists(ctx context.Context, db *sql.DB, tableName string) (bool, error) {
func tableExists(ctx context.Context, db migrationConnection, tableName string) (bool, error) {
var exists bool
err := db.QueryRowContext(ctx, `
SELECT EXISTS (
@@ -524,7 +539,12 @@ func stripSQLLineComment(s string) string {
// pgAdvisoryLock 获取 PostgreSQL Advisory Lock。
// Advisory Lock 是一种轻量级的锁机制,不与任何特定的数据库对象关联。
// 它非常适合用于应用层面的分布式锁场景,如迁移序列化。
func pgAdvisoryLock(ctx context.Context, db *sql.DB) error {
type advisoryLockConnection interface {
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
}
func pgAdvisoryLock(ctx context.Context, db advisoryLockConnection) error {
ticker := time.NewTicker(migrationsLockRetryInterval)
defer ticker.Stop()
@@ -546,7 +566,7 @@ func pgAdvisoryLock(ctx context.Context, db *sql.DB) error {
// pgAdvisoryUnlock 释放 PostgreSQL Advisory Lock。
// 必须在获取锁后确保释放,否则会阻塞其他实例的迁移操作。
func pgAdvisoryUnlock(ctx context.Context, db *sql.DB) error {
func pgAdvisoryUnlock(ctx context.Context, db advisoryLockConnection) error {
_, err := db.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", migrationsAdvisoryLockID)
if err != nil {
return fmt.Errorf("release migrations lock: %w", err)
@@ -275,6 +275,9 @@ func TestApplyMigrationsFS_TransactionalMigration(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
// The advisory lock and all migration work must share one session. This also
// proves startup cannot self-deadlock when deployments cap the pool at one.
db.SetMaxOpenConns(1)
prepareMigrationsBootstrapExpectations(mock)
mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1").
@@ -5,11 +5,32 @@ package repository
import (
"context"
"database/sql"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestMigrationsRunner_ConcurrentInstancesSerializeOnSessionLock(t *testing.T) {
const instances = 2
errorsByInstance := make([]error, instances)
var wg sync.WaitGroup
for i := 0; i < instances; i++ {
wg.Add(1)
go func(index int) {
defer wg.Done()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
errorsByInstance[index] = ApplyMigrations(ctx, integrationDB)
}(i)
}
wg.Wait()
for i, err := range errorsByInstance {
require.NoErrorf(t, err, "migration instance %d", i)
}
}
func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) {
tx := testTx(t)
@@ -111,6 +132,13 @@ func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) {
requireColumn(t, tx, "ops_system_logs", "api_key_id", "bigint", 0, true)
requireIndex(t, tx, "ops_system_logs", "idx_ops_system_logs_api_key_id_created_at")
// Bounded ingress rejection security aggregates.
requireColumn(t, tx, "ops_ingress_reject_aggregates", "bucket_start", "timestamp with time zone", 0, false)
requireColumn(t, tx, "ops_ingress_reject_aggregates", "client_ip", "inet", 0, false)
requireColumn(t, tx, "ops_ingress_reject_aggregates", "request_count", "bigint", 0, false)
requireIndex(t, tx, "ops_ingress_reject_aggregates", "idx_ops_ingress_reject_aggregates_bucket")
requireIndex(t, tx, "ops_ingress_reject_aggregates", "idx_ops_ingress_reject_aggregates_ip_bucket")
// user_allowed_groups table should exist
var uagRegclass sql.NullString
require.NoError(t, tx.QueryRowContext(context.Background(), "SELECT to_regclass('public.user_allowed_groups')").Scan(&uagRegclass))
@@ -137,26 +137,17 @@ func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) {
}
}
func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
func TestBuildOpsErrorLogsWhere_UserOwnershipIsDirectOnly(t *testing.T) {
uid := int64(42)
// 开关开启 → 归属放宽为 OR(user_id 或 deleted_key_owner_user_id),且共用同一占位符
on := &service.OpsErrorLogFilter{UserID: &uid, MatchDeletedKeyOwner: true}
whereOn, argsOn := buildOpsErrorLogsWhere(on)
if !strings.Contains(whereOn, "(e.user_id = $1 OR e.deleted_key_owner_user_id = $1)") {
t.Fatalf("MatchDeletedKeyOwner=true should widen to OR, got: %s", whereOn)
filter := &service.OpsErrorLogFilter{UserID: &uid}
where, args := buildOpsErrorLogsWhere(filter)
if !strings.Contains(where, "e.user_id = $1") {
t.Fatalf("user scope should match user_id exactly, got: %s", where)
}
if len(argsOn) != 1 || argsOn[0] != uid {
t.Fatalf("expected single reused arg %d, got %v", uid, argsOn)
if len(args) != 1 || args[0] != uid {
t.Fatalf("expected user id arg %d, got %v", uid, args)
}
// 开关关闭(默认)→ 仅精确 user_id,绝不出现 deleted_key_owner_user_id(admin 回归)
off := &service.OpsErrorLogFilter{UserID: &uid}
whereOff, _ := buildOpsErrorLogsWhere(off)
if !strings.Contains(whereOff, "e.user_id = $1") {
t.Fatalf("default should match user_id exactly, got: %s", whereOff)
}
if strings.Contains(whereOff, "deleted_key_owner_user_id") {
t.Fatalf("default must NOT include deleted_key_owner_user_id, got: %s", whereOff)
if strings.Contains(where, "deleted_key_owner_user_id") {
t.Fatalf("user ownership must not depend on deleted-key attribution: %s", where)
}
}
@@ -0,0 +1,158 @@
package repository
import (
"context"
"fmt"
"strings"
"github.com/Wei-Shaw/sub2api/internal/service"
)
const ingressRejectUpsertChunkSize = 500
func (r *opsRepository) BatchUpsertIngressRejects(ctx context.Context, items []*service.OpsIngressRejectAggregate) error {
if r == nil || r.db == nil || len(items) == 0 {
return nil
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
for start := 0; start < len(items); start += ingressRejectUpsertChunkSize {
end := start + ingressRejectUpsertChunkSize
if end > len(items) {
end = len(items)
}
valid := make([]*service.OpsIngressRejectAggregate, 0, end-start)
for _, item := range items[start:end] {
if item != nil && item.RequestCount > 0 {
valid = append(valid, item)
}
}
if len(valid) == 0 {
continue
}
var query strings.Builder
_, _ = query.WriteString(`INSERT INTO ops_ingress_reject_aggregates
(bucket_start, reject_reason, route_family, protocol, client_ip, user_id, api_key_id, request_count, first_seen, last_seen)
VALUES `)
args := make([]any, 0, len(valid)*10)
for i, item := range valid {
if i > 0 {
_ = query.WriteByte(',')
}
base := len(args)
fmt.Fprintf(&query, "($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d)",
base+1, base+2, base+3, base+4, base+5, base+6, base+7, base+8, base+9, base+10)
var userID, apiKeyID int64
if item.UserID != nil {
userID = *item.UserID
}
if item.APIKeyID != nil {
apiKeyID = *item.APIKeyID
}
args = append(args, item.BucketStart.UTC(), item.RejectReason, item.RouteFamily, item.Protocol,
item.ClientIP, userID, apiKeyID, item.RequestCount, item.FirstSeen.UTC(), item.LastSeen.UTC())
}
_, _ = query.WriteString(`
ON CONFLICT (bucket_start, reject_reason, route_family, protocol, client_ip, user_id, api_key_id)
DO UPDATE SET request_count = ops_ingress_reject_aggregates.request_count + EXCLUDED.request_count,
first_seen = LEAST(ops_ingress_reject_aggregates.first_seen, EXCLUDED.first_seen),
last_seen = GREATEST(ops_ingress_reject_aggregates.last_seen, EXCLUDED.last_seen),
updated_at = NOW()`)
if _, err := tx.ExecContext(ctx, query.String(), args...); err != nil {
return err
}
}
return tx.Commit()
}
func (r *opsRepository) ListIngressRejects(ctx context.Context, filter *service.OpsIngressRejectFilter) (*service.OpsIngressRejectList, error) {
if r == nil || r.db == nil {
return nil, fmt.Errorf("nil ops repository")
}
if filter == nil {
filter = &service.OpsIngressRejectFilter{}
}
page, pageSize := filter.Page, filter.PageSize
if page <= 0 {
page = 1
}
if pageSize <= 0 {
pageSize = 50
}
if pageSize > 200 {
pageSize = 200
}
clauses := []string{"1=1"}
args := make([]any, 0)
add := func(expr string, value any) {
args = append(args, value)
clauses = append(clauses, fmt.Sprintf(expr, len(args)))
}
if filter.StartTime != nil {
add("bucket_start >= $%d", filter.StartTime.UTC())
}
if filter.EndTime != nil {
add("bucket_start < $%d", filter.EndTime.UTC())
}
if value := strings.TrimSpace(filter.RejectReason); value != "" {
add("reject_reason = $%d", value)
}
if value := strings.TrimSpace(filter.RouteFamily); value != "" {
add("route_family = $%d", value)
}
if value := strings.TrimSpace(filter.Protocol); value != "" {
add("protocol = $%d", value)
}
if value := strings.TrimSpace(filter.ClientIP); value != "" {
add("client_ip = $%d::inet", value)
}
if filter.UserID != nil {
add("user_id = $%d", *filter.UserID)
}
if filter.APIKeyID != nil {
add("api_key_id = $%d", *filter.APIKeyID)
}
where := "WHERE " + strings.Join(clauses, " AND ")
var total int
if err := r.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM ops_ingress_reject_aggregates "+where, args...).Scan(&total); err != nil {
return nil, err
}
args = append(args, pageSize, (page-1)*pageSize)
query := fmt.Sprintf(`SELECT id,bucket_start,reject_reason,route_family,protocol,host(client_ip),user_id,api_key_id,request_count,first_seen,last_seen
FROM ops_ingress_reject_aggregates %s ORDER BY bucket_start DESC,id DESC LIMIT $%d OFFSET $%d`, where, len(args)-1, len(args))
rows, err := r.db.QueryContext(ctx, query, args...)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
result := &service.OpsIngressRejectList{
Items: make([]*service.OpsIngressRejectAggregate, 0, pageSize), Total: total, Page: page, PageSize: pageSize,
}
for rows.Next() {
item := &service.OpsIngressRejectAggregate{}
var userID, apiKeyID int64
if err := rows.Scan(&item.ID, &item.BucketStart, &item.RejectReason, &item.RouteFamily, &item.Protocol,
&item.ClientIP, &userID, &apiKeyID, &item.RequestCount, &item.FirstSeen, &item.LastSeen); err != nil {
return nil, err
}
if userID > 0 {
item.UserID = &userID
}
if apiKeyID > 0 {
item.APIKeyID = &apiKeyID
}
result.Items = append(result.Items, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
}
@@ -0,0 +1,32 @@
package repository
import (
"context"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func TestBatchUpsertIngressRejectsUsesFixedMultiRowChunks(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
t.Cleanup(func() { _ = db.Close() })
repo := &opsRepository{db: db}
now := time.Now().UTC().Truncate(time.Minute)
items := make([]*service.OpsIngressRejectAggregate, ingressRejectUpsertChunkSize+1)
for i := range items {
items[i] = &service.OpsIngressRejectAggregate{
BucketStart: now, RejectReason: "invalid_api_key", RouteFamily: "messages",
Protocol: "anthropic", ClientIP: "192.0.2.1", RequestCount: 1, FirstSeen: now, LastSeen: now,
}
}
mock.ExpectBegin()
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, int64(ingressRejectUpsertChunkSize)))
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
require.NoError(t, repo.BatchUpsertIngressRejects(context.Background(), items))
require.NoError(t, mock.ExpectationsWereMet())
}
+7 -82
View File
@@ -4,7 +4,6 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
@@ -56,12 +55,9 @@ INSERT INTO ops_error_logs (
response_latency_ms,
time_to_first_token_ms,
created_at,
attempted_key_prefix,
deleted_key_owner_user_id,
deleted_key_name,
api_key_prefix
) VALUES (
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38,$39,$40,$41
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38
)`
func NewOpsRepository(db *sql.DB) service.OpsRepository {
@@ -170,9 +166,6 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
opsNullInt64(input.ResponseLatencyMs),
opsNullInt64(input.TimeToFirstTokenMs),
input.CreatedAt,
opsNullString(input.AttemptedKeyPrefix),
opsNullInt64(input.DeletedKeyOwnerUserID),
opsNullString(input.DeletedKeyName),
opsNullString(input.APIKeyPrefix),
}
}
@@ -274,16 +267,12 @@ SELECT
COALESCE(e.user_agent, ''),
e.request_type,
COALESCE(ak.name, ''),
ak.deleted_at,
COALESCE(e.deleted_key_name, ''),
e.deleted_key_owner_user_id,
COALESCE(du.email, '')
ak.deleted_at
FROM ops_error_logs e
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
` + where + `
ORDER BY ` + opsErrorLogsOrderBy(filter) + `
@@ -313,9 +302,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
var requestType sql.NullInt64
var apiKeyName string
var apiKeyDeletedAt sql.NullTime
var deletedKeyName string
var deletedKeyOwnerID sql.NullInt64
var deletedKeyOwnerEmail string
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
@@ -352,9 +338,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
&requestType,
&apiKeyName,
&apiKeyDeletedAt,
&deletedKeyName,
&deletedKeyOwnerID,
&deletedKeyOwnerEmail,
); err != nil {
return nil, err
}
@@ -395,21 +378,8 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
v := int16(requestType.Int64)
item.RequestType = &v
}
// Key 名称:优先关联到的 ak.name(已软删的 key name 仍保留);
// 关联不到(api_key_id 为空 / 历史硬删)时回退错误记录里快照的 deleted_key_name。
if apiKeyName != "" {
item.APIKeyName = apiKeyName
} else {
item.APIKeyName = deletedKeyName
}
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
// 已删除 KEY 所有者快照:认证失败行 user_id 为空,列表用户列以此回退。
if deletedKeyOwnerID.Valid {
v := deletedKeyOwnerID.Int64
item.DeletedKeyOwnerUserID = &v
item.DeletedKeyOwnerEmail = deletedKeyOwnerEmail
}
item.APIKeyName = apiKeyName
item.APIKeyDeleted = apiKeyDeletedAt.Valid
out = append(out, &item)
}
if err := rows.Err(); err != nil {
@@ -477,10 +447,6 @@ SELECT
e.upstream_latency_ms,
e.response_latency_ms,
e.time_to_first_token_ms,
COALESCE(e.attempted_key_prefix, ''),
e.deleted_key_owner_user_id,
COALESCE(du.email, ''),
COALESCE(e.deleted_key_name, ''),
COALESCE(e.api_key_prefix, ''),
COALESCE(ak.name, ''),
ak.deleted_at
@@ -488,7 +454,6 @@ FROM ops_error_logs e
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
WHERE e.id = $1
LIMIT 1`
@@ -509,7 +474,6 @@ LIMIT 1`
var responseLatency sql.NullInt64
var ttft sql.NullInt64
var requestType sql.NullInt64
var deletedKeyOwnerUserID sql.NullInt64
var detailAPIKeyName string
var detailAPIKeyDeletedAt sql.NullTime
@@ -557,10 +521,6 @@ LIMIT 1`
&upstreamLatency,
&responseLatency,
&ttft,
&out.AttemptedKeyPrefix,
&deletedKeyOwnerUserID,
&out.DeletedKeyOwnerEmail,
&out.DeletedKeyName,
&out.APIKeyPrefix,
&detailAPIKeyName,
&detailAPIKeyDeletedAt,
@@ -626,18 +586,8 @@ LIMIT 1`
v := int16(requestType.Int64)
out.RequestType = &v
}
if deletedKeyOwnerUserID.Valid {
v := deletedKeyOwnerUserID.Int64
out.DeletedKeyOwnerUserID = &v
}
// Key 名称:优先关联到的 ak.name;关联不到时回退快照的 deleted_key_name。
if detailAPIKeyName != "" {
out.APIKeyName = detailAPIKeyName
} else {
out.APIKeyName = out.DeletedKeyName
}
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid || (detailAPIKeyName == "" && out.DeletedKeyName != "")
out.APIKeyName = detailAPIKeyName
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid
// Normalize upstream_errors to empty string when stored as JSON null.
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
@@ -648,26 +598,6 @@ LIMIT 1`
return &out, nil
}
// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计。
// 同一 key 可能有多条历史(反复创建/删除),取 deleted_at 最近一条(id 作同毫秒 tiebreaker)。
// 未命中返回 (nil, nil)。
func (r *opsRepository) LookupDeletedKeyAudit(ctx context.Context, key string) (*service.DeletedKeyAuditResult, error) {
var res service.DeletedKeyAuditResult
err := r.db.QueryRowContext(ctx, `
SELECT user_id, key_name
FROM deleted_api_key_audits
WHERE key = $1
ORDER BY deleted_at DESC, id DESC
LIMIT 1`, key).Scan(&res.UserID, &res.KeyName)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return nil, err
}
return &res, nil
}
func (r *opsRepository) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64, resolvedAt *time.Time) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil ops repository")
@@ -1082,12 +1012,7 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
if filter.UserID != nil && *filter.UserID > 0 {
args = append(args, *filter.UserID)
n := itoa(len(args))
if filter.MatchDeletedKeyOwner {
// 用户侧:把「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录也纳入。
clauses = append(clauses, "(e.user_id = $"+n+" OR e.deleted_key_owner_user_id = $"+n+")")
} else {
clauses = append(clauses, "e.user_id = $"+n)
}
clauses = append(clauses, "e.user_id = $"+n)
}
if filter.APIKeyID != nil && *filter.APIKeyID > 0 {
args = append(args, *filter.APIKeyID)
@@ -14,7 +14,7 @@ func TestOpsInsertErrorLogArgsPreservesExplicitZeroUpstreamStatus(t *testing.T)
zero := 0
args := opsInsertErrorLogArgs(&service.OpsInsertErrorLogInput{UpstreamStatusCode: &zero})
require.Len(t, args, 41)
require.Len(t, args, 38)
encoded, ok := args[27].(sql.NullInt64)
require.True(t, ok)
require.True(t, encoded.Valid)
@@ -11,47 +11,13 @@ import (
"github.com/stretchr/testify/require"
)
// TestGetErrorLogByID_DeletedKeyOwner 验证:
// 1. 带 deleted_key_owner_user_id 的记录能正确 JOIN users 返回 DeletedKeyOwnerEmail
// 2. 新列全为 NULL 的普通记录 Scan 不报错,这些字段为空/nil
func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
func TestGetErrorLogByID_APIKeyPrefixAndUpstreamStatus(t *testing.T) {
ctx := context.Background()
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE ops_error_logs RESTART IDENTITY CASCADE")
repo := NewOpsRepository(integrationDB).(*opsRepository)
// ── Case 1: 带 deleted_key_owner 信息的记录 ──────────────────────────────
owner := mustCreateUser(t, integrationEntClient, &service.User{
Email: "deleted-key-owner-" + time.Now().Format("150405.000000000") + "@example.com",
})
var insertedID int64
err := integrationDB.QueryRowContext(ctx, `
INSERT INTO ops_error_logs (
error_phase, error_type, severity, status_code, created_at,
attempted_key_prefix, deleted_key_owner_user_id, deleted_key_name
) VALUES (
'auth', 'INVALID_API_KEY', 'error', 401, NOW(),
'sk-test-abc', $1, 'my-deleted-key'
) RETURNING id`,
owner.ID,
).Scan(&insertedID)
require.NoError(t, err)
require.Positive(t, insertedID)
detail, err := repo.GetErrorLogByID(ctx, insertedID)
require.NoError(t, err)
require.NotNil(t, detail)
require.Equal(t, "sk-test-abc", detail.AttemptedKeyPrefix)
require.NotNil(t, detail.DeletedKeyOwnerUserID)
require.Equal(t, owner.ID, *detail.DeletedKeyOwnerUserID)
require.Equal(t, owner.Email, detail.DeletedKeyOwnerEmail)
require.Equal(t, "my-deleted-key", detail.DeletedKeyName)
// ── Case 2: 新列全为 NULL 的普通错误记录 ──────────────────────────────────
var plainID int64
err = integrationDB.QueryRowContext(ctx, `
err := integrationDB.QueryRowContext(ctx, `
INSERT INTO ops_error_logs (
error_phase, error_type, severity, status_code, created_at
) VALUES (
@@ -59,20 +25,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
) RETURNING id`,
).Scan(&plainID)
require.NoError(t, err)
require.Positive(t, plainID)
plain, err := repo.GetErrorLogByID(ctx, plainID)
require.NoError(t, err)
require.NotNil(t, plain)
require.Empty(t, plain.APIKeyPrefix)
require.Empty(t, plain.AttemptedKeyPrefix, "no prefix for plain error")
require.Nil(t, plain.DeletedKeyOwnerUserID, "no owner for plain error")
require.Empty(t, plain.DeletedKeyOwnerEmail, "no owner email for plain error")
require.Empty(t, plain.DeletedKeyName, "no key name for plain error")
require.Empty(t, plain.APIKeyPrefix, "no api key prefix for plain error")
// ── Case 3: 有效(未删除)key 报错,经 InsertErrorLog 快照 api_key_prefix ──────
// 走真实 InsertErrorLog 写入路径(覆盖新列 + $41 占位符),再 GetErrorLogByID 读回。
validID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
ErrorPhase: "request",
ErrorType: "api_error",
@@ -82,17 +39,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
APIKeyPrefix: "sk-valid",
})
require.NoError(t, err)
require.Positive(t, validID)
valid, err := repo.GetErrorLogByID(ctx, validID)
require.NoError(t, err)
require.NotNil(t, valid)
require.Equal(t, "sk-valid", valid.APIKeyPrefix)
require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive")
require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner")
// ── Case 4: account_auth with no inference attempt preserves explicit 0 ──
zero := 0
credentialFailureID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
ErrorPhase: "account_auth",
@@ -1,36 +0,0 @@
//go:build integration
package repository
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestOpsRepositoryLookupDeletedKeyAudit(t *testing.T) {
ctx := context.Background()
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE deleted_api_key_audits RESTART IDENTITY")
repo := NewOpsRepository(integrationDB).(*opsRepository)
// 同一 key 两条审计,取最近一条(deleted_at DESC, id DESC)
_, err := integrationDB.ExecContext(ctx, `
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
VALUES ('sk-lookup-1', 10, 100, 'old', $1),
('sk-lookup-1', 11, 200, 'new', $2)`,
time.Now().Add(-time.Hour), time.Now())
require.NoError(t, err)
res, err := repo.LookupDeletedKeyAudit(ctx, "sk-lookup-1")
require.NoError(t, err)
require.NotNil(t, res)
require.Equal(t, int64(200), res.UserID)
require.Equal(t, "new", res.KeyName)
// 未命中返回 nil
miss, err := repo.LookupDeletedKeyAudit(ctx, "sk-never-existed")
require.NoError(t, err)
require.Nil(t, miss)
}
+27 -7
View File
@@ -1039,24 +1039,44 @@ func (r *userRepository) syncUserAllowedGroupsWithClient(ctx context.Context, cl
return nil
}
// Keep join table as the source of truth for reads.
if _, err := client.UserAllowedGroup.Delete().Where(userallowedgroup.UserIDEQ(userID)).Exec(ctx); err != nil {
existingRows, err := client.UserAllowedGroup.Query().
Where(userallowedgroup.UserIDEQ(userID)).
All(ctx)
if err != nil {
return err
}
unique := make(map[int64]struct{}, len(groupIDs))
desired := make(map[int64]struct{}, len(groupIDs))
for _, id := range groupIDs {
if id <= 0 {
continue
}
unique[id] = struct{}{}
desired[id] = struct{}{}
}
if len(unique) > 0 {
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(unique))
for groupID := range unique {
existing := make(map[int64]struct{}, len(existingRows))
removed := make([]int64, 0)
for _, row := range existingRows {
existing[row.GroupID] = struct{}{}
if _, keep := desired[row.GroupID]; !keep {
removed = append(removed, row.GroupID)
}
}
if len(removed) > 0 {
if _, err := client.UserAllowedGroup.Delete().
Where(userallowedgroup.UserIDEQ(userID), userallowedgroup.GroupIDIn(removed...)).
Exec(ctx); err != nil {
return err
}
}
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(desired))
for groupID := range desired {
if _, present := existing[groupID]; !present {
creates = append(creates, client.UserAllowedGroup.Create().SetUserID(userID).SetGroupID(groupID))
}
}
if len(creates) > 0 {
if err := client.UserAllowedGroup.
CreateBulk(creates...).
OnConflictColumns(userallowedgroup.FieldUserID, userallowedgroup.FieldGroupID).
@@ -14,7 +14,7 @@ import (
)
// TestUserRepository_DeleteUser_AtomicWithAPIKeys 复现 AdminService.DeleteUser 的事务编排场景:
// 把"删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
// 把"tombstone 并删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
// userRepo.Delete 必须复用 context 中的事务,而不是用 base client 自起一个独立事务并提前提交。
//
// 用例用"回滚外层事务"来模拟 commit 失败 / 中止:
@@ -91,5 +91,5 @@ func TestUserRepository_DeleteUser_AtomicWithAPIKeys(t *testing.T) {
require.NoError(t, integrationDB.QueryRowContext(ctx,
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE user_id = $1`, user.ID).Scan(&auditCount))
require.Equal(t, 2, auditCount, "提交后应为每个被删 Key 写入一行审计")
require.Zero(t, auditCount, "提交后也不得保留被删 Key 的凭据材料")
}
+1
View File
@@ -126,6 +126,7 @@ var ProviderSet = wire.NewSet(
NewLeaderLockCache,
ProvideSchedulerCache,
NewSchedulerOutboxRepository,
NewAuthCacheInvalidationOutboxRepository,
NewProxyLatencyCache,
NewTotpCache,
NewRefreshTokenCache,
@@ -16,6 +16,7 @@ var ProviderSet = wire.NewSet(
wire.Bind(new(ConfigStore), new(*ConfigManager)),
NewPromptService,
wire.Bind(new(PromptEngine), new(*PromptService)),
wire.Bind(new(PromptAdminService), new(*PromptService)),
NewLegacyModerationAdapter,
NewCoordinator,
NewPromptAdminHandler,
@@ -1876,6 +1876,10 @@ func (s *stubAccountRepo) ListSchedulableUngroupedByPlatforms(ctx context.Contex
return nil, errors.New("not implemented")
}
func (s *stubAccountRepo) ListModelAvailabilityCandidates(ctx context.Context, groupID *int64, platforms []string, includeGrouped bool) ([]service.Account, error) {
return nil, errors.New("not implemented")
}
func (s *stubAccountRepo) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
return errors.New("not implemented")
}
+3 -2
View File
@@ -103,8 +103,9 @@ func ProvideRouter(
func ProvideHTTPServer(cfg *config.Config, router *gin.Engine) *http.Server {
httpHandler := http.Handler(router)
server := &http.Server{
Addr: cfg.Server.Address(),
Handler: httpHandler,
Addr: cfg.Server.Address(),
Handler: httpHandler,
MaxHeaderBytes: cfg.Server.MaxHeaderBytes,
// ReadHeaderTimeout: 读取请求头的超时时间,防止慢速请求头攻击
ReadHeaderTimeout: time.Duration(cfg.Server.ReadHeaderTimeout) * time.Second,
// IdleTimeout: 空闲连接超时时间,释放不活跃的连接资源
@@ -0,0 +1,126 @@
//go:build unit
package server
import (
"bufio"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func ingressTestConfig() *config.Config {
return &config.Config{
Server: config.ServerConfig{
Host: "127.0.0.1",
ReadHeaderTimeout: 1,
IdleTimeout: 5,
MaxHeaderBytes: 8 * 1024,
MaxRequestBodySize: 1024,
},
Gateway: config.GatewayConfig{MaxBodySize: 1024},
}
}
func TestProvideHTTPServerAppliesIngressLimits(t *testing.T) {
srv := ProvideHTTPServer(ingressTestConfig(), gin.New())
require.Equal(t, 8*1024, srv.MaxHeaderBytes)
require.Equal(t, time.Second, srv.ReadHeaderTimeout)
require.Equal(t, 5*time.Second, srv.IdleTimeout)
}
func TestProvideHTTPServerEnablesBoundedH2C(t *testing.T) {
cfg := ingressTestConfig()
cfg.Server.H2C = config.H2CConfig{
Enabled: true,
MaxConcurrentStreams: 25,
IdleTimeout: 30,
MaxReadFrameSize: 64 * 1024,
MaxUploadBufferPerConnection: 1024 * 1024,
MaxUploadBufferPerStream: 256 * 1024,
}
srv := ProvideHTTPServer(cfg, gin.New())
require.NotNil(t, srv.Protocols)
require.True(t, srv.Protocols.UnencryptedHTTP2())
require.True(t, srv.Protocols.HTTP1())
}
func TestHTTPServerRejectsOversizedHTTP1Header(t *testing.T) {
r := gin.New()
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
srv := ProvideHTTPServer(ingressTestConfig(), r)
addr, stop := serveIngressTestServer(t, srv)
defer stop()
conn, err := net.DialTimeout("tcp", addr, time.Second)
require.NoError(t, err)
defer func() { _ = conn.Close() }()
_ = conn.SetDeadline(time.Now().Add(3 * time.Second))
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Fill: "+strings.Repeat("a", 32*1024)+"\r\n\r\n")
require.NoError(t, err)
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusRequestHeaderFieldsTooLarge, resp.StatusCode)
}
func TestHTTPServerClosesSlowIncompleteHeader(t *testing.T) {
r := gin.New()
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
srv := ProvideHTTPServer(ingressTestConfig(), r)
addr, stop := serveIngressTestServer(t, srv)
defer stop()
conn, err := net.DialTimeout("tcp", addr, time.Second)
require.NoError(t, err)
defer func() { _ = conn.Close() }()
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Slow:")
require.NoError(t, err)
time.Sleep(1200 * time.Millisecond)
_ = conn.SetReadDeadline(time.Now().Add(time.Second))
_, err = bufio.NewReader(conn).ReadByte()
require.Error(t, err)
}
func TestHTTPServerGlobalBodyLimit(t *testing.T) {
r := gin.New()
r.POST("/", func(c *gin.Context) {
_, err := io.ReadAll(c.Request.Body)
if err != nil {
var maxErr *http.MaxBytesError
if errors.As(err, &maxErr) {
c.Status(http.StatusRequestEntityTooLarge)
return
}
}
c.Status(http.StatusOK)
})
srv := ProvideHTTPServer(ingressTestConfig(), r)
req, err := http.NewRequest(http.MethodPost, "/", strings.NewReader(strings.Repeat("x", 1025)))
require.NoError(t, err)
rec := httptest.NewRecorder()
srv.Handler.ServeHTTP(rec, req)
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code)
}
func serveIngressTestServer(t *testing.T, srv *http.Server) (string, func()) {
t.Helper()
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
go func() { _ = srv.Serve(ln) }()
return ln.Addr().String(), func() {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = srv.Shutdown(ctx)
}
}
@@ -15,6 +15,8 @@ import (
"github.com/gin-gonic/gin"
)
const maxAPIKeyAuthorizationHeaderBytes = service.MaxAPIKeyCredentialBytes + 128
// NewAPIKeyAuthMiddleware 创建 API Key 认证中间件
func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) APIKeyAuthMiddleware {
return APIKeyAuthMiddleware(apiKeyAuthWithSubscription(apiKeyService, subscriptionService, cfg))
@@ -32,10 +34,23 @@ func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionS
func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
// ── 1. 提取 API Key ──────────────────────────────────────────
if rejectInvalidAuthAbuse(c, apiKeyService) {
AbortWithError(c, http.StatusTooManyRequests, "INVALID_AUTH_RATE_LIMITED", "Too many invalid authentication attempts; retry later")
return
}
if apiKeyHeadersTooLarge(c) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
return
}
queryKey := strings.TrimSpace(c.Query("key"))
queryApiKey := strings.TrimSpace(c.Query("api_key"))
if queryKey != "" || queryApiKey != "" {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
AbortWithError(c, 400, "api_key_in_query_deprecated", "API key in query parameter is deprecated. Please use Authorization header instead.")
return
}
@@ -56,6 +71,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
if apiKeyString == "" {
apiKeyString = c.GetHeader("x-api-key")
}
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
return
}
// 如果x-api-key header中没有,尝试从x-goog-api-key header中提取(Gemini CLI兼容)
if apiKeyString == "" {
@@ -64,6 +85,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
// 如果所有header都没有API key
if apiKeyString == "" {
recordInvalidAuthFailure(c, apiKeyService)
if hasAPIKeyCredentialInput(c) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
} else {
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
}
AbortWithError(c, 401, "API_KEY_REQUIRED", "API key is required in Authorization header (Bearer scheme), x-api-key header, or x-goog-api-key header")
return
}
@@ -73,9 +100,16 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
if err != nil {
if errors.Is(err, service.ErrAPIKeyNotFound) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, 401, "INVALID_API_KEY", "Invalid API key")
return
}
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
AbortWithError(c, http.StatusServiceUnavailable, "API_KEY_AUTH_OVERLOADED", "API key authentication is temporarily unavailable")
return
}
AbortWithError(c, 500, "INTERNAL_ERROR", "Failed to validate API key")
return
}
@@ -90,6 +124,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
if !apiKey.IsActive() &&
apiKey.Status != service.StatusAPIKeyExpired &&
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
AbortWithError(c, 401, "API_KEY_DISABLED", "API key is disabled")
return
}
@@ -104,6 +139,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
clientIP = "unknown"
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
MarkIngressRejected(c, IngressRejectIPRestricted)
AbortWithError(c, 403, "ACCESS_DENIED", fmt.Sprintf("Access denied. Your IP is %s", clientIP))
return
}
@@ -117,6 +153,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
// 检查用户状态
if !apiKey.User.IsActive() {
MarkIngressRejected(c, IngressRejectUserInactive)
AbortWithError(c, 401, "USER_INACTIVE", "User account is not active")
return
}
@@ -250,6 +287,24 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
}
}
func apiKeyHeadersTooLarge(c *gin.Context) bool {
if c == nil {
return false
}
return len(c.GetHeader("Authorization")) > maxAPIKeyAuthorizationHeaderBytes ||
len(c.GetHeader("x-api-key")) > service.MaxAPIKeyCredentialBytes ||
len(c.GetHeader("x-goog-api-key")) > service.MaxAPIKeyCredentialBytes
}
func hasAPIKeyCredentialInput(c *gin.Context) bool {
if c == nil {
return false
}
return c.GetHeader("Authorization") != "" ||
c.GetHeader("x-api-key") != "" ||
c.GetHeader("x-goog-api-key") != ""
}
func isAsyncImageTaskRead(method, path string) bool {
if method != http.MethodGet {
return false
@@ -321,6 +376,11 @@ func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool
return false
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
if code == "GROUP_DELETED" {
MarkIngressRejected(c, IngressRejectGroupDeleted)
} else {
MarkIngressRejected(c, IngressRejectGroupDisabled)
}
AbortWithError(c, 403, code, message)
return true
}
@@ -330,6 +390,7 @@ func abortIfAPIKeyGroupNotAllowed(c *gin.Context, apiKey *service.APIKey) bool {
return false
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
AbortWithError(c, 403, "GROUP_NOT_ALLOWED", "API Key 所属专属分组不再允许当前用户使用")
return true
}
@@ -24,22 +24,53 @@ func APIKeyAuthGoogle(apiKeyService *service.APIKeyService, cfg *config.Config)
// It is intended for Gemini native endpoints (/v1beta) to match Gemini SDK expectations.
func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
if rejectInvalidAuthAbuse(c, apiKeyService) {
abortWithGoogleError(c, 429, "Too many invalid authentication attempts; retry later")
return
}
if apiKeyHeadersTooLarge(c) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
if v := strings.TrimSpace(c.Query("api_key")); v != "" {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
abortWithGoogleError(c, 400, "Query parameter api_key is deprecated. Use Authorization header or key instead.")
return
}
apiKeyString := extractAPIKeyForGoogle(c)
if apiKeyString == "" {
recordInvalidAuthFailure(c, apiKeyService)
if hasAPIKeyCredentialInput(c) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
} else {
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
}
abortWithGoogleError(c, 401, "API key is required")
return
}
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
if err != nil {
if errors.Is(err, service.ErrAPIKeyNotFound) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
abortWithGoogleError(c, 503, "API key authentication is temporarily unavailable")
return
}
abortWithGoogleError(c, 500, "Failed to validate API key")
return
}
@@ -53,6 +84,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
if !apiKey.IsActive() &&
apiKey.Status != service.StatusAPIKeyExpired &&
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
abortWithGoogleError(c, 401, "API key is disabled")
return
}
@@ -66,6 +98,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
clientIP = "unknown"
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
MarkIngressRejected(c, IngressRejectIPRestricted)
abortWithGoogleError(c, 403, fmt.Sprintf("Access denied. Your IP is %s", clientIP))
return
}
@@ -76,17 +109,24 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
return
}
if !apiKey.User.IsActive() {
MarkIngressRejected(c, IngressRejectUserInactive)
abortWithGoogleError(c, 401, "User account is not active")
return
}
if _, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
if code, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
if code == "GROUP_DELETED" {
MarkIngressRejected(c, IngressRejectGroupDeleted)
} else {
MarkIngressRejected(c, IngressRejectGroupDisabled)
}
abortWithGoogleError(c, 403, message)
return
}
// 专属分组授权校验:用户对该专属分组的授权被撤销后应拒绝(与主中间件一致,防止越权)。
if !validateAPIKeyGroupAllowed(apiKey) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
abortWithGoogleError(c, 403, "API Key 所属专属分组不再允许当前用户使用")
return
}
@@ -6,6 +6,8 @@ import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
@@ -18,6 +20,59 @@ import (
"github.com/stretchr/testify/require"
)
func TestGoogleAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
gin.SetMode(gin.TestMode)
var calls atomic.Int32
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
calls.Add(1)
return nil, service.ErrAPIKeyNotFound
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthGoogle(svc, cfg))
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
req.Header.Set("x-goog-api-key", strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1))
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Zero(t, calls.Load())
require.True(t, rejected)
require.Equal(t, IngressRejectInvalidAPIKey, reason)
}
func TestGoogleAPIKeyAuthMarksLookupBulkheadRejection(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
return nil, service.ErrAPIKeyAuthOverloaded
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthGoogle(svc, cfg))
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
req.Header.Set("x-goog-api-key", "valid-shape")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusServiceUnavailable, w.Code)
require.True(t, rejected)
require.Equal(t, IngressRejectAPIKeyAuthOverloaded, reason)
}
type fakeAPIKeyRepo struct {
getByKey func(ctx context.Context, key string) (*service.APIKey, error)
updateLastUsed func(ctx context.Context, id int64, usedAt time.Time) error
@@ -376,6 +431,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
return nil, service.ErrAPIKeyNotFound
},
})
var rejectReason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
rejectReason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, &config.Config{}))
r.GET("/v1beta/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
@@ -390,6 +451,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
require.Equal(t, http.StatusUnauthorized, resp.Error.Code)
require.Equal(t, "Invalid API key", resp.Error.Message)
require.Equal(t, "UNAUTHENTICATED", resp.Error.Status)
require.True(t, rejected)
require.Equal(t, IngressRejectInvalidAPIKey, rejectReason)
}
func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t *testing.T) {
@@ -422,9 +485,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
r := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -452,6 +518,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
require.Equal(t, "API Key 所属分组已删除", resp.Error.Message)
require.True(t, markedBusinessLimited)
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
require.True(t, rejected)
require.Equal(t, IngressRejectGroupDeleted, rejectReason)
}
func TestApiKeyAuthWithSubscriptionGoogle_RepoError(t *testing.T) {
@@ -8,6 +8,8 @@ import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
@@ -19,6 +21,35 @@ import (
"github.com/stretchr/testify/require"
)
func TestAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
gin.SetMode(gin.TestMode)
var calls atomic.Int32
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
calls.Add(1)
return nil, service.ErrAPIKeyNotFound
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
for _, headers := range []map[string]string{
{"x-api-key": strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
{"Authorization": "Bearer " + strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
{"Authorization": strings.Repeat("x", maxAPIKeyAuthorizationHeaderBytes+1)},
} {
r := gin.New()
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/t", nil)
for name, value := range headers {
req.Header.Set(name, value)
}
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
}
require.Zero(t, calls.Load())
}
func TestSimpleModeBypassesQuotaCheck(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -436,6 +467,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus int
wantCode string
wantMarked bool
wantReject IngressRejectReason
}{
{
name: "active group passes",
@@ -460,6 +492,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DISABLED",
wantMarked: true,
wantReject: IngressRejectGroupDisabled,
},
{
name: "deleted status group is forbidden",
@@ -473,6 +506,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DELETED",
wantMarked: true,
wantReject: IngressRejectGroupDeleted,
},
{
name: "missing group edge is forbidden",
@@ -480,6 +514,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DELETED",
wantMarked: true,
wantReject: IngressRejectGroupDeleted,
},
}
@@ -508,9 +543,12 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
router := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -530,6 +568,8 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
require.Contains(t, w.Body.String(), tt.wantCode)
}
require.Equal(t, tt.wantMarked, markedBusinessLimited)
require.Equal(t, tt.wantReject != "", rejected)
require.Equal(t, tt.wantReject, rejectReason)
if tt.wantMarked {
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
}
@@ -537,6 +577,112 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
}
}
func TestAPIKeyAuthMarksOnlyExpectedIngressRejections(t *testing.T) {
gin.SetMode(gin.TestMode)
tests := []struct {
name string
path string
key string
authHeader string
repoErr error
wantStatus int
wantCode string
wantReason IngressRejectReason
}{
{
name: "query key deprecated",
path: "/t?key=legacy",
wantStatus: http.StatusBadRequest,
wantCode: "api_key_in_query_deprecated",
wantReason: IngressRejectQueryAPIKeyDeprecated,
},
{
name: "missing key",
path: "/t",
wantStatus: http.StatusUnauthorized,
wantCode: "API_KEY_REQUIRED",
wantReason: IngressRejectAPIKeyRequired,
},
{
name: "malformed authorization",
path: "/t",
authHeader: "Basic not-a-bearer-key",
wantStatus: http.StatusUnauthorized,
wantCode: "API_KEY_REQUIRED",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "oversized key",
path: "/t",
key: strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1),
wantStatus: http.StatusUnauthorized,
wantCode: "INVALID_API_KEY",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "invalid key",
path: "/t",
key: "invalid",
repoErr: service.ErrAPIKeyNotFound,
wantStatus: http.StatusUnauthorized,
wantCode: "INVALID_API_KEY",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "repository failure remains operational error",
path: "/t",
key: "valid-shape",
repoErr: errors.New("database unavailable"),
wantStatus: http.StatusInternalServerError,
wantCode: "INTERNAL_ERROR",
},
{
name: "auth lookup bulkhead rejection is an admission rejection",
path: "/t",
key: "valid-shape",
repoErr: service.ErrAPIKeyAuthOverloaded,
wantStatus: http.StatusServiceUnavailable,
wantCode: "API_KEY_AUTH_OVERLOADED",
wantReason: IngressRejectAPIKeyAuthOverloaded,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
return nil, tt.repoErr
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
apiKeyService := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
var reason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, tt.path, nil)
if tt.key != "" {
req.Header.Set("x-api-key", tt.key)
}
if tt.authHeader != "" {
req.Header.Set("Authorization", tt.authHeader)
}
router.ServeHTTP(w, req)
require.Equal(t, tt.wantStatus, w.Code)
require.Contains(t, w.Body.String(), tt.wantCode)
require.Equal(t, tt.wantReason != "", rejected)
require.Equal(t, tt.wantReason, reason)
})
}
}
func TestAPIKeyAuthSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -686,9 +832,12 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
router := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -708,6 +857,8 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
require.Equal(t, http.StatusForbidden, w.Code)
require.Contains(t, w.Body.String(), "not assigned to any group")
require.True(t, rejected)
require.Equal(t, IngressRejectGroupUnassigned, rejectReason)
require.True(t, markedBusinessLimited)
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned, businessLimitedReason)
}
@@ -822,7 +973,7 @@ func TestAPIKeyAuthIPRestrictionIncludesClientIPForBlacklistDenial(t *testing.T)
requireAPIKeyAuthError(t, w, "ACCESS_DENIED", "Access denied. Your IP is 9.9.9.9")
}
func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *testing.T) {
func TestAPIKeyAuthIPRestrictionUsesConfiguredTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
user := &service.User{
@@ -855,7 +1006,7 @@ func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *test
cfg.SetTrustForwardedIPForAPIKeyACL(true)
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
require.NoError(t, router.SetTrustedProxies(nil))
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -906,7 +1057,7 @@ func TestAPIKeyAuthIPRestrictionUsesForwardedClientIPInDenialWhenTrusted(t *test
cfg.SetTrustForwardedIPForAPIKeyACL(true)
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
require.NoError(t, router.SetTrustedProxies(nil))
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -24,7 +24,14 @@ func ClientRequestID() gin.HandlerFunc {
}
if v, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(v) != "" {
c.Header(clientRequestIDHeader, strings.TrimSpace(v))
var valid bool
v, valid = normalizeCorrelationID(v)
if !valid {
v = uuid.New().String()
}
c.Header(clientRequestIDHeader, v)
ctx := context.WithValue(c.Request.Context(), ctxkey.ClientRequestID, v)
c.Request = c.Request.WithContext(ctx)
c.Next()
return
}
@@ -4,6 +4,7 @@ import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
@@ -30,6 +31,23 @@ func TestClientRequestIDGeneratesAndExposesID(t *testing.T) {
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
}
func TestClientRequestIDBoundsExistingContextID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(ClientRequestID())
router.GET("/", func(c *gin.Context) {
value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
c.String(http.StatusOK, value)
})
req := httptest.NewRequest(http.MethodGet, "/", nil)
req = req.WithContext(context.WithValue(req.Context(), ctxkey.ClientRequestID, strings.Repeat("x", 200)))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Len(t, w.Body.String(), 36)
require.NotEqual(t, strings.Repeat("x", maxPersistentRequestIDBytes), w.Body.String())
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
}
func TestClientRequestIDPreservesExistingContextID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
@@ -0,0 +1,166 @@
package middleware
import (
"math"
"net/netip"
"strconv"
"strings"
"sync/atomic"
"time"
"github.com/gin-gonic/gin"
)
// IngressRejectReason identifies expected gateway admission failures that must
// not be treated as operational request errors.
type IngressRejectReason string
const (
IngressRejectQueryAPIKeyDeprecated IngressRejectReason = "query_api_key_deprecated"
IngressRejectAPIKeyRequired IngressRejectReason = "api_key_required"
IngressRejectInvalidAPIKey IngressRejectReason = "invalid_api_key"
IngressRejectAPIKeyDisabled IngressRejectReason = "api_key_disabled"
IngressRejectIPRestricted IngressRejectReason = "ip_restricted"
IngressRejectUserInactive IngressRejectReason = "user_inactive"
IngressRejectGroupDeleted IngressRejectReason = "group_deleted"
IngressRejectGroupDisabled IngressRejectReason = "group_disabled"
IngressRejectGroupNotAllowed IngressRejectReason = "group_not_allowed"
IngressRejectGroupUnassigned IngressRejectReason = "group_unassigned"
IngressRejectInvalidAuthRateLimited IngressRejectReason = "invalid_auth_rate_limited"
IngressRejectAPIKeyAuthOverloaded IngressRejectReason = "api_key_auth_overloaded"
)
const ingressRejectReasonContextKey = "ingress_reject_reason"
type IngressRejectRecorder interface {
RecordIngressReject(reason, routeFamily, protocol, clientIP string, userID, apiKeyID int64)
}
func invalidAuthClientKey(c *gin.Context) string {
return normalizeIngressRejectIP(SecurityClientIP(c))
}
func rejectInvalidAuthAbuse(c *gin.Context, apiKeyService interface {
CheckInvalidAuthAbuse(string) (time.Duration, bool)
}) bool {
if c == nil || apiKeyService == nil {
return false
}
retry, blocked := apiKeyService.CheckInvalidAuthAbuse(invalidAuthClientKey(c))
if !blocked {
return false
}
retrySeconds := int(math.Ceil(retry.Seconds()))
if retrySeconds < 1 {
retrySeconds = 1
}
c.Header("Retry-After", strconv.Itoa(retrySeconds))
MarkIngressRejected(c, IngressRejectInvalidAuthRateLimited)
return true
}
func recordInvalidAuthFailure(c *gin.Context, apiKeyService interface {
RecordInvalidAuthFailure(string)
}) {
if c == nil || apiKeyService == nil {
return
}
apiKeyService.RecordInvalidAuthFailure(invalidAuthClientKey(c))
}
type ingressRejectRecorderHolder struct{ recorder IngressRejectRecorder }
var activeIngressRejectRecorder atomic.Pointer[ingressRejectRecorderHolder]
func SetIngressRejectRecorder(recorder IngressRejectRecorder) {
if recorder == nil {
activeIngressRejectRecorder.Store(nil)
return
}
activeIngressRejectRecorder.Store(&ingressRejectRecorderHolder{recorder: recorder})
}
// MarkIngressRejected marks a request as rejected before gateway admission.
func MarkIngressRejected(c *gin.Context, reason IngressRejectReason) {
if c == nil || reason == "" {
return
}
c.Set(ingressRejectReasonContextKey, reason)
}
// GetIngressRejectReason returns the admission rejection reason, if any.
func GetIngressRejectReason(c *gin.Context) (IngressRejectReason, bool) {
if c == nil {
return "", false
}
value, exists := c.Get(ingressRejectReasonContextKey)
if !exists {
return "", false
}
reason, ok := value.(IngressRejectReason)
return reason, ok && reason != ""
}
func recordIngressReject(c *gin.Context, reason IngressRejectReason) {
holder := activeIngressRejectRecorder.Load()
if holder == nil || holder.recorder == nil || c == nil || c.Request == nil {
return
}
routeFamily, protocol := ingressRejectRoute(c.Request.URL.Path)
clientIP := normalizeIngressRejectIP(SecurityClientIP(c))
var userID, apiKeyID int64
if apiKey, ok := GetAPIKeyFromContext(c); ok && apiKey != nil {
apiKeyID = apiKey.ID
if apiKey.User != nil {
userID = apiKey.User.ID
}
} else if apiKey, ok := GetOpsFallbackAPIKey(c); ok && apiKey != nil {
apiKeyID = apiKey.ID
if apiKey.User != nil {
userID = apiKey.User.ID
}
}
holder.recorder.RecordIngressReject(string(reason), routeFamily, protocol, clientIP, userID, apiKeyID)
}
func normalizeIngressRejectIP(raw string) string {
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
if err != nil {
return "0.0.0.0"
}
addr = addr.Unmap()
if addr.Is6() {
return netip.PrefixFrom(addr, 64).Masked().Addr().String()
}
return addr.String()
}
func ingressRejectRoute(path string) (string, string) {
path = strings.ToLower(strings.TrimSpace(path))
switch {
case strings.HasPrefix(path, "/antigravity/v1beta"):
return "antigravity", "google"
case strings.HasPrefix(path, "/v1beta"):
return "gemini", "google"
case strings.HasPrefix(path, "/backend-api/codex"):
return "codex", "openai"
case strings.HasPrefix(path, "/antigravity"):
return "antigravity", "anthropic"
case strings.Contains(path, "/messages"):
return "messages", "anthropic"
case strings.Contains(path, "/responses"):
return "responses", "openai"
case strings.Contains(path, "/chat/completions"):
return "chat_completions", "openai"
case strings.Contains(path, "/images"):
return "images", "openai"
case strings.Contains(path, "/videos"):
return "videos", "openai"
case strings.Contains(path, "/embeddings"):
return "embeddings", "openai"
case strings.Contains(path, "/models"):
return "models", "openai"
default:
return "other", "gateway"
}
}
@@ -0,0 +1,58 @@
package middleware
import (
"sync"
"time"
)
const (
ingressRejectAccessLogLimit = 20
ingressRejectAccessLogWindow = time.Second
ingressRejectDroppedSummaryPeriod = 30 * time.Second
)
type ingressRejectAccessSampler struct {
mu sync.Mutex
limit int
window time.Duration
summaryPeriod time.Duration
windowStart time.Time
emitted int
dropped uint64
lastSummary time.Time
}
func newIngressRejectAccessSampler(limit int, window, summaryPeriod time.Duration) *ingressRejectAccessSampler {
return &ingressRejectAccessSampler{limit: limit, window: window, summaryPeriod: summaryPeriod}
}
// allow applies one process-wide fixed-window budget. It stores no attacker
// dimensions, so memory remains constant even for rotating keys and addresses.
func (s *ingressRejectAccessSampler) allow(now time.Time) (allowed bool, droppedSummary uint64) {
if s == nil || s.limit <= 0 || s.window <= 0 {
return false, 0
}
s.mu.Lock()
defer s.mu.Unlock()
if s.windowStart.IsZero() || now.Sub(s.windowStart) >= s.window || now.Before(s.windowStart) {
s.windowStart = now
s.emitted = 0
}
if s.emitted < s.limit {
s.emitted++
return true, 0
}
s.dropped++
if s.summaryPeriod > 0 && (s.lastSummary.IsZero() || now.Sub(s.lastSummary) >= s.summaryPeriod) {
droppedSummary = s.dropped
s.dropped = 0
s.lastSummary = now
}
return false, droppedSummary
}
var globalIngressRejectAccessSampler = newIngressRejectAccessSampler(
ingressRejectAccessLogLimit,
ingressRejectAccessLogWindow,
ingressRejectDroppedSummaryPeriod,
)
@@ -0,0 +1,81 @@
package middleware
import (
"net/http"
"net/http/httptest"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestIngressRejectAccessSamplerConcurrentGlobalLimit(t *testing.T) {
sampler := newIngressRejectAccessSampler(10, time.Hour, time.Minute)
now := time.Now()
var allowed atomic.Int64
var wg sync.WaitGroup
for i := 0; i < 200; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if ok, _ := sampler.allow(now); ok {
allowed.Add(1)
}
}()
}
wg.Wait()
require.Equal(t, int64(10), allowed.Load())
}
func TestLoggerIngressRejectSamplingIsBoundedAndSummarySkipsOpsSink(t *testing.T) {
gin.SetMode(gin.TestMode)
original := globalIngressRejectAccessSampler
globalIngressRejectAccessSampler = newIngressRejectAccessSampler(2, time.Hour, time.Hour)
t.Cleanup(func() { globalIngressRejectAccessSampler = original })
sink := initMiddlewareTestLogger(t)
router := gin.New()
router.Use(Logger())
router.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
for i := 0; i < 20; i++ {
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/v1/messages", nil))
}
var accessEvents, summaries int
for _, event := range sink.list() {
switch event.Message {
case "http request completed":
accessEvents++
case "ingress rejection access logs dropped":
summaries++
if skipped, _ := event.Fields[logger.OpsSystemLogSkipField].(bool); !skipped {
t.Fatalf("dropped summary must skip ops system log sink")
}
}
}
require.Equal(t, 2, accessEvents)
require.Equal(t, 1, summaries)
}
func TestIngressRejectAccessSamplerDroppedSummaryIsLowFrequency(t *testing.T) {
sampler := newIngressRejectAccessSampler(1, time.Hour, time.Second)
now := time.Now()
allowed, summary := sampler.allow(now)
require.True(t, allowed)
require.Zero(t, summary)
allowed, summary = sampler.allow(now.Add(100 * time.Millisecond))
require.False(t, allowed)
require.Equal(t, uint64(1), summary)
allowed, summary = sampler.allow(now.Add(200 * time.Millisecond))
require.False(t, allowed)
require.Zero(t, summary)
allowed, summary = sampler.allow(now.Add(2 * time.Second))
require.False(t, allowed)
require.Equal(t, uint64(2), summary)
}
@@ -0,0 +1,50 @@
package middleware
import (
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type ingressRejectRecorderStub struct {
mu sync.Mutex
calls int
clientIP string
}
func (r *ingressRejectRecorderStub) RecordIngressReject(_, _, _, clientIP string, _, _ int64) {
r.mu.Lock()
defer r.mu.Unlock()
r.calls++
r.clientIP = clientIP
}
func TestNormalizeIngressRejectIP(t *testing.T) {
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::1"))
require.Equal(t, "192.0.2.4", normalizeIngressRejectIP("::ffff:192.0.2.4"))
require.Equal(t, "0.0.0.0", normalizeIngressRejectIP("not-an-ip"))
}
func TestLoggerRecordsIngressRejectOnce(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := &ingressRejectRecorderStub{}
SetIngressRejectRecorder(recorder)
t.Cleanup(func() { SetIngressRejectRecorder(nil) })
router := gin.New()
router.Use(Logger())
router.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
request := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
request.RemoteAddr = "[2001:db8:abcd:1234:ffff::1]:1234"
router.ServeHTTP(httptest.NewRecorder(), request)
recorder.mu.Lock()
require.Equal(t, 1, recorder.calls)
require.Equal(t, "2001:db8:abcd:1234::", recorder.clientIP)
recorder.mu.Unlock()
}
@@ -0,0 +1,147 @@
//go:build unit
package middleware
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func invalidAuthAbuseTestConfig(threshold int) *config.Config {
return &config.Config{
RunMode: config.RunModeSimple,
APIKeyAuth: config.APIKeyAuthCacheConfig{InvalidAbuse: config.InvalidAuthAbuseConfig{
Enabled: true, Threshold: threshold, WindowSeconds: 60, BlockSeconds: 60, Capacity: 256,
}},
}
}
func TestAPIKeyAuthInvalidAbuseReturns429BeforeRepository(t *testing.T) {
gin.SetMode(gin.TestMode)
repoCalls := 0
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
repoCalls++
return nil, service.ErrAPIKeyNotFound
}}
cfg := invalidAuthAbuseTestConfig(3)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.POST("/v1/messages", func(c *gin.Context) { c.Status(http.StatusOK) })
requests := []*http.Request{
httpRequest(t, "/v1/messages", "", ""),
httpRequest(t, "/v1/messages", "Basic malformed", ""),
httpRequest(t, "/v1/messages", "", "random-invalid-key"),
}
for _, req := range requests {
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
require.NotEqual(t, http.StatusTooManyRequests, w.Code)
}
w := httptest.NewRecorder()
r.ServeHTTP(w, httpRequest(t, "/v1/messages", "", "another-random-key"))
require.Equal(t, http.StatusTooManyRequests, w.Code)
require.Equal(t, "60", w.Header().Get("Retry-After"))
require.Contains(t, w.Body.String(), "INVALID_AUTH_RATE_LIMITED")
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
require.Equal(t, 1, repoCalls, "rate-limited request must not reach the repository")
}
func TestGoogleAPIKeyAuthInvalidAbuseReturnsProtocol429(t *testing.T) {
gin.SetMode(gin.TestMode)
repoCalls := 0
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
repoCalls++
return nil, service.ErrAPIKeyNotFound
}}
cfg := invalidAuthAbuseTestConfig(2)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
r.Use(APIKeyAuthGoogle(svc, cfg))
r.POST("/v1beta/models/test:generateContent", func(c *gin.Context) { c.Status(http.StatusOK) })
for _, key := range []string{"random-1", "random-2"} {
w := httptest.NewRecorder()
req := httpRequest(t, "/v1beta/models/test:generateContent", "", key)
req.Header.Del("x-api-key")
req.Header.Set("x-goog-api-key", key)
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
}
w := httptest.NewRecorder()
req := httpRequest(t, "/v1beta/models/test:generateContent", "", "random-3")
req.Header.Del("x-api-key")
req.Header.Set("x-goog-api-key", "random-3")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusTooManyRequests, w.Code)
require.Equal(t, "60", w.Header().Get("Retry-After"))
require.Contains(t, w.Body.String(), "RESOURCE_EXHAUSTED")
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
require.Equal(t, 2, repoCalls)
}
func TestInvalidAuthAbuseDoesNotCountValidOrOperationalFailures(t *testing.T) {
gin.SetMode(gin.TestMode)
user := &service.User{ID: 1, Status: service.StatusActive, Role: service.RoleUser, Balance: 1}
repo := &stubApiKeyRepo{getByKey: func(_ context.Context, key string) (*service.APIKey, error) {
switch key {
case "valid-key":
return &service.APIKey{ID: 1, UserID: 1, Key: key, Status: service.StatusActive, User: user}, nil
case "db-error":
return nil, errors.New("database unavailable")
default:
return nil, service.ErrAPIKeyNotFound
}
}}
cfg := invalidAuthAbuseTestConfig(10)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.POST("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
for _, tc := range []struct {
key string
want int
}{{"invalid", 401}, {"valid-key", 200}, {"db-error", 500}, {"db-error", 500}} {
w := httptest.NewRecorder()
r.ServeHTTP(w, httpRequest(t, "/t", "", tc.key))
require.Equal(t, tc.want, w.Code)
}
w := httptest.NewRecorder()
req := httpRequest(t, "/t", "", "")
req.Header.Set("x-goog-api-key", "valid-key")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
require.Equal(t, uint64(1), svc.InvalidAuthAbuseHealth().Recorded)
}
func TestNormalizeIngressRejectIPGroupsIPv6By64(t *testing.T) {
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"))
require.Equal(t, normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"), normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::2"))
}
func httpRequest(t *testing.T, path, authorization, apiKey string) *http.Request {
t.Helper()
req := httptest.NewRequest(http.MethodPost, path, nil)
req.RemoteAddr = "203.0.113.10:12345"
if authorization != "" {
req.Header.Set("Authorization", authorization)
}
if apiKey != "" {
req.Header.Set("x-api-key", apiKey)
}
return req
}
@@ -37,6 +37,21 @@ func Logger() gin.HandlerFunc {
accountID, hasAccountID := c.Request.Context().Value(ctxkey.AccountID).(int64)
platform, _ := c.Request.Context().Value(ctxkey.Platform).(string)
model, _ := c.Request.Context().Value(ctxkey.Model).(string)
reason, rejected := GetIngressRejectReason(c)
if rejected {
recordIngressReject(c, reason)
allowed, droppedSummary := globalIngressRejectAccessSampler.allow(endTime)
if droppedSummary > 0 {
logger.FromContext(c.Request.Context()).Info("ingress rejection access logs dropped",
zap.String("component", "http.access"),
zap.Uint64("dropped_count", droppedSummary),
zap.Bool(logger.OpsSystemLogSkipField, true),
)
}
if !allowed {
return
}
}
fields := []zap.Field{
zap.String("component", "http.access"),
@@ -47,6 +62,12 @@ func Logger() gin.HandlerFunc {
zap.String("method", method),
zap.String("path", path),
}
if rejected {
fields = append(fields,
zap.String("ingress_reject_reason", string(reason)),
zap.Bool(logger.OpsSystemLogSkipField, true),
)
}
if hasAccountID && accountID > 0 {
fields = append(fields, zap.Int64("account_id", accountID))
}
@@ -121,6 +121,7 @@ func RequireGroupAssignment(settingService *service.SettingService, writeError G
return
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned)
MarkIngressRejected(c, IngressRejectGroupUnassigned)
writeError(c, http.StatusForbidden, "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.")
c.Abort()
}
@@ -112,6 +112,26 @@ func TestRequestLogger_KeepIncomingRequestID(t *testing.T) {
}
}
func TestRequestLoggerBoundsIncomingRequestID(t *testing.T) {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(RequestLogger())
r.GET("/t", func(c *gin.Context) {
reqID, _ := c.Request.Context().Value(ctxkey.RequestID).(string)
if len(reqID) != 36 {
t.Fatalf("request_id length=%d", len(reqID))
}
c.Status(http.StatusOK)
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/t", nil)
req.Header.Set(requestIDHeader, strings.Repeat("r", 1024))
r.ServeHTTP(w, req)
if got := len(w.Header().Get(requestIDHeader)); got != 36 {
t.Fatalf("response request_id length=%d", got)
}
}
func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
@@ -180,11 +200,43 @@ func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
}
}
func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
func TestLogger_IngressRejectRemainsInStandardAccessLog(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
r := gin.New()
r.Use(Logger())
r.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("status=%d", w.Code)
}
events := sink.list()
if len(events) != 1 {
t.Fatalf("events=%d, want 1", len(events))
}
if got := events[0].Fields["ingress_reject_reason"]; got != string(IngressRejectInvalidAPIKey) {
t.Fatalf("ingress_reject_reason=%v", got)
}
if got, _ := events[0].Fields[logger.OpsSystemLogSkipField].(bool); !got {
t.Fatalf("%s must be true", logger.OpsSystemLogSkipField)
}
}
func TestLogger_AccessLogUsesForwardedClientIPFromTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
r := gin.New()
if err := r.SetTrustedProxies([]string{"104.23.251.120"}); err != nil {
t.Fatalf("set trusted proxies: %v", err)
}
r.Use(Logger())
r.GET("/api/test", func(c *gin.Context) {
c.Status(http.StatusOK)
@@ -193,7 +245,7 @@ func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/test", nil)
req.RemoteAddr = "104.23.251.120:443"
req.Header.Set("CF-Connecting-IP", "203.0.113.42")
req.Header.Set("X-Forwarded-For", "203.0.113.42")
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status=%d", w.Code)
@@ -21,14 +21,15 @@ func RequestLogger() gin.HandlerFunc {
return
}
requestID := strings.TrimSpace(c.GetHeader(requestIDHeader))
if requestID == "" {
requestID, validRequestID := normalizeCorrelationID(c.GetHeader(requestIDHeader))
if !validRequestID {
requestID = uuid.NewString()
}
c.Header(requestIDHeader, requestID)
ctx := context.WithValue(c.Request.Context(), ctxkey.RequestID, requestID)
clientRequestID, _ := ctx.Value(ctxkey.ClientRequestID).(string)
clientRequestID, _ = normalizeCorrelationID(clientRequestID)
requestLogger := logger.With(
zap.String("component", "http"),
@@ -0,0 +1,30 @@
package middleware
import (
"strings"
"unicode/utf8"
)
const (
maxPersistentRequestIDBytes = 64
maxPersistentUserAgentBytes = 512
)
// normalizePersistentText bounds attacker-controlled metadata before it reaches
// logs or database columns while preserving valid UTF-8 content.
func normalizePersistentText(value string, maxBytes int) string {
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
if maxBytes <= 0 || len(value) <= maxBytes {
return value
}
value = value[:maxBytes]
for !utf8.ValidString(value) {
value = value[:len(value)-1]
}
return value
}
func normalizeCorrelationID(value string) (string, bool) {
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
return value, value != "" && len(value) <= maxPersistentRequestIDBytes
}
@@ -13,14 +13,15 @@ import (
// SessionBindingContext 全局中间件:将请求的客户端 IP 与 User-Agent 注入
// request context,供 token 签发路径(登录 / 刷新 / OAuth 回调)读取并写入会话绑定,
// 同时作为审计日志、会话绑定校验的统一客户端 IP 来源。
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关:
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
// 关闭时走 trusted_proxies 解析链,避免不可信头伪造绕过绑定。
// IP 取值与 API Key IP 限制共用 Gin trusted_proxies 解析链;旧设置开关
// 仅为配置兼容保留,不能单独使直连请求的转发头变为可信。
func SessionBindingContext(cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
userAgent := normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes)
c.Request.Header.Set("User-Agent", userAgent)
binding := &service.SessionBinding{
IP: ip.GetSecurityClientIP(c, cfg.TrustForwardedIPForAPIKeyACL()),
UserAgent: c.Request.UserAgent(),
UserAgent: userAgent,
}
c.Request = c.Request.WithContext(service.WithSessionBinding(c.Request.Context(), binding))
c.Next()
@@ -36,7 +37,7 @@ func requestSessionBinding(c *gin.Context) *service.SessionBinding {
}
return &service.SessionBinding{
IP: ip.GetTrustedClientIP(c),
UserAgent: c.Request.UserAgent(),
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
}
}
@@ -93,7 +94,7 @@ func enforceSessionBinding(
Method: c.Request.Method,
Path: path,
ClientIP: binding.IP,
UserAgent: c.Request.UserAgent(),
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
StatusCode: 401,
})
}
@@ -4,6 +4,7 @@ package middleware
import (
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -13,9 +14,7 @@ import (
"github.com/stretchr/testify/require"
)
// 反代场景:RemoteAddr 为 127.0.0.1,真实客户端 IP 在 X-Real-IP 中。
// 会话绑定注入与审计 IP 必须与 API Key IP 限制共用「信任反代传递的客户端 IP」开关语义。
func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
func TestSessionBindingContextDoesNotTrustHeadersWithoutTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, tc := range []struct {
@@ -24,7 +23,7 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
wantIP string
}{
{name: "trust disabled records proxy address", trustForwarded: false, wantIP: "127.0.0.1"},
{name: "trust enabled records forwarded client IP", trustForwarded: true, wantIP: "1.2.3.4"},
{name: "legacy trust toggle cannot bypass trusted proxies", trustForwarded: true, wantIP: "127.0.0.1"},
} {
t.Run(tc.name, func(t *testing.T) {
cfg := &config.Config{}
@@ -54,6 +53,23 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
}
}
func TestSessionBindingContextBoundsPersistedUserAgent(t *testing.T) {
cfg := &config.Config{}
r := gin.New()
r.Use(SessionBindingContext(cfg))
r.GET("/t", func(c *gin.Context) {
binding := service.SessionBindingFromContext(c.Request.Context())
require.Len(t, binding.UserAgent, maxPersistentUserAgentBytes)
require.Equal(t, binding.UserAgent, c.Request.UserAgent())
c.Status(200)
})
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/t", nil)
req.Header.Set("User-Agent", strings.Repeat("u", 2048))
r.ServeHTTP(w, req)
require.Equal(t, 200, w.Code)
}
// 未经过 SessionBindingContext 注入时(异常挂载顺序/单测直调),回退 trusted_proxies 链,
// 等价于开关关闭时的历史行为。
func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
@@ -75,8 +91,6 @@ func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
require.Equal(t, "9.9.9.9", w.Body.String())
}
// requestSessionBinding 优先取注入值:开关开启时校验哈希必须基于注入的转发 IP 计算,
// 与 token 签发路径取值一致,否则同一客户端会被误判为指纹变化。
func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -84,7 +98,7 @@ func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
cfg.SetTrustForwardedIPForAPIKeyACL(true)
r := gin.New()
require.NoError(t, r.SetTrustedProxies(nil))
require.NoError(t, r.SetTrustedProxies([]string{"127.0.0.1"}))
r.Use(SessionBindingContext(cfg))
r.GET("/t", func(c *gin.Context) {
issued := &service.SessionBinding{IP: "1.2.3.4", UserAgent: "test-agent"}
+2 -1
View File
@@ -35,6 +35,7 @@ func SetupRouter(
cfg *config.Config,
redisClient *redis.Client,
) *gin.Engine {
middleware2.SetIngressRejectRecorder(opsService)
// 缓存 iframe 页面的 origin 列表,用于动态注入 CSP frame-src
var cachedFrameOrigins atomic.Pointer[[]string]
emptyOrigins := []string{}
@@ -55,7 +56,7 @@ func SetupRouter(
// 应用中间件
r.Use(middleware2.RequestLogger())
// 将客户端 IP + UA 注入 request context,供 token 签发/会话绑定/审计日志统一读取。
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关。
// IP 取值与 API Key IP 限制共用 server.trusted_proxies 信任链。
r.Use(middleware2.SessionBindingContext(cfg))
r.Use(middleware2.Logger())
r.Use(middleware2.CORS(cfg.CORS))
+5
View File
@@ -235,6 +235,11 @@ func registerOpsRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
ops.GET("/request-errors/:id/upstream-errors", h.Admin.Ops.ListRequestErrorUpstreamErrors)
ops.PUT("/request-errors/:id/resolve", h.Admin.Ops.ResolveRequestError)
// Bounded ingress-admission rejection aggregates.
ops.GET("/ingress-rejections", h.Admin.Ops.ListIngressRejects)
ops.GET("/ingress-rejections/health", h.Admin.Ops.GetIngressRejectHealth)
ops.GET("/auth-cache-invalidation/health", h.Admin.Ops.GetAuthCacheInvalidationHealth)
// Upstream errors (independent upstream failures)
ops.GET("/upstream-errors", h.Admin.Ops.ListUpstreamErrors)
ops.GET("/upstream-errors/:id", h.Admin.Ops.GetUpstreamError)
+6 -5
View File
@@ -23,6 +23,7 @@ func RegisterGatewayRoutes(
cfg *config.Config,
) {
bodyLimit := middleware.RequestBodyLimit(cfg.Gateway.MaxBodySize)
textBodyLimit := middleware.RequestBodyLimit(cfg.Gateway.TextMaxBodySize)
clientRequestID := middleware.ClientRequestID()
opsErrorLogger := handler.OpsErrorLoggerMiddleware(opsService)
endpointNorm := handler.InboundEndpointMiddleware()
@@ -178,7 +179,7 @@ func RegisterGatewayRoutes(
}
h.Gateway.Responses(c)
})
gateway.POST("/alpha/search", h.OpenAIGateway.AlphaSearch)
gateway.POST("/alpha/search", textBodyLimit, h.OpenAIGateway.AlphaSearch)
gateway.GET("/responses", func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -190,7 +191,7 @@ func RegisterGatewayRoutes(
}
h.Gateway.ChatCompletions(c)
})
gateway.POST("/embeddings", func(c *gin.Context) {
gateway.POST("/embeddings", textBodyLimit, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
@@ -250,7 +251,7 @@ func RegisterGatewayRoutes(
}
r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/alpha/search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch)
r.POST("/alpha/search", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -260,7 +261,7 @@ func RegisterGatewayRoutes(
{
codexDirect.POST("/responses", responsesHandler)
codexDirect.POST("/responses/*subpath", responsesHandler)
codexDirect.POST("/alpha/search", h.OpenAIGateway.AlphaSearch)
codexDirect.POST("/alpha/search", textBodyLimit, h.OpenAIGateway.AlphaSearch)
codexDirect.GET("/responses", func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -274,7 +275,7 @@ func RegisterGatewayRoutes(
}
h.Gateway.ChatCompletions(c)
})
r.POST("/embeddings", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
r.POST("/embeddings", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
@@ -0,0 +1,53 @@
package routes
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/handler"
adminhandler "github.com/Wei-Shaw/sub2api/internal/handler/admin"
servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestIngressRejectAdminRoutesRequireAdminAuthentication(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
handlers := &handler.Handlers{Admin: &handler.AdminHandlers{Ops: adminhandler.NewOpsHandler(nil)}}
adminAuth := servermiddleware.AdminAuthMiddleware(func(c *gin.Context) {
if c.GetHeader("Authorization") == "" {
servermiddleware.AbortWithError(c, http.StatusUnauthorized, "UNAUTHORIZED", "Authorization required")
return
}
servermiddleware.AbortWithError(c, http.StatusForbidden, "FORBIDDEN", "Admin access required")
})
auditLog := servermiddleware.AuditLogMiddleware(func(c *gin.Context) { c.Next() })
stepUp := servermiddleware.StepUpAuthMiddleware(func(c *gin.Context) { c.Next() })
RegisterAdminRoutes(router.Group("/api/v1"), handlers, adminAuth, auditLog, stepUp, nil)
for _, path := range []string{
"/api/v1/admin/ops/ingress-rejections",
"/api/v1/admin/ops/ingress-rejections/health",
} {
for _, tc := range []struct {
name string
auth string
wantStatus int
}{
{name: "unauthenticated", wantStatus: http.StatusUnauthorized},
{name: "non-admin", auth: "Bearer user-token", wantStatus: http.StatusForbidden},
} {
t.Run(path+"/"+tc.name, func(t *testing.T) {
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, path, nil)
if tc.auth != "" {
request.Header.Set("Authorization", tc.auth)
}
router.ServeHTTP(recorder, request)
require.Equal(t, tc.wantStatus, recorder.Code)
})
}
}
}
+16 -6
View File
@@ -1423,8 +1423,12 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa
case OpenAIEndpointCapabilityChatCompletions:
return true
case OpenAIEndpointCapabilityGrokMediaGeneration:
eligible, _ := a.GrokMediaGenerationEligibility()
return eligible
eligible, reason := a.GrokMediaGenerationEligibility()
// Unobserved OAuth accounts remain scheduler candidates only so the
// request path can run the billing probe before forwarding. The
// forwarding gate itself fails closed if that probe is unavailable or
// cannot produce positive paid-entitlement evidence.
return eligible || reason == "billing_unobserved"
default:
return false
}
@@ -1469,9 +1473,9 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa
}
// GrokMediaGenerationEligibility reports whether a Grok account may receive
// new image/video generation requests. Missing observations preserve legacy
// routing; operators can fail closed for known-bad accounts with the explicit
// override. A successful override takes precedence over stale probe data.
// new image/video generation requests. OAuth media fails closed unless billing
// observations provide positive paid-entitlement evidence. An explicit
// operator override takes precedence over probe data.
func (a *Account) GrokMediaGenerationEligibility() (bool, string) {
if a == nil || !a.IsGrok() {
return false, "not_grok"
@@ -1488,11 +1492,17 @@ func (a *Account) GrokMediaGenerationEligibility() (bool, string) {
billing, err := grokBillingSnapshotFromExtra(a.Extra)
if err != nil || billing == nil {
return true, "billing_unobserved"
return false, "billing_unobserved"
}
if billing.StatusCode == 403 || billing.WeeklyStatusCode == 403 || billing.MonthlyStatusCode == 403 {
return false, "billing_forbidden"
}
if isKnownGrokFreeAccount(a) {
return false, "billing_free_tier"
}
if !grokBillingHasAuthoritativeQuota(billing) {
return false, "billing_inconclusive"
}
return true, "eligible"
}
@@ -13,6 +13,7 @@ import (
)
func TestGrokMediaGenerationEligibility(t *testing.T) {
weeklyUsagePercent := 12.5
forbiddenBilling := &xai.BillingSummary{
StatusCode: http.StatusForbidden,
WeeklyStatusCode: http.StatusForbidden,
@@ -20,9 +21,24 @@ func TestGrokMediaGenerationEligibility(t *testing.T) {
}
weeklyAllowance := &xai.BillingSummary{
PeriodType: "weekly",
UsagePercent: &weeklyUsagePercent,
StatusCode: http.StatusOK,
WeeklyStatusCode: http.StatusOK,
}
freeBilling := &xai.BillingSummary{
PeriodType: "monthly",
StatusCode: http.StatusOK,
WeeklyStatusCode: http.StatusOK,
MonthlyStatusCode: http.StatusOK,
MonthlyUpdatedAt: "2026-07-17T00:00:00Z",
}
inconclusiveBilling := &xai.BillingSummary{
StatusCode: http.StatusOK,
WeeklyStatusCode: http.StatusOK,
MonthlyStatusCode: http.StatusBadGateway,
Partial: true,
FailedWindows: []string{"monthly"},
}
weeklyForbidden := &xai.BillingSummary{
StatusCode: http.StatusOK,
WeeklyStatusCode: http.StatusForbidden,
@@ -43,12 +59,14 @@ func TestGrokMediaGenerationEligibility(t *testing.T) {
{name: "nil account", account: nil, want: false, wantReason: "not_grok"},
{name: "non grok account", account: &Account{Platform: PlatformOpenAI}, want: false, wantReason: "not_grok"},
{name: "non oauth grok account stays eligible", account: &Account{Platform: PlatformGrok, Type: AccountTypeAPIKey}, want: true, wantReason: "non_oauth"},
{name: "unobserved oauth preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: true, wantReason: "billing_unobserved"},
{name: "weekly allowance is not treated as weekly subscription", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"},
{name: "unobserved oauth fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: false, wantReason: "billing_unobserved"},
{name: "weekly paid usage is eligible without inferring from period type", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"},
{name: "observed free account is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: freeBilling}}, want: false, wantReason: "billing_free_tier"},
{name: "inconclusive billing fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: inconclusiveBilling}}, want: false, wantReason: "billing_inconclusive"},
{name: "billing forbidden is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: forbiddenBilling}}, want: false, wantReason: "billing_forbidden"},
{name: "weekly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyForbidden}}, want: false, wantReason: "billing_forbidden"},
{name: "monthly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: monthlyForbidden}}, want: false, wantReason: "billing_forbidden"},
{name: "malformed billing observation preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: true, wantReason: "billing_unobserved"},
{name: "malformed billing observation fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: false, wantReason: "billing_unobserved"},
{name: "malformed override falls back to observations", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: "false", grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"},
{name: "explicit disable wins", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: false}}, want: false, wantReason: "override_disabled"},
{name: "explicit enable wins over forbidden probe", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: true, grokBillingExtraKey: forbiddenBilling}}, want: true, wantReason: "override_enabled"},
@@ -63,6 +81,24 @@ func TestGrokMediaGenerationEligibility(t *testing.T) {
}
}
func TestGrokMediaCapabilityKeepsOnlyUnobservedOAuthAsProbeCandidate(t *testing.T) {
unobserved := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}
eligible, reason := unobserved.GrokMediaGenerationEligibility()
require.False(t, eligible)
require.Equal(t, "billing_unobserved", reason)
require.True(t, unobserved.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration))
inconclusive := &Account{
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Extra: map[string]any{grokBillingExtraKey: &xai.BillingSummary{
StatusCode: http.StatusOK,
Partial: true,
}},
}
require.False(t, inconclusive.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration))
}
func TestGrokMediaCapabilityFiltersOnlyGeneration(t *testing.T) {
account := &Account{
ID: 1,
@@ -91,6 +91,13 @@ type AccountRepository interface {
ListSchedulableByGroupIDAndPlatforms(ctx context.Context, groupID int64, platforms []string) ([]Account, error)
ListSchedulableUngroupedByPlatform(ctx context.Context, platform string) ([]Account, error)
ListSchedulableUngroupedByPlatforms(ctx context.Context, platforms []string) ([]Account, error)
// ListModelAvailabilityCandidates returns accounts that are enabled by
// persistent configuration (active + schedulable) for model-support
// diagnosis. It deliberately does not filter transient runtime state such
// as rate-limit, overload, temporary-unschedulable, or expiry windows.
// When groupID is nil, includeGrouped controls whether the query scans all
// matching accounts or only accounts without a group binding.
ListModelAvailabilityCandidates(ctx context.Context, groupID *int64, platforms []string, includeGrouped bool) ([]Account, error)
SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error
SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error
@@ -144,6 +151,7 @@ type AccountBulkUpdate struct {
Schedulable *bool
Credentials map[string]any
Extra map[string]any
ProbeEnabled *bool
}
// CreateAccountRequest 创建账号请求
@@ -159,6 +159,10 @@ func (s *accountRepoStub) ListSchedulableUngroupedByPlatforms(ctx context.Contex
panic("unexpected ListSchedulableUngroupedByPlatforms call")
}
func (s *accountRepoStub) ListModelAvailabilityCandidates(ctx context.Context, groupID *int64, platforms []string, includeGrouped bool) ([]Account, error) {
panic("unexpected ListModelAvailabilityCandidates call")
}
func (s *accountRepoStub) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
panic("unexpected SetRateLimited call")
}
+37 -4
View File
@@ -469,6 +469,15 @@ func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]an
Status: StatusActive,
Schedulable: true,
}
if input.ProbeEnabled != nil && *input.ProbeEnabled {
if !isUpstreamBillingProbeAccount(account) {
return nil, ErrUpstreamBillingProbeAccountInvalid
}
if account.Extra == nil {
account.Extra = make(map[string]any)
}
account.Extra[UpstreamBillingProbeEnabledExtraKey] = true
}
// 预计算固定时间重置的下次重置时间
if account.Extra != nil {
if err := ValidateQuotaResetConfig(account.Extra); err != nil {
@@ -842,7 +851,7 @@ func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, upd
// BulkUpdateAccounts updates multiple accounts in one request.
// It merges credentials/extra keys instead of overwriting the whole object.
func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUpdateAccountsInput) (*BulkUpdateAccountsResult, error) {
// Probe state is updated only through its dedicated endpoints.
// Managed probe state may only enter through the dedicated typed field below.
delete(input.Extra, UpstreamBillingProbeEnabledExtraKey)
delete(input.Extra, UpstreamBillingProbeExtraKey)
@@ -874,13 +883,30 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp
// 预取所有目标账号,供凭据守卫/代理守卫/混合渠道检查共用,避免多次 DB 查询。
var cachedTargets []*Account
if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate {
if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate || input.ProbeEnabled != nil {
loaded, err := s.accountRepo.GetByIDs(ctx, input.AccountIDs)
if err != nil {
return nil, err
}
cachedTargets = loaded
}
if input.ProbeEnabled != nil {
targetsByID := make(map[int64]*Account, len(cachedTargets))
for _, account := range cachedTargets {
if account != nil {
targetsByID[account.ID] = account
}
}
for _, accountID := range input.AccountIDs {
account, ok := targetsByID[accountID]
if !ok {
return nil, ErrAccountNotFound
}
if !isUpstreamBillingProbeAccount(account) {
return nil, ErrUpstreamBillingProbeAccountInvalid
}
}
}
if hasLongContextBillingUpdate {
for _, account := range cachedTargets {
if account == nil || account.Platform != PlatformOpenAI {
@@ -952,8 +978,15 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp
// Prepare bulk updates for columns and JSONB fields.
repoUpdates := AccountBulkUpdate{
Credentials: input.Credentials,
Extra: input.Extra,
Credentials: input.Credentials,
Extra: input.Extra,
ProbeEnabled: input.ProbeEnabled,
}
if input.ProbeEnabled != nil {
if repoUpdates.Extra == nil {
repoUpdates.Extra = make(map[string]any)
}
repoUpdates.Extra[UpstreamBillingProbeEnabledExtraKey] = *input.ProbeEnabled
}
if updatesUpstreamBillingProbeIdentity(input.Credentials) || input.ProxyID != nil {
if repoUpdates.Extra == nil {
@@ -38,6 +38,32 @@ func TestCreateAccountDropsManagedUpstreamBillingProbeState(t *testing.T) {
require.NotContains(t, created.Extra, UpstreamBillingProbeExtraKey)
}
func TestCreateAccountAcceptsDedicatedUpstreamBillingProbeSetting(t *testing.T) {
enabled := true
repo := &upstreamBillingProbeAccountRepo{}
created, err := (&adminServiceImpl{accountRepo: repo}).CreateAccount(context.Background(), &CreateAccountInput{
Name: "upstream",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Credentials: map[string]any{"api_key": "sk-test"},
ProbeEnabled: &enabled,
SkipDefaultGroupBind: true,
})
require.NoError(t, err)
require.Equal(t, true, created.Extra[UpstreamBillingProbeEnabledExtraKey])
_, err = (&adminServiceImpl{accountRepo: repo}).CreateAccount(context.Background(), &CreateAccountInput{
Name: "oauth",
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Credentials: map[string]any{"access_token": "token"},
ProbeEnabled: &enabled,
SkipDefaultGroupBind: true,
})
require.ErrorIs(t, err, ErrUpstreamBillingProbeAccountInvalid)
}
func TestUpdateAccountPreservesManagedUpstreamBillingProbeStateForUnrelatedEdit(t *testing.T) {
accountID := int64(110)
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
@@ -360,6 +386,63 @@ func TestBulkUpdateAccountsDropsManagedUpstreamBillingProbeState(t *testing.T) {
require.NotContains(t, repo.bulkUpdates[0].Extra, UpstreamBillingProbeExtraKey)
}
func TestBulkUpdateAccountsAcceptsDedicatedUpstreamBillingProbeSetting(t *testing.T) {
for _, enabled := range []bool{true, false} {
t.Run(map[bool]string{true: "enable", false: "disable"}[enabled], func(t *testing.T) {
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
1: {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
2: {ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
}}
result, err := (&adminServiceImpl{accountRepo: repo}).BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
AccountIDs: []int64{1, 2},
ProbeEnabled: &enabled,
})
require.NoError(t, err)
require.Equal(t, 2, result.Success)
require.Len(t, repo.bulkUpdates, 1)
require.Equal(t, enabled, repo.bulkUpdates[0].Extra[UpstreamBillingProbeEnabledExtraKey])
require.NotNil(t, repo.bulkUpdates[0].ProbeEnabled)
require.Equal(t, enabled, *repo.bulkUpdates[0].ProbeEnabled)
})
}
}
func TestBulkUpdateAccountsRejectsProbeSettingForIneligibleTargetBeforeWrite(t *testing.T) {
for _, enabled := range []bool{true, false} {
t.Run(map[bool]string{true: "enable", false: "disable"}[enabled], func(t *testing.T) {
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
1: {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
2: {ID: 2, Platform: PlatformOpenAI, Type: AccountTypeOAuth},
}}
_, err := (&adminServiceImpl{accountRepo: repo}).BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
AccountIDs: []int64{1, 2},
ProbeEnabled: &enabled,
})
require.ErrorIs(t, err, ErrUpstreamBillingProbeAccountInvalid)
require.Empty(t, repo.bulkUpdates)
})
}
}
func TestBulkUpdateAccountsRejectsProbeSettingWhenTargetIsMissing(t *testing.T) {
enabled := true
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
1: {ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
}}
_, err := (&adminServiceImpl{accountRepo: repo}).BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
AccountIDs: []int64{1, 2},
ProbeEnabled: &enabled,
})
require.ErrorIs(t, err, ErrAccountNotFound)
require.Empty(t, repo.bulkUpdates)
}
func TestBulkUpdateAccountsInvalidatesProbeSnapshotForIdentityCredentials(t *testing.T) {
repo := &upstreamBillingProbeAccountRepo{}
input := &BulkUpdateAccountsInput{
@@ -329,6 +329,7 @@ type CreateAccountInput struct {
GroupIDs []int64
ExpiresAt *int64
AutoPauseOnExpired *bool
ProbeEnabled *bool
// SkipDefaultGroupBind prevents auto-binding to platform default group when GroupIDs is empty.
SkipDefaultGroupBind bool
// SkipMixedChannelCheck skips the mixed channel risk check when binding groups.
@@ -378,6 +379,7 @@ type BulkUpdateAccountsInput struct {
GroupIDs *[]int64
Credentials map[string]any
Extra map[string]any
ProbeEnabled *bool
// SkipMixedChannelCheck skips the mixed channel risk check when binding groups.
// This should only be set when the caller has explicitly confirmed the risk.
SkipMixedChannelCheck bool

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