mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:08:02 +08:00
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:
+19
-5
@@ -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 \
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 的凭据材料")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user