fix: 过滤入口拒绝日志并强化鉴权边界

This commit is contained in:
benjamin
2026-07-18 00:11:18 +08:00
parent 57914967cb
commit b92bbf0299
101 changed files with 5872 additions and 828 deletions
@@ -0,0 +1,20 @@
# Ingress rejection log cleanup
This maintenance command removes historical admission rejections from
`ops_error_logs` without matching unrelated authentication or upstream errors.
It is a dry run unless `--execute` is supplied, and always requires an explicit
RFC3339 cutoff.
```sh
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z --execute
```
Run the execute form only after every application instance has been upgraded so
older instances cannot add new ingress rejection rows below the chosen cutoff.
The classifier intentionally retains invariant failures such as
`USER_NOT_FOUND`, database errors, quota/billing errors, and upstream failures.
After the rollout and cleanup are verified, run
`backend/scripts/finalize-ingress-reject-cleanup.sql` in a maintenance window to
remove the deprecated plaintext-key audit table and attribution columns.
@@ -0,0 +1,218 @@
package main
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"flag"
"fmt"
"log"
"sort"
"strings"
"time"
_ "github.com/Wei-Shaw/sub2api/ent/runtime"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/repository"
"github.com/lib/pq"
)
const classifierVersion = "ingress-reject-v1"
type candidate struct {
id int64
statusCode int
message string
body string
}
func main() {
beforeRaw := flag.String("before", "", "required RFC3339 cutoff; only older rows are considered")
execute := flag.Bool("execute", false, "delete matched rows (default is dry-run)")
batchSize := flag.Int("batch-size", 5000, "scan/delete batch size (1-5000)")
flag.Parse()
if *beforeRaw == "" {
log.Fatal("--before is required")
}
before, err := time.Parse(time.RFC3339, *beforeRaw)
if err != nil {
log.Fatalf("invalid --before: %v", err)
}
if *batchSize < 1 || *batchSize > 5000 {
log.Fatal("--batch-size must be between 1 and 5000")
}
cfg, err := config.LoadForBootstrap()
if err != nil {
log.Fatalf("load config: %v", err)
}
client, db, err := repository.InitEnt(cfg)
if err != nil {
log.Fatalf("initialize database: %v", err)
}
defer func() { _ = client.Close() }()
ctx := context.Background()
counts, scanned, matched, deleted, err := cleanup(ctx, db, before, *batchSize, *execute)
if err != nil {
log.Fatalf("cleanup failed: %v", err)
}
digest := sha256.Sum256([]byte(classifierVersion))
mode := "dry-run"
if *execute {
mode = "execute"
}
fmt.Printf("mode=%s before=%s classifier=%s scanned=%d matched=%d deleted=%d\n",
mode, before.UTC().Format(time.RFC3339), hex.EncodeToString(digest[:]), scanned, matched, deleted)
reasons := make([]string, 0, len(counts))
for reason := range counts {
reasons = append(reasons, reason)
}
sort.Strings(reasons)
for _, reason := range reasons {
fmt.Printf("reason=%s count=%d\n", reason, counts[reason])
}
if *execute && deleted > 0 {
fmt.Println("cleanup complete; schedule VACUUM (ANALYZE) ops_error_logs during normal maintenance")
}
}
func cleanup(ctx context.Context, db *sql.DB, before time.Time, batchSize int, execute bool) (map[string]int64, int64, int64, int64, error) {
counts := make(map[string]int64)
var cursor, scanned, matched, deleted int64
for {
rows, err := db.QueryContext(ctx, `
SELECT id, COALESCE(status_code, 0), COALESCE(error_message, ''), COALESCE(error_body, '')
FROM ops_error_logs
WHERE id > $1
AND created_at < $2
AND error_phase = 'auth'
AND account_id IS NULL
AND upstream_status_code IS NULL
AND COALESCE(upstream_error_message, '') = ''
AND COALESCE(upstream_error_detail, '') = ''
ORDER BY id ASC
LIMIT $3`, cursor, before, batchSize)
if err != nil {
return nil, scanned, matched, deleted, err
}
batch := make([]candidate, 0, batchSize)
for rows.Next() {
var item candidate
if err := rows.Scan(&item.id, &item.statusCode, &item.message, &item.body); err != nil {
_ = rows.Close()
return nil, scanned, matched, deleted, err
}
batch = append(batch, item)
cursor = item.id
}
if err := rows.Err(); err != nil {
_ = rows.Close()
return nil, scanned, matched, deleted, err
}
_ = rows.Close()
if len(batch) == 0 {
break
}
ids := make([]int64, 0, len(batch))
for _, item := range batch {
scanned++
if reason, ok := historicalIngressRejectReason(item); ok {
matched++
counts[reason]++
ids = append(ids, item.id)
}
}
if execute && len(ids) > 0 {
result, err := db.ExecContext(ctx,
`DELETE FROM ops_error_logs WHERE id = ANY($1) AND created_at < $2`, pq.Array(ids), before)
if err != nil {
return nil, scanned, matched, deleted, err
}
n, err := result.RowsAffected()
if err != nil {
return nil, scanned, matched, deleted, err
}
deleted += n
}
}
return counts, scanned, matched, deleted, nil
}
func historicalIngressRejectReason(item candidate) (string, bool) {
code, message := parseErrorIdentity(item.body, item.message)
switch code {
case "API_KEY_REQUIRED":
return "missing_key", true
case "INVALID_API_KEY":
return "invalid_key", true
case "API_KEY_DISABLED":
return "key_disabled", true
case "USER_INACTIVE":
return "user_inactive", true
case "GROUP_DELETED":
return "group_deleted", true
case "GROUP_DISABLED":
return "group_disabled", true
case "GROUP_NOT_ALLOWED":
return "group_forbidden", true
case "ACCESS_DENIED":
return "ip_acl_denied", true
case "api_key_in_query_deprecated":
return "query_key_deprecated", true
}
normalized := strings.TrimSpace(message)
switch {
case normalized == "API key is required":
return "missing_key", true
case normalized == "Invalid API key":
return "invalid_key", true
case normalized == "API key is disabled":
return "key_disabled", true
case normalized == "User account is not active":
return "user_inactive", true
case normalized == "API Key 所属分组已删除":
return "group_deleted", true
case normalized == "API Key 所属分组已停用":
return "group_disabled", true
case normalized == "API Key 所属专属分组不再允许当前用户使用":
return "group_forbidden", true
case normalized == "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.":
return "group_unassigned", true
case strings.HasPrefix(normalized, "Access denied. Your IP is "):
return "ip_acl_denied", true
case normalized == "Query parameter api_key is deprecated. Use Authorization header or key instead.":
return "query_key_deprecated", true
default:
return "", false
}
}
func parseErrorIdentity(body, fallbackMessage string) (string, string) {
var payload struct {
Code string `json:"code"`
Message string `json:"message"`
Error struct {
Code json.RawMessage `json:"code"`
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal([]byte(body), &payload); err != nil {
return "", fallbackMessage
}
message := payload.Message
if message == "" {
message = payload.Error.Message
}
if message == "" {
message = fallbackMessage
}
return strings.TrimSpace(payload.Code), message
}
@@ -0,0 +1,29 @@
package main
import "testing"
func TestHistoricalIngressRejectReason(t *testing.T) {
tests := []struct {
name string
item candidate
reason string
match bool
}{
{name: "standard invalid key", item: candidate{body: `{"code":"INVALID_API_KEY","message":"Invalid API key"}`}, reason: "invalid_key", match: true},
{name: "google missing key", item: candidate{body: `{"error":{"code":401,"message":"API key is required","status":"UNAUTHENTICATED"}}`}, reason: "missing_key", match: true},
{name: "google group deleted", item: candidate{body: `{"error":{"code":403,"message":"API Key 所属分组已删除","status":"PERMISSION_DENIED"}}`}, reason: "group_deleted", match: true},
{name: "ip acl", item: candidate{body: `{"code":"ACCESS_DENIED","message":"Access denied. Your IP is 192.0.2.1"}`}, reason: "ip_acl_denied", match: true},
{name: "user not found remains", item: candidate{body: `{"code":"USER_NOT_FOUND","message":"User associated with API key not found"}`}, match: false},
{name: "quota remains", item: candidate{body: `{"code":"API_KEY_QUOTA_EXHAUSTED","message":"quota"}`}, match: false},
{name: "database failure remains", item: candidate{statusCode: 500, message: "Failed to validate API key", body: `{"code":"INTERNAL_ERROR","message":"Failed to validate API key"}`}, match: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
reason, ok := historicalIngressRejectReason(tt.item)
if ok != tt.match || reason != tt.reason {
t.Fatalf("got (%q, %v), want (%q, %v)", reason, ok, tt.reason, tt.match)
}
})
}
}
+28
View File
@@ -81,6 +81,10 @@ func provideCleanup(
opsCleanup *service.OpsCleanupService,
opsScheduledReport *service.OpsScheduledReportService,
opsSystemLogSink *service.OpsSystemLogSink,
opsService *service.OpsService,
opsIngressReject *service.OpsIngressRejectAggregator,
apiKeyService *service.APIKeyService,
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
schedulerSnapshot *service.SchedulerSnapshotService,
tokenRefresh *service.TokenRefreshService,
accountExpiry *service.AccountExpiryService,
@@ -121,6 +125,30 @@ func provideCleanup(
// 应用层清理步骤可并行执行,基础设施资源(Redis/Ent)最后按顺序关闭。
parallelSteps := []cleanupStep{
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
}
return nil
}},
{"AuthCacheInvalidationWorker", func() error {
if authCacheInvalidationWorker != nil {
authCacheInvalidationWorker.Stop()
}
return nil
}},
{"AuthCacheInvalidationSubscriber", func() error {
if apiKeyService != nil {
apiKeyService.StopAuthCacheInvalidationSubscriber()
}
return nil
}},
{"OpsRuntimeSettingsRefresh", func() error {
if opsService != nil {
opsService.StopRuntimeSettingsRefresh()
}
return nil
}},
{"PromptAuditService", func() error {
if promptAudit != nil {
return promptAudit.Shutdown(ctx)
+33 -2
View File
@@ -157,7 +157,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
antigravityGatewayService := service.NewAntigravityGatewayService(accountRepository, gatewayCache, schedulerSnapshotService, antigravityTokenProvider, rateLimitService, httpUpstream, settingService, internal500CounterCache)
geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
authCacheInvalidationOutboxRepository := repository.NewAuthCacheInvalidationOutboxRepository(db)
authCacheInvalidationWorker := service.ProvideAuthCacheInvalidationWorker(authCacheInvalidationOutboxRepository, apiKeyCache, apiKeyService)
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService, authCacheInvalidationWorker, apiKeyService)
usageHandler := handler.NewUsageHandler(usageService, apiKeyService, opsService, settingService)
redeemHandler := handler.NewRedeemHandler(redeemService)
subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
@@ -306,6 +308,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
opsAlertEvaluatorService := service.ProvideOpsAlertEvaluatorService(opsService, opsRepository, emailService, redisClient, configConfig, proxyRepository)
opsCleanupService := service.ProvideOpsCleanupService(opsRepository, db, redisClient, configConfig, channelMonitorService, settingRepository, opsService)
opsScheduledReportService := service.ProvideOpsScheduledReportService(opsService, userService, emailService, redisClient, configConfig)
opsIngressRejectAggregator := service.ProvideOpsIngressRejectAggregator(opsRepository, opsService)
accountExpiryService := service.ProvideAccountExpiryService(accountRepository)
proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository)
subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db)
@@ -314,7 +317,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db)
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
application := &Application{
Server: httpServer,
PromptAudit: promptService,
@@ -351,6 +354,10 @@ func provideCleanup(
opsCleanup *service.OpsCleanupService,
opsScheduledReport *service.OpsScheduledReportService,
opsSystemLogSink *service.OpsSystemLogSink,
opsService *service.OpsService,
opsIngressReject *service.OpsIngressRejectAggregator,
apiKeyService *service.APIKeyService,
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
schedulerSnapshot *service.SchedulerSnapshotService,
tokenRefresh *service.TokenRefreshService,
accountExpiry *service.AccountExpiryService,
@@ -390,6 +397,30 @@ func provideCleanup(
}
parallelSteps := []cleanupStep{
{"OpsIngressRejectAggregator", func() error {
if opsIngressReject != nil {
opsIngressReject.Stop()
}
return nil
}},
{"AuthCacheInvalidationWorker", func() error {
if authCacheInvalidationWorker != nil {
authCacheInvalidationWorker.Stop()
}
return nil
}},
{"AuthCacheInvalidationSubscriber", func() error {
if apiKeyService != nil {
apiKeyService.StopAuthCacheInvalidationSubscriber()
}
return nil
}},
{"OpsRuntimeSettingsRefresh", func() error {
if opsService != nil {
opsService.StopRuntimeSettingsRefresh()
}
return nil
}},
{"PromptAuditService", func() error {
if promptAudit != nil {
return promptAudit.Shutdown(ctx)
+4
View File
@@ -58,6 +58,10 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
&service.OpsCleanupService{},
&service.OpsScheduledReportService{},
opsSystemLogSinkSvc,
nil, // opsService
nil, // opsIngressRejectAggregator
nil, // apiKeyService
nil, // authCacheInvalidationWorker
schedulerSnapshotSvc,
tokenRefreshSvc,
accountExpirySvc,
+75 -8
View File
@@ -645,6 +645,7 @@ type ServerConfig struct {
EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header
FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接
ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒)
MaxHeaderBytes int `mapstructure:"max_header_bytes"` // 请求头最大字节数(HTTP/2 映射为 header-list 上限)
IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒)
TrustedProxies []string `mapstructure:"trusted_proxies"` // 可信代理列表(CIDR/IP)
MaxRequestBodySize int64 `mapstructure:"max_request_body_size"` // 全局最大请求体限制
@@ -796,6 +797,8 @@ type GatewayConfig struct {
OpenAIHighEffortFirstOutputTimeoutSeconds int `mapstructure:"openai_high_effort_first_output_timeout_seconds"`
// 请求体最大字节数,用于网关请求体大小限制
MaxBodySize int64 `mapstructure:"max_body_size"`
// TextMaxBodySize limits endpoints that cannot carry inline image/video payloads.
TextMaxBodySize int64 `mapstructure:"text_max_body_size"`
// 非流式上游响应体读取上限(字节),用于防止无界读取导致内存放大
UpstreamResponseReadMaxBytes int64 `mapstructure:"upstream_response_read_max_bytes"`
// 代理探测响应体读取上限(字节)
@@ -1419,12 +1422,22 @@ type RateLimitConfig struct {
// APIKeyAuthCacheConfig API Key 认证缓存配置
type APIKeyAuthCacheConfig struct {
L1Size int `mapstructure:"l1_size"`
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
JitterPercent int `mapstructure:"jitter_percent"`
Singleflight bool `mapstructure:"singleflight"`
L1Size int `mapstructure:"l1_size"`
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
JitterPercent int `mapstructure:"jitter_percent"`
Singleflight bool `mapstructure:"singleflight"`
LookupConcurrency int `mapstructure:"lookup_concurrency"`
InvalidAbuse InvalidAuthAbuseConfig `mapstructure:"invalid_abuse"`
}
type InvalidAuthAbuseConfig struct {
Enabled bool `mapstructure:"enabled"`
Threshold int `mapstructure:"threshold"`
WindowSeconds int `mapstructure:"window_seconds"`
BlockSeconds int `mapstructure:"block_seconds"`
Capacity int `mapstructure:"capacity"`
}
// SubscriptionCacheConfig 订阅认证 L1 缓存配置
@@ -1698,8 +1711,9 @@ func setDefaults() {
viper.SetDefault("server.mode", "release")
viper.SetDefault("server.enable_server_timing", false)
viper.SetDefault("server.frontend_url", "")
viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
viper.SetDefault("server.read_header_timeout", 10) // 10秒读取请求头
viper.SetDefault("server.max_header_bytes", 64*1024)
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
viper.SetDefault("server.trusted_proxies", []string{})
viper.SetDefault("server.max_request_body_size", int64(256*1024*1024))
// H2C 默认配置
@@ -1983,6 +1997,12 @@ func setDefaults() {
viper.SetDefault("api_key_auth_cache.negative_ttl_seconds", 30)
viper.SetDefault("api_key_auth_cache.jitter_percent", 10)
viper.SetDefault("api_key_auth_cache.singleflight", true)
viper.SetDefault("api_key_auth_cache.lookup_concurrency", 64)
viper.SetDefault("api_key_auth_cache.invalid_abuse.enabled", true)
viper.SetDefault("api_key_auth_cache.invalid_abuse.threshold", 120)
viper.SetDefault("api_key_auth_cache.invalid_abuse.window_seconds", 60)
viper.SetDefault("api_key_auth_cache.invalid_abuse.block_seconds", 60)
viper.SetDefault("api_key_auth_cache.invalid_abuse.capacity", 16384)
// Subscription auth L1 cache
viper.SetDefault("subscription_cache.l1_size", 16384)
@@ -2111,6 +2131,7 @@ func setDefaults() {
viper.SetDefault("gateway.antigravity_fallback_cooldown_minutes", 1)
viper.SetDefault("gateway.antigravity_extra_retries", 10)
viper.SetDefault("gateway.max_body_size", int64(256*1024*1024))
viper.SetDefault("gateway.text_max_body_size", int64(32*1024*1024))
viper.SetDefault("gateway.upstream_response_read_max_bytes", DefaultUpstreamResponseReadMaxBytes)
viper.SetDefault("gateway.proxy_probe_response_read_max_bytes", int64(1024*1024))
viper.SetDefault("gateway.gemini_debug_response_headers", false)
@@ -2208,6 +2229,49 @@ func setDefaults() {
}
func (c *Config) Validate() error {
if c.Server.ReadHeaderTimeout < 1 || c.Server.ReadHeaderTimeout > 60 {
return fmt.Errorf("server.read_header_timeout must be between 1 and 60 seconds")
}
if c.Server.MaxHeaderBytes < 8*1024 || c.Server.MaxHeaderBytes > 1024*1024 {
return fmt.Errorf("server.max_header_bytes must be between 8192 and 1048576 bytes")
}
if c.Server.IdleTimeout <= 0 {
return fmt.Errorf("server.idle_timeout must be positive")
}
if c.Server.MaxRequestBodySize < 0 {
return fmt.Errorf("server.max_request_body_size must be non-negative")
}
if c.Server.H2C.Enabled {
if c.Server.H2C.MaxConcurrentStreams == 0 {
return fmt.Errorf("server.h2c.max_concurrent_streams must be positive")
}
if c.Server.H2C.IdleTimeout <= 0 {
return fmt.Errorf("server.h2c.idle_timeout must be positive")
}
if c.Server.H2C.MaxReadFrameSize < 16*1024 || c.Server.H2C.MaxReadFrameSize > 16*1024*1024-1 {
return fmt.Errorf("server.h2c.max_read_frame_size must be between 16384 and 16777215 bytes")
}
if c.Server.H2C.MaxUploadBufferPerConnection < 65535 {
return fmt.Errorf("server.h2c.max_upload_buffer_per_connection must be at least 65535 bytes")
}
if c.Server.H2C.MaxUploadBufferPerStream <= 0 {
return fmt.Errorf("server.h2c.max_upload_buffer_per_stream must be positive")
}
}
if c.APIKeyAuth.InvalidAbuse.Enabled {
if c.APIKeyAuth.InvalidAbuse.Threshold < 10 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.threshold must be at least 10")
}
if c.APIKeyAuth.InvalidAbuse.WindowSeconds < 1 || c.APIKeyAuth.InvalidAbuse.WindowSeconds > 3600 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.window_seconds must be between 1 and 3600")
}
if c.APIKeyAuth.InvalidAbuse.BlockSeconds < 1 || c.APIKeyAuth.InvalidAbuse.BlockSeconds > 3600 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.block_seconds must be between 1 and 3600")
}
if c.APIKeyAuth.InvalidAbuse.Capacity < 256 || c.APIKeyAuth.InvalidAbuse.Capacity > 1_000_000 {
return fmt.Errorf("api_key_auth_cache.invalid_abuse.capacity must be between 256 and 1000000")
}
}
jwtSecret := strings.TrimSpace(c.JWT.Secret)
if jwtSecret == "" {
return fmt.Errorf("jwt.secret is required")
@@ -2733,6 +2797,9 @@ func (c *Config) Validate() error {
if c.Gateway.MaxBodySize <= 0 {
return fmt.Errorf("gateway.max_body_size must be positive")
}
if c.Gateway.TextMaxBodySize <= 0 || c.Gateway.TextMaxBodySize > c.Gateway.MaxBodySize {
return fmt.Errorf("gateway.text_max_body_size must be positive and no greater than gateway.max_body_size")
}
if c.Gateway.UpstreamResponseReadMaxBytes <= 0 {
return fmt.Errorf("gateway.upstream_response_read_max_bytes must be positive")
}
+62
View File
@@ -35,6 +35,18 @@ func TestLoadServerTimingConfig(t *testing.T) {
})
}
func TestLoadHTTPIngressSafetyDefaults(t *testing.T) {
resetViperWithJWTSecret(t)
cfg, err := Load()
require.NoError(t, err)
require.Equal(t, 10, cfg.Server.ReadHeaderTimeout)
require.Equal(t, 64*1024, cfg.Server.MaxHeaderBytes)
require.Equal(t, int64(32*1024*1024), cfg.Gateway.TextMaxBodySize)
require.True(t, cfg.APIKeyAuth.InvalidAbuse.Enabled)
require.Equal(t, 120, cfg.APIKeyAuth.InvalidAbuse.Threshold)
require.Equal(t, 16384, cfg.APIKeyAuth.InvalidAbuse.Capacity)
}
func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) {
viper.Reset()
t.Setenv("JWT_SECRET", "")
@@ -1185,6 +1197,51 @@ func TestValidateConfigErrors(t *testing.T) {
mutate func(*Config)
wantErr string
}{
{
name: "server read header timeout",
mutate: func(c *Config) { c.Server.ReadHeaderTimeout = 0 },
wantErr: "server.read_header_timeout",
},
{
name: "server max header bytes too small",
mutate: func(c *Config) { c.Server.MaxHeaderBytes = 4096 },
wantErr: "server.max_header_bytes",
},
{
name: "server max request body size",
mutate: func(c *Config) { c.Server.MaxRequestBodySize = -1 },
wantErr: "server.max_request_body_size",
},
{
name: "h2c zero concurrent streams",
mutate: func(c *Config) {
c.Server.H2C.Enabled = true
c.Server.H2C.MaxConcurrentStreams = 0
},
wantErr: "server.h2c.max_concurrent_streams",
},
{
name: "h2c oversized read frame",
mutate: func(c *Config) {
c.Server.H2C.Enabled = true
c.Server.H2C.MaxReadFrameSize = 16 * 1024 * 1024
},
wantErr: "server.h2c.max_read_frame_size",
},
{
name: "invalid auth abuse threshold too small",
mutate: func(c *Config) {
c.APIKeyAuth.InvalidAbuse.Threshold = 9
},
wantErr: "api_key_auth_cache.invalid_abuse.threshold",
},
{
name: "invalid auth abuse capacity too small",
mutate: func(c *Config) {
c.APIKeyAuth.InvalidAbuse.Capacity = 255
},
wantErr: "api_key_auth_cache.invalid_abuse.capacity",
},
{
name: "jwt secret required",
mutate: func(c *Config) { c.JWT.Secret = "" },
@@ -1386,6 +1443,11 @@ func TestValidateConfigErrors(t *testing.T) {
mutate: func(c *Config) { c.Gateway.MaxBodySize = 0 },
wantErr: "gateway.max_body_size",
},
{
name: "gateway text body exceeds media body",
mutate: func(c *Config) { c.Gateway.TextMaxBodySize = c.Gateway.MaxBodySize + 1 },
wantErr: "gateway.text_max_body_size",
},
{
name: "gateway response header timeout",
mutate: func(c *Config) { c.Gateway.ResponseHeaderTimeout = -1 },
@@ -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`)
}
@@ -1622,6 +1622,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 {
+130 -75
View File
@@ -72,24 +72,10 @@ const (
opsErrorLogMinQueueSize = 256
opsErrorLogMaxQueueSize = 8192
opsErrorLogBatchSize = 32
opsErrorLogMaxQueueBytes = 32 * 1024 * 1024
opsErrorLogMaxUserAgentBytes = 512
)
// looksLikeSystemKey 粗筛"形似本系统 key"的输入:长度 16-128 且仅含 [a-zA-Z0-9_-]。
// 不用前缀匹配(APIKeyPrefix 可配置)。用于反查审计表前挡掉随机扫描的乱码输入。
func looksLikeSystemKey(key string) bool {
if len(key) < 16 || len(key) > 128 {
return false
}
for _, c := range key {
allowed := (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
(c >= '0' && c <= '9') || c == '_' || c == '-'
if !allowed {
return false
}
}
return true
}
// keyPrefix 返回脱敏前缀(前 n 个字符);不足 n 则原样返回。
func keyPrefix(key string, n int) string {
if len(key) <= n {
@@ -98,43 +84,26 @@ func keyPrefix(key string, n int) string {
return key[:n]
}
// extractAttemptedKey 按认证中间件同样的顺序从请求头提取提交的 key 明文。
// 与 api_key_auth.go:43-59 一致:Authorization 仅取 Bearer scheme,非 Bearer 则忽略并继续 x-api-key → x-goog-api-key。
func extractAttemptedKey(c *gin.Context) string {
if h := c.GetHeader("Authorization"); h != "" {
parts := strings.SplitN(h, " ", 2)
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
return strings.TrimSpace(parts[1])
}
// 非 Bearer:与中间件一致,忽略 Authorization,继续尝试其它 header(不在此 return)。
}
if k := c.GetHeader("x-api-key"); k != "" {
return strings.TrimSpace(k)
}
if k := c.GetHeader("x-goog-api-key"); k != "" {
return strings.TrimSpace(k)
}
return ""
}
type opsErrorLogJob struct {
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
queuedBytes int64
}
var (
opsErrorLogOnce sync.Once
opsErrorLogQueue chan opsErrorLogJob
opsErrorLogStopOnce sync.Once
opsErrorLogWorkersWg sync.WaitGroup
opsErrorLogMu sync.RWMutex
opsErrorLogStopping bool
opsErrorLogQueueLen atomic.Int64
opsErrorLogEnqueued atomic.Int64
opsErrorLogDropped atomic.Int64
opsErrorLogProcessed atomic.Int64
opsErrorLogSanitized atomic.Int64
opsErrorLogStopOnce sync.Once
opsErrorLogWorkersWg sync.WaitGroup
opsErrorLogMu sync.RWMutex
opsErrorLogStopping bool
opsErrorLogQueueLen atomic.Int64
opsErrorLogQueueBytes atomic.Int64
opsErrorLogEnqueued atomic.Int64
opsErrorLogDropped atomic.Int64
opsErrorLogProcessed atomic.Int64
opsErrorLogSanitized atomic.Int64
opsErrorLogLastDropLogAt atomic.Int64
@@ -154,6 +123,7 @@ func startOpsErrorLogWorkers() {
workerCount, queueSize := opsErrorLogConfig()
opsErrorLogQueue = make(chan opsErrorLogJob, queueSize)
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
opsErrorLogWorkersWg.Add(workerCount)
for i := 0; i < workerCount; i++ {
@@ -165,6 +135,7 @@ func startOpsErrorLogWorkers() {
return
}
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-job.queuedBytes)
batch := make([]opsErrorLogJob, 0, opsErrorLogBatchSize)
batch = append(batch, job)
@@ -184,6 +155,7 @@ func startOpsErrorLogWorkers() {
return
}
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-nextJob.queuedBytes)
batch = append(batch, nextJob)
case <-timer.C:
break batchLoop
@@ -239,6 +211,20 @@ func enqueueOpsErrorLog(ops *service.OpsService, entry *service.OpsInsertErrorLo
if ops == nil || entry == nil {
return
}
entry.UserAgent = normalizeOpsPersistentUserAgent(entry.UserAgent)
if entry.ErrorBody != "" {
originalBody := entry.ErrorBody
body, truncated := service.SanitizeOpsErrorBodyForQueue(originalBody)
entry.ErrorBody = body
if truncated || body != originalBody {
opsErrorLogSanitized.Add(1)
}
}
if err := service.SanitizeOpsUpstreamErrorsForQueue(entry); err != nil {
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
return
}
select {
case <-opsErrorLogShutdownCh:
return
@@ -259,18 +245,29 @@ func enqueueOpsErrorLog(ops *service.OpsService, entry *service.OpsInsertErrorLo
if opsErrorLogStopping || opsErrorLogQueue == nil {
return
}
queuedBytes := estimateOpsErrorLogJobBytes(entry)
if !reserveOpsErrorLogQueueBytes(queuedBytes) {
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
return
}
select {
case opsErrorLogQueue <- opsErrorLogJob{ops: ops, entry: entry}:
opsErrorLogQueueLen.Add(1)
case opsErrorLogQueue <- opsErrorLogJob{ops: ops, entry: entry, queuedBytes: queuedBytes}:
opsErrorLogEnqueued.Add(1)
default:
opsErrorLogQueueLen.Add(-1)
opsErrorLogQueueBytes.Add(-queuedBytes)
// Queue is full; drop to avoid blocking request handling.
opsErrorLogDropped.Add(1)
maybeLogOpsErrorLogDrop()
}
}
func normalizeOpsPersistentUserAgent(value string) string {
return truncateString(strings.TrimSpace(strings.ToValidUTF8(value, "")), opsErrorLogMaxUserAgentBytes)
}
func StopOpsErrorLogWorkers() bool {
opsErrorLogStopOnce.Do(func() {
opsErrorLogShutdownOnce.Do(func() {
@@ -293,6 +290,7 @@ func stopOpsErrorLogWorkers() bool {
if ch == nil {
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
return true
}
@@ -305,6 +303,7 @@ func stopOpsErrorLogWorkers() bool {
select {
case <-done:
opsErrorLogQueueLen.Store(0)
opsErrorLogQueueBytes.Store(0)
return true
case <-time.After(opsErrorLogDrainTimeout):
return false
@@ -315,6 +314,14 @@ func OpsErrorLogQueueLength() int64 {
return opsErrorLogQueueLen.Load()
}
func OpsErrorLogQueueBytes() int64 {
return opsErrorLogQueueBytes.Load()
}
func OpsErrorLogQueueBytesCapacity() int64 {
return opsErrorLogMaxQueueBytes
}
func OpsErrorLogQueueCapacity() int {
opsErrorLogMu.RLock()
ch := opsErrorLogQueue
@@ -355,12 +362,15 @@ func maybeLogOpsErrorLogDrop() {
}
queued := opsErrorLogQueueLen.Load()
queuedBytes := opsErrorLogQueueBytes.Load()
queueCap := OpsErrorLogQueueCapacity()
log.Printf(
"[OpsErrorLogger] queue is full; dropping logs (queued=%d cap=%d enqueued_total=%d dropped_total=%d processed_total=%d sanitized_total=%d)",
"[OpsErrorLogger] queue is full; dropping logs (queued=%d cap=%d queued_bytes=%d bytes_cap=%d enqueued_total=%d dropped_total=%d processed_total=%d sanitized_total=%d)",
queued,
queueCap,
queuedBytes,
opsErrorLogMaxQueueBytes,
opsErrorLogEnqueued.Load(),
opsErrorLogDropped.Load(),
opsErrorLogProcessed.Load(),
@@ -368,6 +378,46 @@ func maybeLogOpsErrorLogDrop() {
)
}
func reserveOpsErrorLogQueueBytes(size int64) bool {
if size < 1 {
size = 1
}
for {
current := opsErrorLogQueueBytes.Load()
if current > opsErrorLogMaxQueueBytes-size {
return false
}
if opsErrorLogQueueBytes.CompareAndSwap(current, current+size) {
opsErrorLogQueueLen.Add(1)
return true
}
}
}
func estimateOpsErrorLogJobBytes(entry *service.OpsInsertErrorLogInput) int64 {
if entry == nil {
return 1
}
const fixedOverhead = 512
size := fixedOverhead + len(entry.RequestID) + len(entry.ClientRequestID) +
len(entry.Platform) + len(entry.Model) + len(entry.RequestPath) +
len(entry.InboundEndpoint) + len(entry.UpstreamEndpoint) +
len(entry.RequestedModel) + len(entry.UpstreamModel) + len(entry.UserAgent) +
len(entry.ErrorPhase) + len(entry.ErrorType) + len(entry.Severity) +
len(entry.ErrorMessage) + len(entry.ErrorBody) + len(entry.ErrorSource) +
len(entry.ErrorOwner) + len(entry.APIKeyPrefix)
if entry.UpstreamErrorMessage != nil {
size += len(*entry.UpstreamErrorMessage)
}
if entry.UpstreamErrorDetail != nil {
size += len(*entry.UpstreamErrorDetail)
}
if entry.UpstreamErrorsJSON != nil {
size += len(*entry.UpstreamErrorsJSON)
}
return int64(size)
}
func opsErrorLogConfig() (workerCount int, queueSize int) {
workerCount = runtime.GOMAXPROCS(0) * 2
if workerCount < opsErrorLogMinWorkerCount {
@@ -470,9 +520,12 @@ type opsCaptureWriter struct {
gin.ResponseWriter
limit int
buf bytes.Buffer
ctx *gin.Context
}
const opsCaptureWriterLimit = 64 * 1024
const opsCaptureWriterLimit = service.OpsErrorLogQueueBodyMaxBytes
const opsCaptureWriterPoolMaxRetainedCapacity = service.OpsErrorLogQueueBodyMaxBytes
var opsCaptureWriterPool = sync.Pool{
New: func() any {
@@ -496,11 +549,19 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) {
return
}
w.ResponseWriter = nil
w.ctx = nil
w.limit = opsCaptureWriterLimit
if !shouldPoolOpsCaptureWriter(w) {
return
}
w.buf.Reset()
opsCaptureWriterPool.Put(w)
}
func shouldPoolOpsCaptureWriter(w *opsCaptureWriter) bool {
return w != nil && w.buf.Cap() <= opsCaptureWriterPoolMaxRetainedCapacity
}
func (w *opsCaptureWriter) Status() int {
if w.ResponseWriter == nil {
return 0
@@ -577,7 +638,7 @@ func (w *opsCaptureWriter) Write(b []byte) (int, error) {
if w.ResponseWriter == nil {
return 0, nil
}
if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
remaining := w.limit - w.buf.Len()
if len(b) > remaining {
_, _ = w.buf.Write(b[:remaining])
@@ -592,7 +653,7 @@ func (w *opsCaptureWriter) WriteString(s string) (int, error) {
if w.ResponseWriter == nil {
return 0, nil
}
if w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit {
remaining := w.limit - w.buf.Len()
if len(s) > remaining {
_, _ = w.buf.WriteString(s[:remaining])
@@ -603,6 +664,14 @@ func (w *opsCaptureWriter) WriteString(s string) (int, error) {
return w.ResponseWriter.WriteString(s)
}
func (w *opsCaptureWriter) shouldCapture() bool {
if w.ctx == nil {
return true
}
_, rejected := middleware2.GetIngressRejectReason(w.ctx)
return !rejected
}
// OpsErrorLoggerMiddleware records error responses (status >= 400) into ops_error_logs.
//
// Notes:
@@ -612,6 +681,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return func(c *gin.Context) {
originalWriter := c.Writer
w := acquireOpsCaptureWriter(originalWriter)
w.ctx = c
defer func() {
// Restore the original writer before returning so outer middlewares
// don't observe a pooled wrapper that has been released.
@@ -623,6 +693,10 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
c.Writer = w
c.Next()
if _, rejected := middleware2.GetIngressRejectReason(c); rejected {
return
}
if ops == nil {
return
}
@@ -989,7 +1063,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
IsCountTokens: isCountTokensRequest(c),
ErrorMessage: parsed.Message,
// Keep the full captured error body (capture is already capped at 64KB) so the
// Keep the captured error body (already capped at the queue-safe limit) so the
// service layer can sanitize JSON before truncating for storage.
ErrorBody: string(body),
ErrorSource: errorSource,
@@ -1002,7 +1076,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
// 有效(未删除)key 报错时快照前缀,key 之后被删也保留;与 INVALID_API_KEY 的 attempted_key_prefix 互斥。
// 有效 key 报错时快照前缀,key 之后被删也保留。
entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
@@ -1022,22 +1096,6 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
entry.ClientIP = &clientIP
}
// 已删除 key 归因:仅 INVALID_API_KEY 才尝试。响应已写出,此处不阻塞客户端。
if parsed.Code == opsCodeInvalidAPIKey {
if attemptedKey := extractAttemptedKey(c); attemptedKey != "" {
entry.AttemptedKeyPrefix = keyPrefix(attemptedKey, 8)
if looksLikeSystemKey(attemptedKey) {
if res, lookupErr := ops.LookupDeletedKeyAudit(c.Request.Context(), attemptedKey); lookupErr != nil {
log.Printf("[OpsErrorLogger] LookupDeletedKeyAudit failed: %v", lookupErr)
} else if res != nil {
owner := res.UserID
entry.DeletedKeyOwnerUserID = &owner
entry.DeletedKeyName = res.KeyName
}
}
}
}
enqueueOpsErrorLog(ops, entry)
}
}
@@ -1694,11 +1752,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。
@@ -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())
}
+7 -41
View File
@@ -8,40 +8,10 @@ import (
"github.com/gin-gonic/gin"
)
// GetClientIP 从 Gin Context 中提取客户端真实 IP 地址。
// 按以下优先级检查 Header:
// 1. CF-Connecting-IP (Cloudflare)
// 2. X-Real-IP (Nginx)
// 3. X-Forwarded-For (取第一个非私有 IP)
// 4. c.ClientIP() (Gin 内置方法)
// GetClientIP resolves a client address only through Gin's configured trusted
// proxy chain. Forwarding headers from a direct or untrusted peer are ignored.
func GetClientIP(c *gin.Context) string {
// 1. Cloudflare
if ip := c.GetHeader("CF-Connecting-IP"); ip != "" {
return normalizeIP(ip)
}
// 2. Nginx X-Real-IP
if ip := c.GetHeader("X-Real-IP"); ip != "" {
return normalizeIP(ip)
}
// 3. X-Forwarded-For (多个 IP 时取第一个公网 IP)
if xff := c.GetHeader("X-Forwarded-For"); xff != "" {
ips := strings.Split(xff, ",")
for _, ip := range ips {
ip = strings.TrimSpace(ip)
if ip != "" && !isPrivateIP(ip) {
return normalizeIP(ip)
}
}
// 如果都是私有 IP,返回第一个
if len(ips) > 0 {
return normalizeIP(strings.TrimSpace(ips[0]))
}
}
// 4. Gin 内置方法
return normalizeIP(c.ClientIP())
return GetTrustedClientIP(c)
}
// GetTrustedClientIP 从 Gin 的可信代理解析链提取客户端 IP。
@@ -54,14 +24,10 @@ func GetTrustedClientIP(c *gin.Context) string {
return normalizeIP(c.ClientIP())
}
// GetSecurityClientIP 返回安全敏感场景(API Key IP 限制、审计日志、会话 IP/UA 绑定)
// 使用的客户端 IP。trustForwarded 对应系统设置「信任反代传递的客户端 IP」:
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
// 关闭时走 Gin trusted_proxies 解析链。
func GetSecurityClientIP(c *gin.Context, trustForwarded bool) string {
if trustForwarded {
return GetClientIP(c)
}
// GetSecurityClientIP returns the address resolved through Gin's configured
// trusted-proxy chain. The legacy toggle is retained for configuration/API
// compatibility, but never makes raw forwarding headers trustworthy by itself.
func GetSecurityClientIP(c *gin.Context, _ bool) string {
return GetTrustedClientIP(c)
}
+18 -3
View File
@@ -95,7 +95,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 +103,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 +124,18 @@ func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
})
}
}
func TestGetSecurityClientIPUsesConfiguredTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
r := gin.New()
require.NoError(t, r.SetTrustedProxies([]string{"9.9.9.9"}))
r.GET("/t", func(c *gin.Context) { c.String(200, GetSecurityClientIP(c, true)) })
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/t", nil)
req.RemoteAddr = "9.9.9.9:12345"
req.Header.Set("X-Forwarded-For", "1.2.3.4")
r.ServeHTTP(w, req)
require.Equal(t, "1.2.3.4", w.Body.String())
}
+4
View File
@@ -26,6 +26,10 @@ const (
LevelWarn = zapcore.WarnLevel
LevelError = zapcore.ErrorLevel
LevelFatal = zapcore.FatalLevel
// OpsSystemLogSkipField keeps an event in the standard logger while
// preventing the database-backed Ops system-log sink from indexing it.
OpsSystemLogSkipField = "ops_system_log_skip"
)
type Sink interface {
+18 -21
View File
@@ -110,28 +110,25 @@ func (c *apiKeyCache) SubscribeAuthCacheInvalidation(ctx context.Context, handle
return fmt.Errorf("subscribe to auth cache invalidation: %w", err)
}
go func() {
defer func() {
if err := pubsub.Close(); err != nil {
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
}
}()
ch := pubsub.Channel()
for {
select {
case <-ctx.Done():
return
case msg, ok := <-ch:
if !ok {
return
}
if msg != nil {
handler(msg.Payload)
}
}
defer func() {
if err := pubsub.Close(); err != nil {
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
}
}()
service.NotifyAuthCacheSubscriptionReady(ctx)
return nil
ch := pubsub.Channel()
for {
select {
case <-ctx.Done():
return ctx.Err()
case msg, ok := <-ch:
if !ok {
return errors.New("auth cache invalidation pubsub channel closed")
}
if msg != nil {
handler(msg.Payload)
}
}
}
}
@@ -0,0 +1,49 @@
package repository
import (
"context"
"errors"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func TestAPIKeyCacheSubscriber_BlocksUntilContextCancellation(t *testing.T) {
server := miniredis.RunT(t)
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
defer func() { _ = client.Close() }()
cache := NewAPIKeyCache(client)
ctx, cancel := context.WithCancel(context.Background())
received := make(chan string, 1)
returned := make(chan error, 1)
go func() {
returned <- cache.SubscribeAuthCacheInvalidation(ctx, func(value string) { received <- value })
}()
var value string
require.Eventually(t, func() bool {
require.NoError(t, client.Publish(context.Background(), authCacheInvalidateChannel, "hash").Err())
select {
case value = <-received:
return true
default:
return false
}
}, time.Second, 10*time.Millisecond)
require.Equal(t, "hash", value)
select {
case err := <-returned:
t.Fatalf("subscriber returned while connection was active: %v", err)
default:
}
cancel()
select {
case err := <-returned:
require.True(t, errors.Is(err, context.Canceled))
case <-time.After(time.Second):
t.Fatal("subscriber did not stop after context cancellation")
}
}
+6 -18
View File
@@ -326,16 +326,14 @@ func (r *apiKeyRepository) Delete(ctx context.Context, id int64) error {
return nil
}
// DeleteWithAudit 在同一事务内:
// 1. 把(明文 key、所有者、key 名称)写入 deleted_api_key_audits;
// 2. 软删除该 key(tombstone 覆盖 key 列以释放唯一约束)。
//
// 保证"被删除的 key 一定能反查到所有者"。事务模式与 group_repo.DeleteCascade 一致。
// DeleteWithAudit keeps the legacy method name for rolling-upgrade compatibility.
// It atomically tombstones and soft-deletes the key without retaining credential
// material. Tombstoning releases the unique key value for safe reuse.
func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error {
tombstoneKey := fmt.Sprintf("__deleted__%d__%d", id, time.Now().UnixNano())
if existingTx := dbent.TxFromContext(ctx); existingTx != nil {
return r.deleteWithAudit(ctx, existingTx.Client(), id, tombstoneKey)
return r.deleteWithTombstone(ctx, existingTx.Client(), id, tombstoneKey)
}
tx, err := r.client.Tx(ctx)
@@ -348,7 +346,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
exec = tx.Client()
}
if err := r.deleteWithAudit(ctx, exec, id, tombstoneKey); err != nil {
if err := r.deleteWithTombstone(ctx, exec, id, tombstoneKey); err != nil {
return err
}
@@ -358,17 +356,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
return nil
}
func (r *apiKeyRepository) deleteWithAudit(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
// 1. 审计:数据源即 api_keys 当前行;WHERE deleted_at IS NULL 保证只对未删除行写一次。
if _, err := exec.ExecContext(ctx, `
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
SELECT key, id, user_id, name, NOW()
FROM api_keys
WHERE id = $1 AND deleted_at IS NULL`, id); err != nil {
return err
}
// 2. 软删除(tombstone 覆盖 key)。
func (r *apiKeyRepository) deleteWithTombstone(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
res, err := exec.ExecContext(ctx, `
UPDATE api_keys
SET key = $1, deleted_at = NOW(), updated_at = NOW()
@@ -556,7 +556,7 @@ func TestIncrementQuotaUsed_Concurrent(t *testing.T) {
"并发递增后总和应为 %v,实际为 %v", float64(goroutines)*increment, got.QuotaUsed)
}
func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
func (s *APIKeyRepoSuite) TestDeleteWithAudit_TombstonesWithoutRetainingCredential() {
user := s.mustCreateUser("delwithaudit@test.com")
key := &service.APIKey{
UserID: user.ID,
@@ -571,18 +571,24 @@ func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
_, err := s.repo.GetByID(s.ctx, key.ID)
s.Require().Error(err)
rows, qErr := s.client.QueryContext(s.ctx,
`SELECT key, key_name, user_id, api_key_id FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
s.Require().NoError(qErr)
defer rows.Close()
s.Require().True(rows.Next(), "expected one audit row")
var auditKey, auditName string
var auditUserID, auditAPIKeyID int64
s.Require().NoError(rows.Scan(&auditKey, &auditName, &auditUserID, &auditAPIKeyID))
s.Require().Equal("sk-del-audit-1", auditKey)
s.Require().Equal("Audit Me", auditName)
s.Require().Equal(user.ID, auditUserID)
s.Require().Equal(key.ID, auditAPIKeyID)
var tombstone string
var deletedAt time.Time
rows, err := s.repo.sql.QueryContext(s.ctx, `SELECT key, deleted_at FROM api_keys WHERE id = $1`, key.ID)
s.Require().NoError(err)
s.Require().True(rows.Next())
s.Require().NoError(rows.Scan(&tombstone, &deletedAt))
s.Require().NoError(rows.Close())
s.Require().NotEqual("sk-del-audit-1", tombstone)
s.Require().Contains(tombstone, "__deleted__")
var auditCount int
auditRows, err := s.repo.sql.QueryContext(s.ctx,
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
s.Require().NoError(err)
s.Require().True(auditRows.Next())
s.Require().NoError(auditRows.Scan(&auditCount))
s.Require().NoError(auditRows.Close())
s.Require().Zero(auditCount, "deleted credentials must not be retained")
}
func (s *APIKeyRepoSuite) TestDeleteWithAudit_RepeatIsIdempotent() {
@@ -0,0 +1,120 @@
//go:build integration
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func TestAuthCacheInvalidationTriggers_CoverSecurityMutationsOnly(t *testing.T) {
ctx := context.Background()
suffix := time.Now().UnixNano()
group := mustCreateGroup(t, integrationEntClient, &service.Group{
Name: fmt.Sprintf("auth-outbox-group-%d", suffix), RateMultiplier: 1, IsExclusive: true,
})
user := mustCreateUser(t, integrationEntClient, &service.User{
Email: fmt.Sprintf("auth-outbox-%d@example.com", suffix), Concurrency: 5,
})
groupID := group.ID
keyValue := fmt.Sprintf("sk-auth-outbox-%d", suffix)
apiKeyRepo := NewAPIKeyRepository(integrationEntClient, integrationDB)
key := &service.APIKey{UserID: user.ID, GroupID: &groupID, Key: keyValue, Name: "outbox", Status: service.StatusActive}
require.NoError(t, apiKeyRepo.Create(ctx, key))
sum := sha256.Sum256([]byte(keyValue))
cacheKey := hex.EncodeToString(sum[:])
clear := func() {
_, err := integrationDB.ExecContext(ctx, "DELETE FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey)
require.NoError(t, err)
}
count := func() int {
var value int
require.NoError(t, integrationDB.QueryRowContext(ctx,
"SELECT COUNT(*) FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey).Scan(&value))
return value
}
clear()
t.Cleanup(clear)
t.Cleanup(func() {
// Keep the shared integration database isolated for suites that assert
// platform-wide group counts. The final clear cleanup runs after this one
// and removes invalidations emitted by these hard deletes.
_, err := integrationDB.ExecContext(ctx, "DELETE FROM user_allowed_groups WHERE user_id = $1 OR group_id = $2", user.ID, group.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM api_keys WHERE id = $1", key.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", user.ID)
require.NoError(t, err)
_, err = integrationDB.ExecContext(ctx, "DELETE FROM groups WHERE id = $1", group.ID)
require.NoError(t, err)
})
_, err := integrationDB.ExecContext(ctx, `
UPDATE api_keys
SET quota_used = quota_used + 1,
usage_5h = usage_5h + 1,
last_used_at = NOW()
WHERE id = $1`, key.ID)
require.NoError(t, err)
require.Zero(t, count(), "usage-only key updates must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'disabled' WHERE id = $1", key.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "key disable must enqueue")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'active' WHERE id = $1", key.ID)
require.NoError(t, err)
clear()
userRepo := NewUserRepository(integrationEntClient, integrationDB)
loadedUser, err := userRepo.GetByID(ctx, user.ID)
require.NoError(t, err)
loadedUser.Balance += 10
require.NoError(t, userRepo.Update(ctx, loadedUser))
require.Zero(t, count(), "balance update with unchanged allowed groups must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'disabled' WHERE id = $1", user.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "user disable must enqueue all active keys")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'active' WHERE id = $1", user.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET name = name || '-cosmetic' WHERE id = $1", group.ID)
require.NoError(t, err)
require.Zero(t, count(), "cosmetic group update must not enqueue")
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'disabled' WHERE id = $1", group.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "group disable must enqueue bound keys")
clear()
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'active' WHERE id = $1", group.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx,
"INSERT INTO user_allowed_groups (user_id, group_id) VALUES ($1, $2)", user.ID, group.ID)
require.NoError(t, err)
clear()
_, err = integrationDB.ExecContext(ctx,
"DELETE FROM user_allowed_groups WHERE user_id = $1 AND group_id = $2", user.ID, group.ID)
require.NoError(t, err)
require.Equal(t, 1, count(), "exclusive-group revocation must enqueue")
clear()
require.NoError(t, apiKeyRepo.DeleteWithAudit(ctx, key.ID))
require.Equal(t, 1, count(), "tombstone delete must hash OLD.key exactly once")
var stored string
require.NoError(t, integrationDB.QueryRowContext(ctx,
"SELECT cache_key FROM auth_cache_invalidation_outbox WHERE cache_key = $1 LIMIT 1", cacheKey).Scan(&stored))
require.Equal(t, cacheKey, stored)
require.NotContains(t, stored, keyValue)
}
@@ -0,0 +1,159 @@
package repository
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
)
type authCacheInvalidationOutboxRepository struct {
db *sql.DB
}
func NewAuthCacheInvalidationOutboxRepository(db *sql.DB) service.AuthCacheInvalidationOutboxRepository {
return &authCacheInvalidationOutboxRepository{db: db}
}
func (r *authCacheInvalidationOutboxRepository) Claim(ctx context.Context, workerID string, limit int, lease time.Duration) ([]service.AuthCacheInvalidationEvent, error) {
if r == nil || r.db == nil {
return nil, errors.New("nil auth cache invalidation outbox database")
}
if limit <= 0 {
limit = 100
}
leaseSeconds := int64(lease / time.Second)
if leaseSeconds < 1 {
leaseSeconds = 30
}
rows, err := r.db.QueryContext(ctx, `
WITH candidates AS (
SELECT id
FROM auth_cache_invalidation_outbox
WHERE available_at <= NOW()
AND (claimed_at IS NULL OR claimed_at < NOW() - ($3 * INTERVAL '1 second'))
ORDER BY id ASC
LIMIT $2
FOR UPDATE SKIP LOCKED
)
UPDATE auth_cache_invalidation_outbox AS o
SET claimed_at = NOW(), claimed_by = $1
FROM candidates AS c
WHERE o.id = c.id
RETURNING o.id, o.cache_key, o.attempts, o.delivery_stage, o.created_at
`, workerID, limit, leaseSeconds)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
events := make([]service.AuthCacheInvalidationEvent, 0, limit)
for rows.Next() {
var event service.AuthCacheInvalidationEvent
if err := rows.Scan(&event.ID, &event.CacheKey, &event.Attempts, &event.Stage, &event.CreatedAt); err != nil {
return nil, err
}
event.CacheKey = strings.TrimSpace(event.CacheKey)
events = append(events, event)
}
if err := rows.Err(); err != nil {
return nil, err
}
return events, nil
}
func (r *authCacheInvalidationOutboxRepository) ScheduleSecondPass(ctx context.Context, id int64, workerID string, availableAt time.Time) error {
result, err := r.db.ExecContext(ctx, `
UPDATE auth_cache_invalidation_outbox
SET delivery_stage = 1,
available_at = $3,
last_error = NULL,
claimed_at = NULL,
claimed_by = NULL
WHERE id = $1 AND claimed_by = $2 AND delivery_stage = 0
`, id, workerID, availableAt)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d cannot schedule second pass", id)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) DeleteClaimed(ctx context.Context, id int64, workerID string) error {
result, err := r.db.ExecContext(ctx, `
DELETE FROM auth_cache_invalidation_outbox
WHERE id = $1 AND claimed_by = $2
`, id, workerID)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) RetryClaimed(ctx context.Context, id int64, workerID string, availableAt time.Time, lastError string) error {
result, err := r.db.ExecContext(ctx, `
UPDATE auth_cache_invalidation_outbox
SET attempts = attempts + 1,
available_at = $3,
last_error = $4,
claimed_at = NULL,
claimed_by = NULL
WHERE id = $1 AND claimed_by = $2
`, id, workerID, availableAt, lastError)
if err != nil {
return err
}
affected, err := result.RowsAffected()
if err != nil {
return err
}
if affected != 1 {
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
}
return nil
}
func (r *authCacheInvalidationOutboxRepository) Stats(ctx context.Context) (service.AuthCacheInvalidationOutboxStats, error) {
var (
stats service.AuthCacheInvalidationOutboxStats
oldest sql.NullTime
lastError sql.NullString
)
err := r.db.QueryRowContext(ctx, `
SELECT COUNT(*), MIN(created_at), COALESCE(MAX(attempts), 0),
(SELECT last_error
FROM auth_cache_invalidation_outbox
WHERE last_error IS NOT NULL
ORDER BY available_at DESC, id DESC
LIMIT 1)
FROM auth_cache_invalidation_outbox
`).Scan(&stats.Pending, &oldest, &stats.MaxAttempts, &lastError)
if err != nil {
return stats, err
}
if oldest.Valid {
value := oldest.Time
stats.OldestCreatedAt = &value
}
if lastError.Valid {
stats.LastError = lastError.String
}
return stats, nil
}
@@ -0,0 +1,123 @@
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"strings"
"testing"
"time"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/Wei-Shaw/sub2api/migrations"
"github.com/stretchr/testify/require"
)
func TestAuthCacheInvalidationOutboxRepository_ClaimUsesLeaseAndSkipLocked(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
created := time.Now().UTC()
mock.ExpectQuery("(?s)claimed_at < NOW\\(\\) - .*FOR UPDATE SKIP LOCKED.*RETURNING").
WithArgs("worker-a", 100, int64(30)).
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}).
AddRow(int64(4), strings.Repeat("a", 64), 2, 1, created))
repo := NewAuthCacheInvalidationOutboxRepository(db)
events, err := repo.Claim(context.Background(), "worker-a", 100, 30*time.Second)
require.NoError(t, err)
require.Len(t, events, 1)
require.Equal(t, int64(4), events[0].ID)
require.Equal(t, 1, events[0].Stage)
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_ClaimIsBoundedByDefault(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
mock.ExpectQuery("(?s)FROM auth_cache_invalidation_outbox.*LIMIT \\$2.*SKIP LOCKED").
WithArgs("worker", 100, int64(30)).
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}))
repo := NewAuthCacheInvalidationOutboxRepository(db)
_, err = repo.Claim(context.Background(), "worker", 0, 0)
require.NoError(t, err)
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_ClaimOwnershipTransitions(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
repo := NewAuthCacheInvalidationOutboxRepository(db)
next := time.Now().UTC().Add(time.Minute)
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
WithArgs(int64(1), "worker", next).
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.ScheduleSecondPass(context.Background(), 1, "worker", next))
retryAt := next.Add(time.Minute)
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
WithArgs(int64(2), "worker", retryAt, "publish failed").
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.RetryClaimed(context.Background(), 2, "worker", retryAt, "publish failed"))
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
WithArgs(int64(3), "worker").
WillReturnResult(sqlmock.NewResult(0, 1))
require.NoError(t, repo.DeleteClaimed(context.Background(), 3, "worker"))
require.NoError(t, mock.ExpectationsWereMet())
}
func TestAuthCacheInvalidationOutboxRepository_RejectsLostClaim(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
WithArgs(int64(3), "old-worker").
WillReturnResult(sqlmock.NewResult(0, 0))
repo := NewAuthCacheInvalidationOutboxRepository(db)
err = repo.DeleteClaimed(context.Background(), 3, "old-worker")
require.ErrorContains(t, err, "no longer owned")
}
func TestAuthCacheInvalidationOutboxRepository_StatsExposeDurableLagAndFailures(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
defer func() { _ = db.Close() }()
oldest := time.Now().UTC().Add(-time.Minute)
mock.ExpectQuery("(?s)SELECT COUNT\\(\\*\\), MIN\\(created_at\\), COALESCE\\(MAX\\(attempts\\), 0\\)").
WillReturnRows(sqlmock.NewRows([]string{"count", "min", "max", "last_error"}).AddRow(5, oldest, 7, "redis down"))
repo := NewAuthCacheInvalidationOutboxRepository(db)
stats, err := repo.Stats(context.Background())
require.NoError(t, err)
require.Equal(t, int64(5), stats.Pending)
require.Equal(t, 7, stats.MaxAttempts)
require.Equal(t, "redis down", stats.LastError)
require.NotNil(t, stats.OldestCreatedAt)
}
func TestAuthCacheInvalidationMigration_SecurityCoverageAndNoPlaintextPayload(t *testing.T) {
content, err := migrations.FS.ReadFile("184_auth_cache_invalidation_outbox.sql")
require.NoError(t, err)
sqlText := string(content)
for _, required := range []string{
"encode(sha256(convert_to(raw_key, 'UTF8')), 'hex')",
"OLD.key", "OLD.status", "OLD.deleted_at", "OLD.user_id", "OLD.group_id",
"OLD.ip_whitelist", "OLD.ip_blacklist", "OLD.expires_at",
"trg_users_auth_cache_invalidation", "trg_groups_auth_cache_invalidation",
"trg_user_allowed_groups_auth_cache_invalidation", "FOR EACH ROW",
"delivery_stage", "claimed_at", "available_at",
} {
require.Contains(t, sqlText, required)
}
require.NotContains(t, sqlText, "quota_used IS DISTINCT")
require.NotContains(t, sqlText, "last_used_at IS DISTINCT")
plaintext := "sk-plaintext-must-not-be-stored"
sum := sha256.Sum256([]byte(plaintext))
require.Len(t, hex.EncodeToString(sum[:]), 64)
require.NotContains(t, sqlText, plaintext)
}
@@ -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 rows.Close()
result := &service.OpsIngressRejectList{
Items: make([]*service.OpsIngressRejectAggregate, 0, pageSize), Total: total, Page: page, PageSize: pageSize,
}
for rows.Next() {
item := &service.OpsIngressRejectAggregate{}
var userID, apiKeyID int64
if err := rows.Scan(&item.ID, &item.BucketStart, &item.RejectReason, &item.RouteFamily, &item.Protocol,
&item.ClientIP, &userID, &apiKeyID, &item.RequestCount, &item.FirstSeen, &item.LastSeen); err != nil {
return nil, err
}
if userID > 0 {
item.UserID = &userID
}
if apiKeyID > 0 {
item.APIKeyID = &apiKeyID
}
result.Items = append(result.Items, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return result, nil
}
@@ -0,0 +1,32 @@
package repository
import (
"context"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func TestBatchUpsertIngressRejectsUsesFixedMultiRowChunks(t *testing.T) {
db, mock, err := sqlmock.New()
require.NoError(t, err)
t.Cleanup(func() { _ = db.Close() })
repo := &opsRepository{db: db}
now := time.Now().UTC().Truncate(time.Minute)
items := make([]*service.OpsIngressRejectAggregate, ingressRejectUpsertChunkSize+1)
for i := range items {
items[i] = &service.OpsIngressRejectAggregate{
BucketStart: now, RejectReason: "invalid_api_key", RouteFamily: "messages",
Protocol: "anthropic", ClientIP: "192.0.2.1", RequestCount: 1, FirstSeen: now, LastSeen: now,
}
}
mock.ExpectBegin()
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, int64(ingressRejectUpsertChunkSize)))
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
require.NoError(t, repo.BatchUpsertIngressRejects(context.Background(), items))
require.NoError(t, mock.ExpectationsWereMet())
}
+7 -82
View File
@@ -4,7 +4,6 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
@@ -56,12 +55,9 @@ INSERT INTO ops_error_logs (
response_latency_ms,
time_to_first_token_ms,
created_at,
attempted_key_prefix,
deleted_key_owner_user_id,
deleted_key_name,
api_key_prefix
) VALUES (
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38,$39,$40,$41
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38
)`
func NewOpsRepository(db *sql.DB) service.OpsRepository {
@@ -170,9 +166,6 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
opsNullInt64(input.ResponseLatencyMs),
opsNullInt64(input.TimeToFirstTokenMs),
input.CreatedAt,
opsNullString(input.AttemptedKeyPrefix),
opsNullInt64(input.DeletedKeyOwnerUserID),
opsNullString(input.DeletedKeyName),
opsNullString(input.APIKeyPrefix),
}
}
@@ -274,16 +267,12 @@ SELECT
COALESCE(e.user_agent, ''),
e.request_type,
COALESCE(ak.name, ''),
ak.deleted_at,
COALESCE(e.deleted_key_name, ''),
e.deleted_key_owner_user_id,
COALESCE(du.email, '')
ak.deleted_at
FROM ops_error_logs e
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
` + where + `
ORDER BY ` + opsErrorLogsOrderBy(filter) + `
@@ -313,9 +302,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
var requestType sql.NullInt64
var apiKeyName string
var apiKeyDeletedAt sql.NullTime
var deletedKeyName string
var deletedKeyOwnerID sql.NullInt64
var deletedKeyOwnerEmail string
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
@@ -352,9 +338,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
&requestType,
&apiKeyName,
&apiKeyDeletedAt,
&deletedKeyName,
&deletedKeyOwnerID,
&deletedKeyOwnerEmail,
); err != nil {
return nil, err
}
@@ -395,21 +378,8 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
v := int16(requestType.Int64)
item.RequestType = &v
}
// Key 名称:优先关联到的 ak.name(已软删的 key name 仍保留);
// 关联不到(api_key_id 为空 / 历史硬删)时回退错误记录里快照的 deleted_key_name。
if apiKeyName != "" {
item.APIKeyName = apiKeyName
} else {
item.APIKeyName = deletedKeyName
}
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
// 已删除 KEY 所有者快照:认证失败行 user_id 为空,列表用户列以此回退。
if deletedKeyOwnerID.Valid {
v := deletedKeyOwnerID.Int64
item.DeletedKeyOwnerUserID = &v
item.DeletedKeyOwnerEmail = deletedKeyOwnerEmail
}
item.APIKeyName = apiKeyName
item.APIKeyDeleted = apiKeyDeletedAt.Valid
out = append(out, &item)
}
if err := rows.Err(); err != nil {
@@ -477,10 +447,6 @@ SELECT
e.upstream_latency_ms,
e.response_latency_ms,
e.time_to_first_token_ms,
COALESCE(e.attempted_key_prefix, ''),
e.deleted_key_owner_user_id,
COALESCE(du.email, ''),
COALESCE(e.deleted_key_name, ''),
COALESCE(e.api_key_prefix, ''),
COALESCE(ak.name, ''),
ak.deleted_at
@@ -488,7 +454,6 @@ FROM ops_error_logs e
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
WHERE e.id = $1
LIMIT 1`
@@ -509,7 +474,6 @@ LIMIT 1`
var responseLatency sql.NullInt64
var ttft sql.NullInt64
var requestType sql.NullInt64
var deletedKeyOwnerUserID sql.NullInt64
var detailAPIKeyName string
var detailAPIKeyDeletedAt sql.NullTime
@@ -557,10 +521,6 @@ LIMIT 1`
&upstreamLatency,
&responseLatency,
&ttft,
&out.AttemptedKeyPrefix,
&deletedKeyOwnerUserID,
&out.DeletedKeyOwnerEmail,
&out.DeletedKeyName,
&out.APIKeyPrefix,
&detailAPIKeyName,
&detailAPIKeyDeletedAt,
@@ -626,18 +586,8 @@ LIMIT 1`
v := int16(requestType.Int64)
out.RequestType = &v
}
if deletedKeyOwnerUserID.Valid {
v := deletedKeyOwnerUserID.Int64
out.DeletedKeyOwnerUserID = &v
}
// Key 名称:优先关联到的 ak.name;关联不到时回退快照的 deleted_key_name。
if detailAPIKeyName != "" {
out.APIKeyName = detailAPIKeyName
} else {
out.APIKeyName = out.DeletedKeyName
}
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid || (detailAPIKeyName == "" && out.DeletedKeyName != "")
out.APIKeyName = detailAPIKeyName
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid
// Normalize upstream_errors to empty string when stored as JSON null.
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
@@ -648,26 +598,6 @@ LIMIT 1`
return &out, nil
}
// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计。
// 同一 key 可能有多条历史(反复创建/删除),取 deleted_at 最近一条(id 作同毫秒 tiebreaker)。
// 未命中返回 (nil, nil)。
func (r *opsRepository) LookupDeletedKeyAudit(ctx context.Context, key string) (*service.DeletedKeyAuditResult, error) {
var res service.DeletedKeyAuditResult
err := r.db.QueryRowContext(ctx, `
SELECT user_id, key_name
FROM deleted_api_key_audits
WHERE key = $1
ORDER BY deleted_at DESC, id DESC
LIMIT 1`, key).Scan(&res.UserID, &res.KeyName)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
return nil, err
}
return &res, nil
}
func (r *opsRepository) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64, resolvedAt *time.Time) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil ops repository")
@@ -1082,12 +1012,7 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
if filter.UserID != nil && *filter.UserID > 0 {
args = append(args, *filter.UserID)
n := itoa(len(args))
if filter.MatchDeletedKeyOwner {
// 用户侧:把「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录也纳入。
clauses = append(clauses, "(e.user_id = $"+n+" OR e.deleted_key_owner_user_id = $"+n+")")
} else {
clauses = append(clauses, "e.user_id = $"+n)
}
clauses = append(clauses, "e.user_id = $"+n)
}
if filter.APIKeyID != nil && *filter.APIKeyID > 0 {
args = append(args, *filter.APIKeyID)
@@ -11,47 +11,13 @@ import (
"github.com/stretchr/testify/require"
)
// TestGetErrorLogByID_DeletedKeyOwner 验证:
// 1. 带 deleted_key_owner_user_id 的记录能正确 JOIN users 返回 DeletedKeyOwnerEmail
// 2. 新列全为 NULL 的普通记录 Scan 不报错,这些字段为空/nil
func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
func TestGetErrorLogByID_APIKeyPrefixAndUpstreamStatus(t *testing.T) {
ctx := context.Background()
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE ops_error_logs RESTART IDENTITY CASCADE")
repo := NewOpsRepository(integrationDB).(*opsRepository)
// ── Case 1: 带 deleted_key_owner 信息的记录 ──────────────────────────────
owner := mustCreateUser(t, integrationEntClient, &service.User{
Email: "deleted-key-owner-" + time.Now().Format("150405.000000000") + "@example.com",
})
var insertedID int64
err := integrationDB.QueryRowContext(ctx, `
INSERT INTO ops_error_logs (
error_phase, error_type, severity, status_code, created_at,
attempted_key_prefix, deleted_key_owner_user_id, deleted_key_name
) VALUES (
'auth', 'INVALID_API_KEY', 'error', 401, NOW(),
'sk-test-abc', $1, 'my-deleted-key'
) RETURNING id`,
owner.ID,
).Scan(&insertedID)
require.NoError(t, err)
require.Positive(t, insertedID)
detail, err := repo.GetErrorLogByID(ctx, insertedID)
require.NoError(t, err)
require.NotNil(t, detail)
require.Equal(t, "sk-test-abc", detail.AttemptedKeyPrefix)
require.NotNil(t, detail.DeletedKeyOwnerUserID)
require.Equal(t, owner.ID, *detail.DeletedKeyOwnerUserID)
require.Equal(t, owner.Email, detail.DeletedKeyOwnerEmail)
require.Equal(t, "my-deleted-key", detail.DeletedKeyName)
// ── Case 2: 新列全为 NULL 的普通错误记录 ──────────────────────────────────
var plainID int64
err = integrationDB.QueryRowContext(ctx, `
err := integrationDB.QueryRowContext(ctx, `
INSERT INTO ops_error_logs (
error_phase, error_type, severity, status_code, created_at
) VALUES (
@@ -59,20 +25,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
) RETURNING id`,
).Scan(&plainID)
require.NoError(t, err)
require.Positive(t, plainID)
plain, err := repo.GetErrorLogByID(ctx, plainID)
require.NoError(t, err)
require.NotNil(t, plain)
require.Empty(t, plain.APIKeyPrefix)
require.Empty(t, plain.AttemptedKeyPrefix, "no prefix for plain error")
require.Nil(t, plain.DeletedKeyOwnerUserID, "no owner for plain error")
require.Empty(t, plain.DeletedKeyOwnerEmail, "no owner email for plain error")
require.Empty(t, plain.DeletedKeyName, "no key name for plain error")
require.Empty(t, plain.APIKeyPrefix, "no api key prefix for plain error")
// ── Case 3: 有效(未删除)key 报错,经 InsertErrorLog 快照 api_key_prefix ──────
// 走真实 InsertErrorLog 写入路径(覆盖新列 + $41 占位符),再 GetErrorLogByID 读回。
validID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
ErrorPhase: "request",
ErrorType: "api_error",
@@ -82,17 +39,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
APIKeyPrefix: "sk-valid",
})
require.NoError(t, err)
require.Positive(t, validID)
valid, err := repo.GetErrorLogByID(ctx, validID)
require.NoError(t, err)
require.NotNil(t, valid)
require.Equal(t, "sk-valid", valid.APIKeyPrefix)
require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive")
require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner")
// ── Case 4: account_auth with no inference attempt preserves explicit 0 ──
zero := 0
credentialFailureID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
ErrorPhase: "account_auth",
@@ -1,36 +0,0 @@
//go:build integration
package repository
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestOpsRepositoryLookupDeletedKeyAudit(t *testing.T) {
ctx := context.Background()
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE deleted_api_key_audits RESTART IDENTITY")
repo := NewOpsRepository(integrationDB).(*opsRepository)
// 同一 key 两条审计,取最近一条(deleted_at DESC, id DESC)
_, err := integrationDB.ExecContext(ctx, `
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
VALUES ('sk-lookup-1', 10, 100, 'old', $1),
('sk-lookup-1', 11, 200, 'new', $2)`,
time.Now().Add(-time.Hour), time.Now())
require.NoError(t, err)
res, err := repo.LookupDeletedKeyAudit(ctx, "sk-lookup-1")
require.NoError(t, err)
require.NotNil(t, res)
require.Equal(t, int64(200), res.UserID)
require.Equal(t, "new", res.KeyName)
// 未命中返回 nil
miss, err := repo.LookupDeletedKeyAudit(ctx, "sk-never-existed")
require.NoError(t, err)
require.Nil(t, miss)
}
+27 -7
View File
@@ -1039,24 +1039,44 @@ func (r *userRepository) syncUserAllowedGroupsWithClient(ctx context.Context, cl
return nil
}
// Keep join table as the source of truth for reads.
if _, err := client.UserAllowedGroup.Delete().Where(userallowedgroup.UserIDEQ(userID)).Exec(ctx); err != nil {
existingRows, err := client.UserAllowedGroup.Query().
Where(userallowedgroup.UserIDEQ(userID)).
All(ctx)
if err != nil {
return err
}
unique := make(map[int64]struct{}, len(groupIDs))
desired := make(map[int64]struct{}, len(groupIDs))
for _, id := range groupIDs {
if id <= 0 {
continue
}
unique[id] = struct{}{}
desired[id] = struct{}{}
}
if len(unique) > 0 {
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(unique))
for groupID := range unique {
existing := make(map[int64]struct{}, len(existingRows))
removed := make([]int64, 0)
for _, row := range existingRows {
existing[row.GroupID] = struct{}{}
if _, keep := desired[row.GroupID]; !keep {
removed = append(removed, row.GroupID)
}
}
if len(removed) > 0 {
if _, err := client.UserAllowedGroup.Delete().
Where(userallowedgroup.UserIDEQ(userID), userallowedgroup.GroupIDIn(removed...)).
Exec(ctx); err != nil {
return err
}
}
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(desired))
for groupID := range desired {
if _, present := existing[groupID]; !present {
creates = append(creates, client.UserAllowedGroup.Create().SetUserID(userID).SetGroupID(groupID))
}
}
if len(creates) > 0 {
if err := client.UserAllowedGroup.
CreateBulk(creates...).
OnConflictColumns(userallowedgroup.FieldUserID, userallowedgroup.FieldGroupID).
@@ -14,7 +14,7 @@ import (
)
// TestUserRepository_DeleteUser_AtomicWithAPIKeys 复现 AdminService.DeleteUser 的事务编排场景:
// 把"删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
// 把"tombstone 并删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
// userRepo.Delete 必须复用 context 中的事务,而不是用 base client 自起一个独立事务并提前提交。
//
// 用例用"回滚外层事务"来模拟 commit 失败 / 中止:
@@ -91,5 +91,5 @@ func TestUserRepository_DeleteUser_AtomicWithAPIKeys(t *testing.T) {
require.NoError(t, integrationDB.QueryRowContext(ctx,
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE user_id = $1`, user.ID).Scan(&auditCount))
require.Equal(t, 2, auditCount, "提交后应为每个被删 Key 写入一行审计")
require.Zero(t, auditCount, "提交后也不得保留被删 Key 的凭据材料")
}
+1
View File
@@ -126,6 +126,7 @@ var ProviderSet = wire.NewSet(
NewLeaderLockCache,
ProvideSchedulerCache,
NewSchedulerOutboxRepository,
NewAuthCacheInvalidationOutboxRepository,
NewProxyLatencyCache,
NewTotpCache,
NewRefreshTokenCache,
@@ -16,6 +16,7 @@ var ProviderSet = wire.NewSet(
wire.Bind(new(ConfigStore), new(*ConfigManager)),
NewPromptService,
wire.Bind(new(PromptEngine), new(*PromptService)),
wire.Bind(new(PromptAdminService), new(*PromptService)),
NewLegacyModerationAdapter,
NewCoordinator,
NewPromptAdminHandler,
+3 -2
View File
@@ -103,8 +103,9 @@ func ProvideRouter(
func ProvideHTTPServer(cfg *config.Config, router *gin.Engine) *http.Server {
httpHandler := http.Handler(router)
server := &http.Server{
Addr: cfg.Server.Address(),
Handler: httpHandler,
Addr: cfg.Server.Address(),
Handler: httpHandler,
MaxHeaderBytes: cfg.Server.MaxHeaderBytes,
// ReadHeaderTimeout: 读取请求头的超时时间,防止慢速请求头攻击
ReadHeaderTimeout: time.Duration(cfg.Server.ReadHeaderTimeout) * time.Second,
// IdleTimeout: 空闲连接超时时间,释放不活跃的连接资源
@@ -0,0 +1,126 @@
//go:build unit
package server
import (
"bufio"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func ingressTestConfig() *config.Config {
return &config.Config{
Server: config.ServerConfig{
Host: "127.0.0.1",
ReadHeaderTimeout: 1,
IdleTimeout: 5,
MaxHeaderBytes: 8 * 1024,
MaxRequestBodySize: 1024,
},
Gateway: config.GatewayConfig{MaxBodySize: 1024},
}
}
func TestProvideHTTPServerAppliesIngressLimits(t *testing.T) {
srv := ProvideHTTPServer(ingressTestConfig(), gin.New())
require.Equal(t, 8*1024, srv.MaxHeaderBytes)
require.Equal(t, time.Second, srv.ReadHeaderTimeout)
require.Equal(t, 5*time.Second, srv.IdleTimeout)
}
func TestProvideHTTPServerEnablesBoundedH2C(t *testing.T) {
cfg := ingressTestConfig()
cfg.Server.H2C = config.H2CConfig{
Enabled: true,
MaxConcurrentStreams: 25,
IdleTimeout: 30,
MaxReadFrameSize: 64 * 1024,
MaxUploadBufferPerConnection: 1024 * 1024,
MaxUploadBufferPerStream: 256 * 1024,
}
srv := ProvideHTTPServer(cfg, gin.New())
require.NotNil(t, srv.Protocols)
require.True(t, srv.Protocols.UnencryptedHTTP2())
require.True(t, srv.Protocols.HTTP1())
}
func TestHTTPServerRejectsOversizedHTTP1Header(t *testing.T) {
r := gin.New()
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
srv := ProvideHTTPServer(ingressTestConfig(), r)
addr, stop := serveIngressTestServer(t, srv)
defer stop()
conn, err := net.DialTimeout("tcp", addr, time.Second)
require.NoError(t, err)
defer func() { _ = conn.Close() }()
_ = conn.SetDeadline(time.Now().Add(3 * time.Second))
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Fill: "+strings.Repeat("a", 32*1024)+"\r\n\r\n")
require.NoError(t, err)
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusRequestHeaderFieldsTooLarge, resp.StatusCode)
}
func TestHTTPServerClosesSlowIncompleteHeader(t *testing.T) {
r := gin.New()
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
srv := ProvideHTTPServer(ingressTestConfig(), r)
addr, stop := serveIngressTestServer(t, srv)
defer stop()
conn, err := net.DialTimeout("tcp", addr, time.Second)
require.NoError(t, err)
defer func() { _ = conn.Close() }()
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Slow:")
require.NoError(t, err)
time.Sleep(1200 * time.Millisecond)
_ = conn.SetReadDeadline(time.Now().Add(time.Second))
_, err = bufio.NewReader(conn).ReadByte()
require.Error(t, err)
}
func TestHTTPServerGlobalBodyLimit(t *testing.T) {
r := gin.New()
r.POST("/", func(c *gin.Context) {
_, err := io.ReadAll(c.Request.Body)
if err != nil {
var maxErr *http.MaxBytesError
if errors.As(err, &maxErr) {
c.Status(http.StatusRequestEntityTooLarge)
return
}
}
c.Status(http.StatusOK)
})
srv := ProvideHTTPServer(ingressTestConfig(), r)
req, err := http.NewRequest(http.MethodPost, "/", strings.NewReader(strings.Repeat("x", 1025)))
require.NoError(t, err)
rec := httptest.NewRecorder()
srv.Handler.ServeHTTP(rec, req)
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code)
}
func serveIngressTestServer(t *testing.T, srv *http.Server) (string, func()) {
t.Helper()
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
go func() { _ = srv.Serve(ln) }()
return ln.Addr().String(), func() {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = srv.Shutdown(ctx)
}
}
@@ -15,6 +15,8 @@ import (
"github.com/gin-gonic/gin"
)
const maxAPIKeyAuthorizationHeaderBytes = service.MaxAPIKeyCredentialBytes + 128
// NewAPIKeyAuthMiddleware 创建 API Key 认证中间件
func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) APIKeyAuthMiddleware {
return APIKeyAuthMiddleware(apiKeyAuthWithSubscription(apiKeyService, subscriptionService, cfg))
@@ -32,10 +34,23 @@ func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionS
func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
// ── 1. 提取 API Key ──────────────────────────────────────────
if rejectInvalidAuthAbuse(c, apiKeyService) {
AbortWithError(c, http.StatusTooManyRequests, "INVALID_AUTH_RATE_LIMITED", "Too many invalid authentication attempts; retry later")
return
}
if apiKeyHeadersTooLarge(c) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
return
}
queryKey := strings.TrimSpace(c.Query("key"))
queryApiKey := strings.TrimSpace(c.Query("api_key"))
if queryKey != "" || queryApiKey != "" {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
AbortWithError(c, 400, "api_key_in_query_deprecated", "API key in query parameter is deprecated. Please use Authorization header instead.")
return
}
@@ -56,6 +71,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
if apiKeyString == "" {
apiKeyString = c.GetHeader("x-api-key")
}
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
return
}
// 如果x-api-key header中没有,尝试从x-goog-api-key header中提取(Gemini CLI兼容)
if apiKeyString == "" {
@@ -64,6 +85,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
// 如果所有header都没有API key
if apiKeyString == "" {
recordInvalidAuthFailure(c, apiKeyService)
if hasAPIKeyCredentialInput(c) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
} else {
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
}
AbortWithError(c, 401, "API_KEY_REQUIRED", "API key is required in Authorization header (Bearer scheme), x-api-key header, or x-goog-api-key header")
return
}
@@ -73,9 +100,16 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
if err != nil {
if errors.Is(err, service.ErrAPIKeyNotFound) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
AbortWithError(c, 401, "INVALID_API_KEY", "Invalid API key")
return
}
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
AbortWithError(c, http.StatusServiceUnavailable, "API_KEY_AUTH_OVERLOADED", "API key authentication is temporarily unavailable")
return
}
AbortWithError(c, 500, "INTERNAL_ERROR", "Failed to validate API key")
return
}
@@ -90,6 +124,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
if !apiKey.IsActive() &&
apiKey.Status != service.StatusAPIKeyExpired &&
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
AbortWithError(c, 401, "API_KEY_DISABLED", "API key is disabled")
return
}
@@ -104,6 +139,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
clientIP = "unknown"
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
MarkIngressRejected(c, IngressRejectIPRestricted)
AbortWithError(c, 403, "ACCESS_DENIED", fmt.Sprintf("Access denied. Your IP is %s", clientIP))
return
}
@@ -117,6 +153,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
// 检查用户状态
if !apiKey.User.IsActive() {
MarkIngressRejected(c, IngressRejectUserInactive)
AbortWithError(c, 401, "USER_INACTIVE", "User account is not active")
return
}
@@ -250,6 +287,24 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
}
}
func apiKeyHeadersTooLarge(c *gin.Context) bool {
if c == nil {
return false
}
return len(c.GetHeader("Authorization")) > maxAPIKeyAuthorizationHeaderBytes ||
len(c.GetHeader("x-api-key")) > service.MaxAPIKeyCredentialBytes ||
len(c.GetHeader("x-goog-api-key")) > service.MaxAPIKeyCredentialBytes
}
func hasAPIKeyCredentialInput(c *gin.Context) bool {
if c == nil {
return false
}
return c.GetHeader("Authorization") != "" ||
c.GetHeader("x-api-key") != "" ||
c.GetHeader("x-goog-api-key") != ""
}
func isAsyncImageTaskRead(method, path string) bool {
if method != http.MethodGet {
return false
@@ -321,6 +376,11 @@ func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool
return false
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
if code == "GROUP_DELETED" {
MarkIngressRejected(c, IngressRejectGroupDeleted)
} else {
MarkIngressRejected(c, IngressRejectGroupDisabled)
}
AbortWithError(c, 403, code, message)
return true
}
@@ -330,6 +390,7 @@ func abortIfAPIKeyGroupNotAllowed(c *gin.Context, apiKey *service.APIKey) bool {
return false
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
AbortWithError(c, 403, "GROUP_NOT_ALLOWED", "API Key 所属专属分组不再允许当前用户使用")
return true
}
@@ -24,22 +24,53 @@ func APIKeyAuthGoogle(apiKeyService *service.APIKeyService, cfg *config.Config)
// It is intended for Gemini native endpoints (/v1beta) to match Gemini SDK expectations.
func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
if rejectInvalidAuthAbuse(c, apiKeyService) {
abortWithGoogleError(c, 429, "Too many invalid authentication attempts; retry later")
return
}
if apiKeyHeadersTooLarge(c) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
if v := strings.TrimSpace(c.Query("api_key")); v != "" {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
abortWithGoogleError(c, 400, "Query parameter api_key is deprecated. Use Authorization header or key instead.")
return
}
apiKeyString := extractAPIKeyForGoogle(c)
if apiKeyString == "" {
recordInvalidAuthFailure(c, apiKeyService)
if hasAPIKeyCredentialInput(c) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
} else {
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
}
abortWithGoogleError(c, 401, "API key is required")
return
}
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
if err != nil {
if errors.Is(err, service.ErrAPIKeyNotFound) {
recordInvalidAuthFailure(c, apiKeyService)
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
abortWithGoogleError(c, 401, "Invalid API key")
return
}
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
abortWithGoogleError(c, 503, "API key authentication is temporarily unavailable")
return
}
abortWithGoogleError(c, 500, "Failed to validate API key")
return
}
@@ -53,6 +84,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
if !apiKey.IsActive() &&
apiKey.Status != service.StatusAPIKeyExpired &&
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
abortWithGoogleError(c, 401, "API key is disabled")
return
}
@@ -66,6 +98,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
clientIP = "unknown"
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
MarkIngressRejected(c, IngressRejectIPRestricted)
abortWithGoogleError(c, 403, fmt.Sprintf("Access denied. Your IP is %s", clientIP))
return
}
@@ -76,17 +109,24 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
return
}
if !apiKey.User.IsActive() {
MarkIngressRejected(c, IngressRejectUserInactive)
abortWithGoogleError(c, 401, "User account is not active")
return
}
if _, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
if code, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
if code == "GROUP_DELETED" {
MarkIngressRejected(c, IngressRejectGroupDeleted)
} else {
MarkIngressRejected(c, IngressRejectGroupDisabled)
}
abortWithGoogleError(c, 403, message)
return
}
// 专属分组授权校验:用户对该专属分组的授权被撤销后应拒绝(与主中间件一致,防止越权)。
if !validateAPIKeyGroupAllowed(apiKey) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
abortWithGoogleError(c, 403, "API Key 所属专属分组不再允许当前用户使用")
return
}
@@ -6,6 +6,8 @@ import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
@@ -18,6 +20,59 @@ import (
"github.com/stretchr/testify/require"
)
func TestGoogleAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
gin.SetMode(gin.TestMode)
var calls atomic.Int32
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
calls.Add(1)
return nil, service.ErrAPIKeyNotFound
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthGoogle(svc, cfg))
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
req.Header.Set("x-goog-api-key", strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1))
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Zero(t, calls.Load())
require.True(t, rejected)
require.Equal(t, IngressRejectInvalidAPIKey, reason)
}
func TestGoogleAPIKeyAuthMarksLookupBulkheadRejection(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
return nil, service.ErrAPIKeyAuthOverloaded
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthGoogle(svc, cfg))
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
req.Header.Set("x-goog-api-key", "valid-shape")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusServiceUnavailable, w.Code)
require.True(t, rejected)
require.Equal(t, IngressRejectAPIKeyAuthOverloaded, reason)
}
type fakeAPIKeyRepo struct {
getByKey func(ctx context.Context, key string) (*service.APIKey, error)
updateLastUsed func(ctx context.Context, id int64, usedAt time.Time) error
@@ -376,6 +431,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
return nil, service.ErrAPIKeyNotFound
},
})
var rejectReason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
rejectReason, rejected = GetIngressRejectReason(c)
})
r.Use(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, &config.Config{}))
r.GET("/v1beta/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
@@ -390,6 +451,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
require.Equal(t, http.StatusUnauthorized, resp.Error.Code)
require.Equal(t, "Invalid API key", resp.Error.Message)
require.Equal(t, "UNAUTHENTICATED", resp.Error.Status)
require.True(t, rejected)
require.Equal(t, IngressRejectInvalidAPIKey, rejectReason)
}
func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t *testing.T) {
@@ -422,9 +485,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
r := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
r.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -452,6 +518,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
require.Equal(t, "API Key 所属分组已删除", resp.Error.Message)
require.True(t, markedBusinessLimited)
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
require.True(t, rejected)
require.Equal(t, IngressRejectGroupDeleted, rejectReason)
}
func TestApiKeyAuthWithSubscriptionGoogle_RepoError(t *testing.T) {
@@ -8,6 +8,8 @@ import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
@@ -19,6 +21,35 @@ import (
"github.com/stretchr/testify/require"
)
func TestAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
gin.SetMode(gin.TestMode)
var calls atomic.Int32
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
calls.Add(1)
return nil, service.ErrAPIKeyNotFound
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
for _, headers := range []map[string]string{
{"x-api-key": strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
{"Authorization": "Bearer " + strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
{"Authorization": strings.Repeat("x", maxAPIKeyAuthorizationHeaderBytes+1)},
} {
r := gin.New()
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/t", nil)
for name, value := range headers {
req.Header.Set(name, value)
}
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
}
require.Zero(t, calls.Load())
}
func TestSimpleModeBypassesQuotaCheck(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -436,6 +467,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus int
wantCode string
wantMarked bool
wantReject IngressRejectReason
}{
{
name: "active group passes",
@@ -460,6 +492,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DISABLED",
wantMarked: true,
wantReject: IngressRejectGroupDisabled,
},
{
name: "deleted status group is forbidden",
@@ -473,6 +506,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DELETED",
wantMarked: true,
wantReject: IngressRejectGroupDeleted,
},
{
name: "missing group edge is forbidden",
@@ -480,6 +514,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
wantStatus: http.StatusForbidden,
wantCode: "GROUP_DELETED",
wantMarked: true,
wantReject: IngressRejectGroupDeleted,
},
}
@@ -508,9 +543,12 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
router := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -530,6 +568,8 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
require.Contains(t, w.Body.String(), tt.wantCode)
}
require.Equal(t, tt.wantMarked, markedBusinessLimited)
require.Equal(t, tt.wantReject != "", rejected)
require.Equal(t, tt.wantReject, rejectReason)
if tt.wantMarked {
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
}
@@ -537,6 +577,112 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
}
}
func TestAPIKeyAuthMarksOnlyExpectedIngressRejections(t *testing.T) {
gin.SetMode(gin.TestMode)
tests := []struct {
name string
path string
key string
authHeader string
repoErr error
wantStatus int
wantCode string
wantReason IngressRejectReason
}{
{
name: "query key deprecated",
path: "/t?key=legacy",
wantStatus: http.StatusBadRequest,
wantCode: "api_key_in_query_deprecated",
wantReason: IngressRejectQueryAPIKeyDeprecated,
},
{
name: "missing key",
path: "/t",
wantStatus: http.StatusUnauthorized,
wantCode: "API_KEY_REQUIRED",
wantReason: IngressRejectAPIKeyRequired,
},
{
name: "malformed authorization",
path: "/t",
authHeader: "Basic not-a-bearer-key",
wantStatus: http.StatusUnauthorized,
wantCode: "API_KEY_REQUIRED",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "oversized key",
path: "/t",
key: strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1),
wantStatus: http.StatusUnauthorized,
wantCode: "INVALID_API_KEY",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "invalid key",
path: "/t",
key: "invalid",
repoErr: service.ErrAPIKeyNotFound,
wantStatus: http.StatusUnauthorized,
wantCode: "INVALID_API_KEY",
wantReason: IngressRejectInvalidAPIKey,
},
{
name: "repository failure remains operational error",
path: "/t",
key: "valid-shape",
repoErr: errors.New("database unavailable"),
wantStatus: http.StatusInternalServerError,
wantCode: "INTERNAL_ERROR",
},
{
name: "auth lookup bulkhead rejection is an admission rejection",
path: "/t",
key: "valid-shape",
repoErr: service.ErrAPIKeyAuthOverloaded,
wantStatus: http.StatusServiceUnavailable,
wantCode: "API_KEY_AUTH_OVERLOADED",
wantReason: IngressRejectAPIKeyAuthOverloaded,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
return nil, tt.repoErr
}}
cfg := &config.Config{RunMode: config.RunModeSimple}
apiKeyService := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
var reason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
reason, rejected = GetIngressRejectReason(c)
})
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, tt.path, nil)
if tt.key != "" {
req.Header.Set("x-api-key", tt.key)
}
if tt.authHeader != "" {
req.Header.Set("Authorization", tt.authHeader)
}
router.ServeHTTP(w, req)
require.Equal(t, tt.wantStatus, w.Code)
require.Contains(t, w.Body.String(), tt.wantCode)
require.Equal(t, tt.wantReason != "", rejected)
require.Equal(t, tt.wantReason, reason)
})
}
}
func TestAPIKeyAuthSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -686,9 +832,12 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
router := gin.New()
var markedBusinessLimited bool
var businessLimitedReason string
var rejectReason IngressRejectReason
var rejected bool
router.Use(func(c *gin.Context) {
c.Next()
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
rejectReason, rejected = GetIngressRejectReason(c)
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
businessLimitedReason, _ = v.(string)
}
@@ -708,6 +857,8 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
require.Equal(t, http.StatusForbidden, w.Code)
require.Contains(t, w.Body.String(), "not assigned to any group")
require.True(t, rejected)
require.Equal(t, IngressRejectGroupUnassigned, rejectReason)
require.True(t, markedBusinessLimited)
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned, businessLimitedReason)
}
@@ -822,7 +973,7 @@ func TestAPIKeyAuthIPRestrictionIncludesClientIPForBlacklistDenial(t *testing.T)
requireAPIKeyAuthError(t, w, "ACCESS_DENIED", "Access denied. Your IP is 9.9.9.9")
}
func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *testing.T) {
func TestAPIKeyAuthIPRestrictionUsesConfiguredTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
user := &service.User{
@@ -855,7 +1006,7 @@ func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *test
cfg.SetTrustForwardedIPForAPIKeyACL(true)
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
require.NoError(t, router.SetTrustedProxies(nil))
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -906,7 +1057,7 @@ func TestAPIKeyAuthIPRestrictionUsesForwardedClientIPInDenialWhenTrusted(t *test
cfg.SetTrustForwardedIPForAPIKeyACL(true)
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
router := gin.New()
require.NoError(t, router.SetTrustedProxies(nil))
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
router.GET("/t", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -24,7 +24,14 @@ func ClientRequestID() gin.HandlerFunc {
}
if v, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(v) != "" {
c.Header(clientRequestIDHeader, strings.TrimSpace(v))
var valid bool
v, valid = normalizeCorrelationID(v)
if !valid {
v = uuid.New().String()
}
c.Header(clientRequestIDHeader, v)
ctx := context.WithValue(c.Request.Context(), ctxkey.ClientRequestID, v)
c.Request = c.Request.WithContext(ctx)
c.Next()
return
}
@@ -4,6 +4,7 @@ import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
@@ -30,6 +31,23 @@ func TestClientRequestIDGeneratesAndExposesID(t *testing.T) {
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
}
func TestClientRequestIDBoundsExistingContextID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(ClientRequestID())
router.GET("/", func(c *gin.Context) {
value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
c.String(http.StatusOK, value)
})
req := httptest.NewRequest(http.MethodGet, "/", nil)
req = req.WithContext(context.WithValue(req.Context(), ctxkey.ClientRequestID, strings.Repeat("x", 200)))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Len(t, w.Body.String(), 36)
require.NotEqual(t, strings.Repeat("x", maxPersistentRequestIDBytes), w.Body.String())
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
}
func TestClientRequestIDPreservesExistingContextID(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
@@ -0,0 +1,166 @@
package middleware
import (
"math"
"net/netip"
"strconv"
"strings"
"sync/atomic"
"time"
"github.com/gin-gonic/gin"
)
// IngressRejectReason identifies expected gateway admission failures that must
// not be treated as operational request errors.
type IngressRejectReason string
const (
IngressRejectQueryAPIKeyDeprecated IngressRejectReason = "query_api_key_deprecated"
IngressRejectAPIKeyRequired IngressRejectReason = "api_key_required"
IngressRejectInvalidAPIKey IngressRejectReason = "invalid_api_key"
IngressRejectAPIKeyDisabled IngressRejectReason = "api_key_disabled"
IngressRejectIPRestricted IngressRejectReason = "ip_restricted"
IngressRejectUserInactive IngressRejectReason = "user_inactive"
IngressRejectGroupDeleted IngressRejectReason = "group_deleted"
IngressRejectGroupDisabled IngressRejectReason = "group_disabled"
IngressRejectGroupNotAllowed IngressRejectReason = "group_not_allowed"
IngressRejectGroupUnassigned IngressRejectReason = "group_unassigned"
IngressRejectInvalidAuthRateLimited IngressRejectReason = "invalid_auth_rate_limited"
IngressRejectAPIKeyAuthOverloaded IngressRejectReason = "api_key_auth_overloaded"
)
const ingressRejectReasonContextKey = "ingress_reject_reason"
type IngressRejectRecorder interface {
RecordIngressReject(reason, routeFamily, protocol, clientIP string, userID, apiKeyID int64)
}
func invalidAuthClientKey(c *gin.Context) string {
return normalizeIngressRejectIP(SecurityClientIP(c))
}
func rejectInvalidAuthAbuse(c *gin.Context, apiKeyService interface {
CheckInvalidAuthAbuse(string) (time.Duration, bool)
}) bool {
if c == nil || apiKeyService == nil {
return false
}
retry, blocked := apiKeyService.CheckInvalidAuthAbuse(invalidAuthClientKey(c))
if !blocked {
return false
}
retrySeconds := int(math.Ceil(retry.Seconds()))
if retrySeconds < 1 {
retrySeconds = 1
}
c.Header("Retry-After", strconv.Itoa(retrySeconds))
MarkIngressRejected(c, IngressRejectInvalidAuthRateLimited)
return true
}
func recordInvalidAuthFailure(c *gin.Context, apiKeyService interface {
RecordInvalidAuthFailure(string)
}) {
if c == nil || apiKeyService == nil {
return
}
apiKeyService.RecordInvalidAuthFailure(invalidAuthClientKey(c))
}
type ingressRejectRecorderHolder struct{ recorder IngressRejectRecorder }
var activeIngressRejectRecorder atomic.Pointer[ingressRejectRecorderHolder]
func SetIngressRejectRecorder(recorder IngressRejectRecorder) {
if recorder == nil {
activeIngressRejectRecorder.Store(nil)
return
}
activeIngressRejectRecorder.Store(&ingressRejectRecorderHolder{recorder: recorder})
}
// MarkIngressRejected marks a request as rejected before gateway admission.
func MarkIngressRejected(c *gin.Context, reason IngressRejectReason) {
if c == nil || reason == "" {
return
}
c.Set(ingressRejectReasonContextKey, reason)
}
// GetIngressRejectReason returns the admission rejection reason, if any.
func GetIngressRejectReason(c *gin.Context) (IngressRejectReason, bool) {
if c == nil {
return "", false
}
value, exists := c.Get(ingressRejectReasonContextKey)
if !exists {
return "", false
}
reason, ok := value.(IngressRejectReason)
return reason, ok && reason != ""
}
func recordIngressReject(c *gin.Context, reason IngressRejectReason) {
holder := activeIngressRejectRecorder.Load()
if holder == nil || holder.recorder == nil || c == nil || c.Request == nil {
return
}
routeFamily, protocol := ingressRejectRoute(c.Request.URL.Path)
clientIP := normalizeIngressRejectIP(SecurityClientIP(c))
var userID, apiKeyID int64
if apiKey, ok := GetAPIKeyFromContext(c); ok && apiKey != nil {
apiKeyID = apiKey.ID
if apiKey.User != nil {
userID = apiKey.User.ID
}
} else if apiKey, ok := GetOpsFallbackAPIKey(c); ok && apiKey != nil {
apiKeyID = apiKey.ID
if apiKey.User != nil {
userID = apiKey.User.ID
}
}
holder.recorder.RecordIngressReject(string(reason), routeFamily, protocol, clientIP, userID, apiKeyID)
}
func normalizeIngressRejectIP(raw string) string {
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
if err != nil {
return "0.0.0.0"
}
addr = addr.Unmap()
if addr.Is6() {
return netip.PrefixFrom(addr, 64).Masked().Addr().String()
}
return addr.String()
}
func ingressRejectRoute(path string) (string, string) {
path = strings.ToLower(strings.TrimSpace(path))
switch {
case strings.HasPrefix(path, "/antigravity/v1beta"):
return "antigravity", "google"
case strings.HasPrefix(path, "/v1beta"):
return "gemini", "google"
case strings.HasPrefix(path, "/backend-api/codex"):
return "codex", "openai"
case strings.HasPrefix(path, "/antigravity"):
return "antigravity", "anthropic"
case strings.Contains(path, "/messages"):
return "messages", "anthropic"
case strings.Contains(path, "/responses"):
return "responses", "openai"
case strings.Contains(path, "/chat/completions"):
return "chat_completions", "openai"
case strings.Contains(path, "/images"):
return "images", "openai"
case strings.Contains(path, "/videos"):
return "videos", "openai"
case strings.Contains(path, "/embeddings"):
return "embeddings", "openai"
case strings.Contains(path, "/models"):
return "models", "openai"
default:
return "other", "gateway"
}
}
@@ -0,0 +1,58 @@
package middleware
import (
"sync"
"time"
)
const (
ingressRejectAccessLogLimit = 20
ingressRejectAccessLogWindow = time.Second
ingressRejectDroppedSummaryPeriod = 30 * time.Second
)
type ingressRejectAccessSampler struct {
mu sync.Mutex
limit int
window time.Duration
summaryPeriod time.Duration
windowStart time.Time
emitted int
dropped uint64
lastSummary time.Time
}
func newIngressRejectAccessSampler(limit int, window, summaryPeriod time.Duration) *ingressRejectAccessSampler {
return &ingressRejectAccessSampler{limit: limit, window: window, summaryPeriod: summaryPeriod}
}
// allow applies one process-wide fixed-window budget. It stores no attacker
// dimensions, so memory remains constant even for rotating keys and addresses.
func (s *ingressRejectAccessSampler) allow(now time.Time) (allowed bool, droppedSummary uint64) {
if s == nil || s.limit <= 0 || s.window <= 0 {
return false, 0
}
s.mu.Lock()
defer s.mu.Unlock()
if s.windowStart.IsZero() || now.Sub(s.windowStart) >= s.window || now.Before(s.windowStart) {
s.windowStart = now
s.emitted = 0
}
if s.emitted < s.limit {
s.emitted++
return true, 0
}
s.dropped++
if s.summaryPeriod > 0 && (s.lastSummary.IsZero() || now.Sub(s.lastSummary) >= s.summaryPeriod) {
droppedSummary = s.dropped
s.dropped = 0
s.lastSummary = now
}
return false, droppedSummary
}
var globalIngressRejectAccessSampler = newIngressRejectAccessSampler(
ingressRejectAccessLogLimit,
ingressRejectAccessLogWindow,
ingressRejectDroppedSummaryPeriod,
)
@@ -0,0 +1,81 @@
package middleware
import (
"net/http"
"net/http/httptest"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestIngressRejectAccessSamplerConcurrentGlobalLimit(t *testing.T) {
sampler := newIngressRejectAccessSampler(10, time.Hour, time.Minute)
now := time.Now()
var allowed atomic.Int64
var wg sync.WaitGroup
for i := 0; i < 200; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if ok, _ := sampler.allow(now); ok {
allowed.Add(1)
}
}()
}
wg.Wait()
require.Equal(t, int64(10), allowed.Load())
}
func TestLoggerIngressRejectSamplingIsBoundedAndSummarySkipsOpsSink(t *testing.T) {
gin.SetMode(gin.TestMode)
original := globalIngressRejectAccessSampler
globalIngressRejectAccessSampler = newIngressRejectAccessSampler(2, time.Hour, time.Hour)
t.Cleanup(func() { globalIngressRejectAccessSampler = original })
sink := initMiddlewareTestLogger(t)
router := gin.New()
router.Use(Logger())
router.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
for i := 0; i < 20; i++ {
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/v1/messages", nil))
}
var accessEvents, summaries int
for _, event := range sink.list() {
switch event.Message {
case "http request completed":
accessEvents++
case "ingress rejection access logs dropped":
summaries++
if skipped, _ := event.Fields[logger.OpsSystemLogSkipField].(bool); !skipped {
t.Fatalf("dropped summary must skip ops system log sink")
}
}
}
require.Equal(t, 2, accessEvents)
require.Equal(t, 1, summaries)
}
func TestIngressRejectAccessSamplerDroppedSummaryIsLowFrequency(t *testing.T) {
sampler := newIngressRejectAccessSampler(1, time.Hour, time.Second)
now := time.Now()
allowed, summary := sampler.allow(now)
require.True(t, allowed)
require.Zero(t, summary)
allowed, summary = sampler.allow(now.Add(100 * time.Millisecond))
require.False(t, allowed)
require.Equal(t, uint64(1), summary)
allowed, summary = sampler.allow(now.Add(200 * time.Millisecond))
require.False(t, allowed)
require.Zero(t, summary)
allowed, summary = sampler.allow(now.Add(2 * time.Second))
require.False(t, allowed)
require.Equal(t, uint64(2), summary)
}
@@ -0,0 +1,50 @@
package middleware
import (
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type ingressRejectRecorderStub struct {
mu sync.Mutex
calls int
clientIP string
}
func (r *ingressRejectRecorderStub) RecordIngressReject(_, _, _, clientIP string, _, _ int64) {
r.mu.Lock()
defer r.mu.Unlock()
r.calls++
r.clientIP = clientIP
}
func TestNormalizeIngressRejectIP(t *testing.T) {
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::1"))
require.Equal(t, "192.0.2.4", normalizeIngressRejectIP("::ffff:192.0.2.4"))
require.Equal(t, "0.0.0.0", normalizeIngressRejectIP("not-an-ip"))
}
func TestLoggerRecordsIngressRejectOnce(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := &ingressRejectRecorderStub{}
SetIngressRejectRecorder(recorder)
t.Cleanup(func() { SetIngressRejectRecorder(nil) })
router := gin.New()
router.Use(Logger())
router.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
request := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
request.RemoteAddr = "[2001:db8:abcd:1234:ffff::1]:1234"
router.ServeHTTP(httptest.NewRecorder(), request)
recorder.mu.Lock()
require.Equal(t, 1, recorder.calls)
require.Equal(t, "2001:db8:abcd:1234::", recorder.clientIP)
recorder.mu.Unlock()
}
@@ -0,0 +1,147 @@
//go:build unit
package middleware
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func invalidAuthAbuseTestConfig(threshold int) *config.Config {
return &config.Config{
RunMode: config.RunModeSimple,
APIKeyAuth: config.APIKeyAuthCacheConfig{InvalidAbuse: config.InvalidAuthAbuseConfig{
Enabled: true, Threshold: threshold, WindowSeconds: 60, BlockSeconds: 60, Capacity: 256,
}},
}
}
func TestAPIKeyAuthInvalidAbuseReturns429BeforeRepository(t *testing.T) {
gin.SetMode(gin.TestMode)
repoCalls := 0
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
repoCalls++
return nil, service.ErrAPIKeyNotFound
}}
cfg := invalidAuthAbuseTestConfig(3)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.POST("/v1/messages", func(c *gin.Context) { c.Status(http.StatusOK) })
requests := []*http.Request{
httpRequest(t, "/v1/messages", "", ""),
httpRequest(t, "/v1/messages", "Basic malformed", ""),
httpRequest(t, "/v1/messages", "", "random-invalid-key"),
}
for _, req := range requests {
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
require.NotEqual(t, http.StatusTooManyRequests, w.Code)
}
w := httptest.NewRecorder()
r.ServeHTTP(w, httpRequest(t, "/v1/messages", "", "another-random-key"))
require.Equal(t, http.StatusTooManyRequests, w.Code)
require.Equal(t, "60", w.Header().Get("Retry-After"))
require.Contains(t, w.Body.String(), "INVALID_AUTH_RATE_LIMITED")
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
require.Equal(t, 1, repoCalls, "rate-limited request must not reach the repository")
}
func TestGoogleAPIKeyAuthInvalidAbuseReturnsProtocol429(t *testing.T) {
gin.SetMode(gin.TestMode)
repoCalls := 0
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
repoCalls++
return nil, service.ErrAPIKeyNotFound
}}
cfg := invalidAuthAbuseTestConfig(2)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
var reason IngressRejectReason
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
r.Use(APIKeyAuthGoogle(svc, cfg))
r.POST("/v1beta/models/test:generateContent", func(c *gin.Context) { c.Status(http.StatusOK) })
for _, key := range []string{"random-1", "random-2"} {
w := httptest.NewRecorder()
req := httpRequest(t, "/v1beta/models/test:generateContent", "", key)
req.Header.Del("x-api-key")
req.Header.Set("x-goog-api-key", key)
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
}
w := httptest.NewRecorder()
req := httpRequest(t, "/v1beta/models/test:generateContent", "", "random-3")
req.Header.Del("x-api-key")
req.Header.Set("x-goog-api-key", "random-3")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusTooManyRequests, w.Code)
require.Equal(t, "60", w.Header().Get("Retry-After"))
require.Contains(t, w.Body.String(), "RESOURCE_EXHAUSTED")
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
require.Equal(t, 2, repoCalls)
}
func TestInvalidAuthAbuseDoesNotCountValidOrOperationalFailures(t *testing.T) {
gin.SetMode(gin.TestMode)
user := &service.User{ID: 1, Status: service.StatusActive, Role: service.RoleUser, Balance: 1}
repo := &stubApiKeyRepo{getByKey: func(_ context.Context, key string) (*service.APIKey, error) {
switch key {
case "valid-key":
return &service.APIKey{ID: 1, UserID: 1, Key: key, Status: service.StatusActive, User: user}, nil
case "db-error":
return nil, errors.New("database unavailable")
default:
return nil, service.ErrAPIKeyNotFound
}
}}
cfg := invalidAuthAbuseTestConfig(10)
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
r := gin.New()
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
r.POST("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
for _, tc := range []struct {
key string
want int
}{{"invalid", 401}, {"valid-key", 200}, {"db-error", 500}, {"db-error", 500}} {
w := httptest.NewRecorder()
r.ServeHTTP(w, httpRequest(t, "/t", "", tc.key))
require.Equal(t, tc.want, w.Code)
}
w := httptest.NewRecorder()
req := httpRequest(t, "/t", "", "")
req.Header.Set("x-goog-api-key", "valid-key")
r.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
require.Equal(t, uint64(1), svc.InvalidAuthAbuseHealth().Recorded)
}
func TestNormalizeIngressRejectIPGroupsIPv6By64(t *testing.T) {
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"))
require.Equal(t, normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"), normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::2"))
}
func httpRequest(t *testing.T, path, authorization, apiKey string) *http.Request {
t.Helper()
req := httptest.NewRequest(http.MethodPost, path, nil)
req.RemoteAddr = "203.0.113.10:12345"
if authorization != "" {
req.Header.Set("Authorization", authorization)
}
if apiKey != "" {
req.Header.Set("x-api-key", apiKey)
}
return req
}
@@ -37,6 +37,21 @@ func Logger() gin.HandlerFunc {
accountID, hasAccountID := c.Request.Context().Value(ctxkey.AccountID).(int64)
platform, _ := c.Request.Context().Value(ctxkey.Platform).(string)
model, _ := c.Request.Context().Value(ctxkey.Model).(string)
reason, rejected := GetIngressRejectReason(c)
if rejected {
recordIngressReject(c, reason)
allowed, droppedSummary := globalIngressRejectAccessSampler.allow(endTime)
if droppedSummary > 0 {
logger.FromContext(c.Request.Context()).Info("ingress rejection access logs dropped",
zap.String("component", "http.access"),
zap.Uint64("dropped_count", droppedSummary),
zap.Bool(logger.OpsSystemLogSkipField, true),
)
}
if !allowed {
return
}
}
fields := []zap.Field{
zap.String("component", "http.access"),
@@ -47,6 +62,12 @@ func Logger() gin.HandlerFunc {
zap.String("method", method),
zap.String("path", path),
}
if rejected {
fields = append(fields,
zap.String("ingress_reject_reason", string(reason)),
zap.Bool(logger.OpsSystemLogSkipField, true),
)
}
if hasAccountID && accountID > 0 {
fields = append(fields, zap.Int64("account_id", accountID))
}
@@ -121,6 +121,7 @@ func RequireGroupAssignment(settingService *service.SettingService, writeError G
return
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned)
MarkIngressRejected(c, IngressRejectGroupUnassigned)
writeError(c, http.StatusForbidden, "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.")
c.Abort()
}
@@ -112,6 +112,26 @@ func TestRequestLogger_KeepIncomingRequestID(t *testing.T) {
}
}
func TestRequestLoggerBoundsIncomingRequestID(t *testing.T) {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(RequestLogger())
r.GET("/t", func(c *gin.Context) {
reqID, _ := c.Request.Context().Value(ctxkey.RequestID).(string)
if len(reqID) != 36 {
t.Fatalf("request_id length=%d", len(reqID))
}
c.Status(http.StatusOK)
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/t", nil)
req.Header.Set(requestIDHeader, strings.Repeat("r", 1024))
r.ServeHTTP(w, req)
if got := len(w.Header().Get(requestIDHeader)); got != 36 {
t.Fatalf("response request_id length=%d", got)
}
}
func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
@@ -180,11 +200,43 @@ func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
}
}
func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
func TestLogger_IngressRejectRemainsInStandardAccessLog(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
r := gin.New()
r.Use(Logger())
r.GET("/v1/messages", func(c *gin.Context) {
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
c.Status(http.StatusUnauthorized)
})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("status=%d", w.Code)
}
events := sink.list()
if len(events) != 1 {
t.Fatalf("events=%d, want 1", len(events))
}
if got := events[0].Fields["ingress_reject_reason"]; got != string(IngressRejectInvalidAPIKey) {
t.Fatalf("ingress_reject_reason=%v", got)
}
if got, _ := events[0].Fields[logger.OpsSystemLogSkipField].(bool); !got {
t.Fatalf("%s must be true", logger.OpsSystemLogSkipField)
}
}
func TestLogger_AccessLogUsesForwardedClientIPFromTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := initMiddlewareTestLogger(t)
r := gin.New()
if err := r.SetTrustedProxies([]string{"104.23.251.120"}); err != nil {
t.Fatalf("set trusted proxies: %v", err)
}
r.Use(Logger())
r.GET("/api/test", func(c *gin.Context) {
c.Status(http.StatusOK)
@@ -193,7 +245,7 @@ func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/test", nil)
req.RemoteAddr = "104.23.251.120:443"
req.Header.Set("CF-Connecting-IP", "203.0.113.42")
req.Header.Set("X-Forwarded-For", "203.0.113.42")
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status=%d", w.Code)
@@ -21,14 +21,15 @@ func RequestLogger() gin.HandlerFunc {
return
}
requestID := strings.TrimSpace(c.GetHeader(requestIDHeader))
if requestID == "" {
requestID, validRequestID := normalizeCorrelationID(c.GetHeader(requestIDHeader))
if !validRequestID {
requestID = uuid.NewString()
}
c.Header(requestIDHeader, requestID)
ctx := context.WithValue(c.Request.Context(), ctxkey.RequestID, requestID)
clientRequestID, _ := ctx.Value(ctxkey.ClientRequestID).(string)
clientRequestID, _ = normalizeCorrelationID(clientRequestID)
requestLogger := logger.With(
zap.String("component", "http"),
@@ -0,0 +1,30 @@
package middleware
import (
"strings"
"unicode/utf8"
)
const (
maxPersistentRequestIDBytes = 64
maxPersistentUserAgentBytes = 512
)
// normalizePersistentText bounds attacker-controlled metadata before it reaches
// logs or database columns while preserving valid UTF-8 content.
func normalizePersistentText(value string, maxBytes int) string {
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
if maxBytes <= 0 || len(value) <= maxBytes {
return value
}
value = value[:maxBytes]
for !utf8.ValidString(value) {
value = value[:len(value)-1]
}
return value
}
func normalizeCorrelationID(value string) (string, bool) {
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
return value, value != "" && len(value) <= maxPersistentRequestIDBytes
}
@@ -13,14 +13,15 @@ import (
// SessionBindingContext 全局中间件:将请求的客户端 IP 与 User-Agent 注入
// request context,供 token 签发路径(登录 / 刷新 / OAuth 回调)读取并写入会话绑定,
// 同时作为审计日志、会话绑定校验的统一客户端 IP 来源。
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关:
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
// 关闭时走 trusted_proxies 解析链,避免不可信头伪造绕过绑定。
// IP 取值与 API Key IP 限制共用 Gin trusted_proxies 解析链;旧设置开关
// 仅为配置兼容保留,不能单独使直连请求的转发头变为可信。
func SessionBindingContext(cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
userAgent := normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes)
c.Request.Header.Set("User-Agent", userAgent)
binding := &service.SessionBinding{
IP: ip.GetSecurityClientIP(c, cfg.TrustForwardedIPForAPIKeyACL()),
UserAgent: c.Request.UserAgent(),
UserAgent: userAgent,
}
c.Request = c.Request.WithContext(service.WithSessionBinding(c.Request.Context(), binding))
c.Next()
@@ -36,7 +37,7 @@ func requestSessionBinding(c *gin.Context) *service.SessionBinding {
}
return &service.SessionBinding{
IP: ip.GetTrustedClientIP(c),
UserAgent: c.Request.UserAgent(),
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
}
}
@@ -93,7 +94,7 @@ func enforceSessionBinding(
Method: c.Request.Method,
Path: path,
ClientIP: binding.IP,
UserAgent: c.Request.UserAgent(),
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
StatusCode: 401,
})
}
@@ -4,6 +4,7 @@ package middleware
import (
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -13,9 +14,7 @@ import (
"github.com/stretchr/testify/require"
)
// 反代场景:RemoteAddr 为 127.0.0.1,真实客户端 IP 在 X-Real-IP 中。
// 会话绑定注入与审计 IP 必须与 API Key IP 限制共用「信任反代传递的客户端 IP」开关语义。
func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
func TestSessionBindingContextDoesNotTrustHeadersWithoutTrustedProxy(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, tc := range []struct {
@@ -24,7 +23,7 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
wantIP string
}{
{name: "trust disabled records proxy address", trustForwarded: false, wantIP: "127.0.0.1"},
{name: "trust enabled records forwarded client IP", trustForwarded: true, wantIP: "1.2.3.4"},
{name: "legacy trust toggle cannot bypass trusted proxies", trustForwarded: true, wantIP: "127.0.0.1"},
} {
t.Run(tc.name, func(t *testing.T) {
cfg := &config.Config{}
@@ -54,6 +53,23 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
}
}
func TestSessionBindingContextBoundsPersistedUserAgent(t *testing.T) {
cfg := &config.Config{}
r := gin.New()
r.Use(SessionBindingContext(cfg))
r.GET("/t", func(c *gin.Context) {
binding := service.SessionBindingFromContext(c.Request.Context())
require.Len(t, binding.UserAgent, maxPersistentUserAgentBytes)
require.Equal(t, binding.UserAgent, c.Request.UserAgent())
c.Status(200)
})
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/t", nil)
req.Header.Set("User-Agent", strings.Repeat("u", 2048))
r.ServeHTTP(w, req)
require.Equal(t, 200, w.Code)
}
// 未经过 SessionBindingContext 注入时(异常挂载顺序/单测直调),回退 trusted_proxies 链,
// 等价于开关关闭时的历史行为。
func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
@@ -75,8 +91,6 @@ func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
require.Equal(t, "9.9.9.9", w.Body.String())
}
// requestSessionBinding 优先取注入值:开关开启时校验哈希必须基于注入的转发 IP 计算,
// 与 token 签发路径取值一致,否则同一客户端会被误判为指纹变化。
func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -84,7 +98,7 @@ func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
cfg.SetTrustForwardedIPForAPIKeyACL(true)
r := gin.New()
require.NoError(t, r.SetTrustedProxies(nil))
require.NoError(t, r.SetTrustedProxies([]string{"127.0.0.1"}))
r.Use(SessionBindingContext(cfg))
r.GET("/t", func(c *gin.Context) {
issued := &service.SessionBinding{IP: "1.2.3.4", UserAgent: "test-agent"}
+2 -1
View File
@@ -35,6 +35,7 @@ func SetupRouter(
cfg *config.Config,
redisClient *redis.Client,
) *gin.Engine {
middleware2.SetIngressRejectRecorder(opsService)
// 缓存 iframe 页面的 origin 列表,用于动态注入 CSP frame-src
var cachedFrameOrigins atomic.Pointer[[]string]
emptyOrigins := []string{}
@@ -55,7 +56,7 @@ func SetupRouter(
// 应用中间件
r.Use(middleware2.RequestLogger())
// 将客户端 IP + UA 注入 request context,供 token 签发/会话绑定/审计日志统一读取。
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关。
// IP 取值与 API Key IP 限制共用 server.trusted_proxies 信任链。
r.Use(middleware2.SessionBindingContext(cfg))
r.Use(middleware2.Logger())
r.Use(middleware2.CORS(cfg.CORS))
+5
View File
@@ -235,6 +235,11 @@ func registerOpsRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
ops.GET("/request-errors/:id/upstream-errors", h.Admin.Ops.ListRequestErrorUpstreamErrors)
ops.PUT("/request-errors/:id/resolve", h.Admin.Ops.ResolveRequestError)
// Bounded ingress-admission rejection aggregates.
ops.GET("/ingress-rejections", h.Admin.Ops.ListIngressRejects)
ops.GET("/ingress-rejections/health", h.Admin.Ops.GetIngressRejectHealth)
ops.GET("/auth-cache-invalidation/health", h.Admin.Ops.GetAuthCacheInvalidationHealth)
// Upstream errors (independent upstream failures)
ops.GET("/upstream-errors", h.Admin.Ops.ListUpstreamErrors)
ops.GET("/upstream-errors/:id", h.Admin.Ops.GetUpstreamError)
+6 -5
View File
@@ -23,6 +23,7 @@ func RegisterGatewayRoutes(
cfg *config.Config,
) {
bodyLimit := middleware.RequestBodyLimit(cfg.Gateway.MaxBodySize)
textBodyLimit := middleware.RequestBodyLimit(cfg.Gateway.TextMaxBodySize)
clientRequestID := middleware.ClientRequestID()
opsErrorLogger := handler.OpsErrorLoggerMiddleware(opsService)
endpointNorm := handler.InboundEndpointMiddleware()
@@ -165,7 +166,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)
})
@@ -177,7 +178,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{
@@ -236,7 +237,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)
})
@@ -246,7 +247,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)
})
@@ -260,7 +261,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)
})
}
}
}
@@ -76,32 +76,119 @@ func (c apiKeyAuthCacheConfig) jitterTTL(ttl time.Duration) time.Duration {
func (s *APIKeyService) initAuthCache(cfg *config.Config) {
s.authCfg = newAPIKeyAuthCacheConfig(cfg)
if !s.authCfg.l1Enabled() {
return
if s.authCfg.negativeEnabled() {
negativeSize := defaultNegativeAuthCacheSize
if s.authCfg.l1Size > 0 && s.authCfg.l1Size < negativeSize {
negativeSize = s.authCfg.l1Size
}
cache, err := ristretto.NewCache(&ristretto.Config{
NumCounters: int64(negativeSize) * 10,
MaxCost: int64(negativeSize),
BufferItems: 64,
})
if err == nil {
s.authNegativeCacheL1 = cache
}
}
cache, err := ristretto.NewCache(&ristretto.Config{
NumCounters: int64(s.authCfg.l1Size) * 10,
MaxCost: int64(s.authCfg.l1Size),
BufferItems: 64,
})
if err != nil {
return
if s.authCfg.l1Enabled() {
cache, err := ristretto.NewCache(&ristretto.Config{
NumCounters: int64(s.authCfg.l1Size) * 10,
MaxCost: int64(s.authCfg.l1Size),
BufferItems: 64,
})
if err == nil {
s.authCacheL1 = cache
}
}
s.authCacheL1 = cache
}
// StartAuthCacheInvalidationSubscriber starts the Pub/Sub subscriber for L1 cache invalidation.
// This should be called after the service is fully initialized.
func (s *APIKeyService) StartAuthCacheInvalidationSubscriber(ctx context.Context) {
if s.cache == nil || s.authCacheL1 == nil {
if s.cache == nil || (s.authCacheL1 == nil && s.authNegativeCacheL1 == nil) {
return
}
if err := s.cache.SubscribeAuthCacheInvalidation(ctx, func(cacheKey string) {
s.authCacheL1.Del(cacheKey)
}); err != nil {
// Log but don't fail - L1 cache will still work, just without cross-instance invalidation
slog.Warn("failed to start auth cache invalidation subscriber", "error", err)
s.authInvalidationStart.Do(func() {
subscriberCtx, cancel := context.WithCancel(ctx)
subscriberCtx = withAuthCacheSubscriptionReady(subscriberCtx, func() {
s.authInvalidationConnected.Store(true)
})
s.authInvalidationCancel = cancel
s.authInvalidationWG.Add(1)
go func() {
defer s.authInvalidationWG.Done()
backoff := time.Second
for {
err := s.cache.SubscribeAuthCacheInvalidation(subscriberCtx, func(cacheKey string) {
s.invalidateLocalAuthCache(cacheKey)
})
wasConnected := s.authInvalidationConnected.Swap(false)
if subscriberCtx.Err() != nil {
return
}
if wasConnected {
backoff = time.Second
}
s.authInvalidationFailures.Add(1)
if err == nil {
err = errors.New("auth cache invalidation subscription closed")
}
slog.Warn("failed to start auth cache invalidation subscriber; retrying", "error", err, "retry_in", backoff)
timer := time.NewTimer(backoff)
select {
case <-subscriberCtx.Done():
timer.Stop()
return
case <-timer.C:
}
if backoff < 30*time.Second {
backoff *= 2
if backoff > 30*time.Second {
backoff = 30 * time.Second
}
}
}
}()
})
}
func (s *APIKeyService) invalidateLocalAuthCache(cacheKey string) {
if s == nil {
return
}
if s.authCacheL1 != nil {
s.authCacheL1.Del(cacheKey)
}
if s.authNegativeCacheL1 != nil {
s.authNegativeCacheL1.Del(cacheKey)
}
}
type AuthCacheInvalidationSubscriberHealth struct {
Connected bool `json:"connected"`
Failures uint64 `json:"failures"`
}
func (s *APIKeyService) AuthCacheInvalidationSubscriberHealth() AuthCacheInvalidationSubscriberHealth {
if s == nil {
return AuthCacheInvalidationSubscriberHealth{}
}
return AuthCacheInvalidationSubscriberHealth{
Connected: s.authInvalidationConnected.Load(),
Failures: s.authInvalidationFailures.Load(),
}
}
func (s *APIKeyService) StopAuthCacheInvalidationSubscriber() {
if s == nil {
return
}
s.authInvalidationStop.Do(func() {
if s.authInvalidationCancel != nil {
s.authInvalidationCancel()
}
s.authInvalidationWG.Wait()
})
}
func (s *APIKeyService) authCacheKey(key string) string {
@@ -117,6 +204,13 @@ func (s *APIKeyService) getAuthCacheEntry(ctx context.Context, cacheKey string)
}
}
}
if s.authNegativeCacheL1 != nil {
if val, ok := s.authNegativeCacheL1.Get(cacheKey); ok {
if entry, ok := val.(*APIKeyAuthCacheEntry); ok && entry.NotFound {
return entry, true
}
}
}
if s.cache == nil || !s.authCfg.l2Enabled() {
return nil, false
}
@@ -129,13 +223,19 @@ func (s *APIKeyService) getAuthCacheEntry(ctx context.Context, cacheKey string)
}
func (s *APIKeyService) setAuthCacheL1(cacheKey string, entry *APIKeyAuthCacheEntry) {
if s.authCacheL1 == nil || entry == nil {
if entry == nil {
return
}
if entry.NotFound {
if s.authNegativeCacheL1 != nil && s.authCfg.negativeTTL > 0 {
_ = s.authNegativeCacheL1.SetWithTTL(cacheKey, entry, 1, s.authCfg.jitterTTL(s.authCfg.negativeTTL))
}
return
}
if s.authCacheL1 == nil {
return
}
ttl := s.authCfg.l1TTL
if entry.NotFound && s.authCfg.negativeTTL > 0 && s.authCfg.negativeTTL < ttl {
ttl = s.authCfg.negativeTTL
}
ttl = s.authCfg.jitterTTL(ttl)
_ = s.authCacheL1.SetWithTTL(cacheKey, entry, 1, ttl)
}
@@ -155,6 +255,9 @@ func (s *APIKeyService) deleteAuthCache(ctx context.Context, cacheKey string) {
if s.authCacheL1 != nil {
s.authCacheL1.Del(cacheKey)
}
if s.authNegativeCacheL1 != nil {
s.authNegativeCacheL1.Del(cacheKey)
}
if s.cache == nil {
return
}
@@ -164,12 +267,15 @@ func (s *APIKeyService) deleteAuthCache(ctx context.Context, cacheKey string) {
}
func (s *APIKeyService) loadAuthCacheEntry(ctx context.Context, key, cacheKey string) (*APIKeyAuthCacheEntry, error) {
apiKey, err := s.apiKeyRepo.GetByKeyForAuth(ctx, key)
apiKey, err := s.lookupAPIKeyForAuth(ctx, key)
if err != nil {
if errors.Is(err, ErrAPIKeyNotFound) {
entry := &APIKeyAuthCacheEntry{NotFound: true}
if s.authCfg.negativeEnabled() {
s.setAuthCacheEntry(ctx, cacheKey, entry, s.authCfg.negativeTTL)
// Invalid keys are attacker-controlled and high-cardinality. Keep their
// negative entries in the bounded process-local cache; do not amplify
// random-key scans into Redis writes on every instance.
s.setAuthCacheL1(cacheKey, entry)
}
return entry, nil
}
@@ -185,6 +291,30 @@ func (s *APIKeyService) loadAuthCacheEntry(ctx context.Context, key, cacheKey st
return entry, nil
}
func (s *APIKeyService) lookupAPIKeyForAuth(ctx context.Context, key string) (*APIKey, error) {
if s == nil || s.apiKeyRepo == nil {
return nil, ErrAPIKeyNotFound
}
if s.authLookupSlots == nil {
return s.apiKeyRepo.GetByKeyForAuth(ctx, key)
}
s.authLookupTotal.Add(1)
select {
case s.authLookupSlots <- struct{}{}:
s.authLookupInFlight.Add(1)
defer func() {
s.authLookupInFlight.Add(-1)
<-s.authLookupSlots
}()
case <-ctx.Done():
return nil, ctx.Err()
default:
s.authLookupRejected.Add(1)
return nil, ErrAPIKeyAuthOverloaded
}
return s.apiKeyRepo.GetByKeyForAuth(ctx, key)
}
func (s *APIKeyService) applyAuthCacheEntry(key string, entry *APIKeyAuthCacheEntry) (*APIKey, bool, error) {
if entry == nil {
return nil, false, nil
+84 -23
View File
@@ -10,6 +10,7 @@ import (
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -22,13 +23,14 @@ import (
)
var (
ErrAPIKeyNotFound = infraerrors.NotFound("API_KEY_NOT_FOUND", "api key not found")
ErrGroupNotAllowed = infraerrors.Forbidden("GROUP_NOT_ALLOWED", "user is not allowed to bind this group")
ErrAPIKeyExists = infraerrors.Conflict("API_KEY_EXISTS", "api key already exists")
ErrAPIKeyTooShort = infraerrors.BadRequest("API_KEY_TOO_SHORT", "api key must be at least 16 characters")
ErrAPIKeyInvalidChars = infraerrors.BadRequest("API_KEY_INVALID_CHARS", "api key can only contain letters, numbers, underscores, and hyphens")
ErrAPIKeyRateLimited = infraerrors.TooManyRequests("API_KEY_RATE_LIMITED", "too many failed attempts, please try again later")
ErrInvalidIPPattern = infraerrors.BadRequest("INVALID_IP_PATTERN", "invalid IP or CIDR pattern")
ErrAPIKeyNotFound = infraerrors.NotFound("API_KEY_NOT_FOUND", "api key not found")
ErrGroupNotAllowed = infraerrors.Forbidden("GROUP_NOT_ALLOWED", "user is not allowed to bind this group")
ErrAPIKeyExists = infraerrors.Conflict("API_KEY_EXISTS", "api key already exists")
ErrAPIKeyTooShort = infraerrors.BadRequest("API_KEY_TOO_SHORT", "api key must be at least 16 characters")
ErrAPIKeyInvalidChars = infraerrors.BadRequest("API_KEY_INVALID_CHARS", "api key can only contain letters, numbers, underscores, and hyphens")
ErrAPIKeyRateLimited = infraerrors.TooManyRequests("API_KEY_RATE_LIMITED", "too many failed attempts, please try again later")
ErrAPIKeyAuthOverloaded = infraerrors.ServiceUnavailable("API_KEY_AUTH_OVERLOADED", "api key authentication is temporarily overloaded")
ErrInvalidIPPattern = infraerrors.BadRequest("INVALID_IP_PATTERN", "invalid IP or CIDR pattern")
// ErrAPIKeyExpired = infraerrors.Forbidden("API_KEY_EXPIRED", "api key has expired")
ErrAPIKeyExpired = infraerrors.Forbidden("API_KEY_EXPIRED", "api key 已过期")
// ErrAPIKeyQuotaExhausted = infraerrors.TooManyRequests("API_KEY_QUOTA_EXHAUSTED", "api key quota exhausted")
@@ -41,6 +43,9 @@ var (
)
const (
MaxAPIKeyCredentialBytes = 128
defaultAuthLookupConcurrency = 64
defaultNegativeAuthCacheSize = 16384
apiKeyMaxErrorsPerHour = 20
apiKeyLastUsedMinTouch = 30 * time.Second
apiKeySortCurrentConcurrency = "current_concurrency"
@@ -58,7 +63,9 @@ type APIKeyRepository interface {
GetByKeyForAuth(ctx context.Context, key string) (*APIKey, error)
Update(ctx context.Context, key *APIKey) error
Delete(ctx context.Context, id int64) error
// DeleteWithAudit 在同一事务内先写 deleted_api_key_audits 审计、再软删除该 key。
// DeleteWithAudit keeps the legacy interface name for rolling-upgrade compatibility.
// Implementations must tombstone the key and soft-delete it atomically without
// retaining the deleted credential material.
DeleteWithAudit(ctx context.Context, id int64) error
ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error)
@@ -149,6 +156,20 @@ type APIKeyCache interface {
SubscribeAuthCacheInvalidation(ctx context.Context, handler func(cacheKey string)) error
}
type authCacheSubscriptionReadyKey struct{}
func withAuthCacheSubscriptionReady(ctx context.Context, ready func()) context.Context {
return context.WithValue(ctx, authCacheSubscriptionReadyKey{}, ready)
}
// NotifyAuthCacheSubscriptionReady lets cache implementations report that the
// server acknowledged the subscription without widening the public cache API.
func NotifyAuthCacheSubscriptionReady(ctx context.Context) {
if ready, ok := ctx.Value(authCacheSubscriptionReadyKey{}).(func()); ok && ready != nil {
ready()
}
}
// APIKeyAuthCacheInvalidator 提供认证缓存失效能力
type APIKeyAuthCacheInvalidator interface {
InvalidateAuthCacheByKey(ctx context.Context, key string)
@@ -202,20 +223,51 @@ type RateLimitCacheInvalidator interface {
}
type APIKeyService struct {
apiKeyRepo APIKeyRepository
userRepo UserRepository
groupRepo GroupRepository
userSubRepo UserSubscriptionRepository
userGroupRateRepo UserGroupRateRepository
cache APIKeyCache
rateLimitCacheInvalid RateLimitCacheInvalidator // optional: invalidate Redis rate limit cache
concurrencyService *ConcurrencyService
cfg *config.Config
authCacheL1 *ristretto.Cache
authCfg apiKeyAuthCacheConfig
authGroup singleflight.Group
lastUsedTouchL1 sync.Map // keyID -> nextAllowedAt(time.Time)
lastUsedTouchSF singleflight.Group
apiKeyRepo APIKeyRepository
userRepo UserRepository
groupRepo GroupRepository
userSubRepo UserSubscriptionRepository
userGroupRateRepo UserGroupRateRepository
cache APIKeyCache
rateLimitCacheInvalid RateLimitCacheInvalidator // optional: invalidate Redis rate limit cache
concurrencyService *ConcurrencyService
cfg *config.Config
authCacheL1 *ristretto.Cache
authNegativeCacheL1 *ristretto.Cache
authCfg apiKeyAuthCacheConfig
authGroup singleflight.Group
authLookupSlots chan struct{}
authLookupTotal atomic.Uint64
authLookupRejected atomic.Uint64
authLookupInFlight atomic.Int64
invalidAuthAbuse *invalidAuthAbuseLimiter
authInvalidationStart sync.Once
authInvalidationStop sync.Once
authInvalidationCancel context.CancelFunc
authInvalidationWG sync.WaitGroup
authInvalidationConnected atomic.Bool
authInvalidationFailures atomic.Uint64
lastUsedTouchL1 sync.Map // keyID -> nextAllowedAt(time.Time)
lastUsedTouchSF singleflight.Group
}
type APIKeyAuthLookupMetrics struct {
Total uint64 `json:"total"`
Rejected uint64 `json:"rejected"`
InFlight int64 `json:"in_flight"`
Capacity int `json:"capacity"`
}
func (s *APIKeyService) AuthLookupMetrics() APIKeyAuthLookupMetrics {
if s == nil {
return APIKeyAuthLookupMetrics{}
}
return APIKeyAuthLookupMetrics{
Total: s.authLookupTotal.Load(),
Rejected: s.authLookupRejected.Load(),
InFlight: s.authLookupInFlight.Load(),
Capacity: cap(s.authLookupSlots),
}
}
// NewAPIKeyService 创建API Key服务实例
@@ -238,6 +290,12 @@ func NewAPIKeyService(
cfg: cfg,
}
svc.initAuthCache(cfg)
lookupConcurrency := defaultAuthLookupConcurrency
if cfg != nil && cfg.APIKeyAuth.LookupConcurrency > 0 {
lookupConcurrency = cfg.APIKeyAuth.LookupConcurrency
}
svc.authLookupSlots = make(chan struct{}, lookupConcurrency)
svc.invalidAuthAbuse = newInvalidAuthAbuseLimiter(cfg)
return svc
}
@@ -581,6 +639,9 @@ func (s *APIKeyService) GetByID(ctx context.Context, id int64) (*APIKey, error)
// GetByKey 根据Key字符串获取API Key(用于认证)
func (s *APIKeyService) GetByKey(ctx context.Context, key string) (*APIKey, error) {
if len(key) == 0 || len(key) > MaxAPIKeyCredentialBytes {
return nil, ErrAPIKeyNotFound
}
cacheKey := s.authCacheKey(key)
if entry, ok := s.getAuthCacheEntry(ctx, cacheKey); ok {
@@ -622,7 +683,7 @@ func (s *APIKeyService) GetByKey(ctx context.Context, key string) (*APIKey, erro
}
}
apiKey, err := s.apiKeyRepo.GetByKeyForAuth(ctx, key)
apiKey, err := s.lookupAPIKeyForAuth(ctx, key)
if err != nil {
return nil, fmt.Errorf("get api key: %w", err)
}
@@ -5,6 +5,7 @@ package service
import (
"context"
"errors"
"strings"
"sync"
"sync/atomic"
"testing"
@@ -509,14 +510,18 @@ func TestAPIKeyService_InvalidateAuthCacheByKey(t *testing.T) {
}
func TestAPIKeyService_GetByKey_CachesNegativeOnRepoMiss(t *testing.T) {
var repoCalls atomic.Int32
cache := &authCacheStub{}
repo := &authRepoStub{
getByKeyForAuth: func(ctx context.Context, key string) (*APIKey, error) {
repoCalls.Add(1)
return nil, ErrAPIKeyNotFound
},
}
cfg := &config.Config{
APIKeyAuth: config.APIKeyAuthCacheConfig{
L1Size: 100,
L1TTLSeconds: 60,
L2TTLSeconds: 60,
NegativeTTLSeconds: 30,
},
@@ -528,7 +533,73 @@ func TestAPIKeyService_GetByKey_CachesNegativeOnRepoMiss(t *testing.T) {
_, err := svc.GetByKey(context.Background(), "missing")
require.ErrorIs(t, err, ErrAPIKeyNotFound)
require.Len(t, cache.setAuthKeys, 1)
require.Empty(t, cache.setAuthKeys, "attacker-controlled misses must not be written to Redis")
svc.authNegativeCacheL1.Wait()
_, err = svc.GetByKey(context.Background(), "missing")
require.ErrorIs(t, err, ErrAPIKeyNotFound)
require.Equal(t, int32(1), repoCalls.Load())
}
func TestAPIKeyService_GetByKeyRejectsInvalidLengthBeforeCaches(t *testing.T) {
var cacheCalls atomic.Int32
cache := &authCacheStub{getAuthCache: func(context.Context, string) (*APIKeyAuthCacheEntry, error) {
cacheCalls.Add(1)
return nil, redis.Nil
}}
repo := &authRepoStub{getByKeyForAuth: func(context.Context, string) (*APIKey, error) {
t.Fatal("invalid credential reached repository")
return nil, nil
}}
svc := NewAPIKeyService(repo, nil, nil, nil, nil, cache, &config.Config{APIKeyAuth: config.APIKeyAuthCacheConfig{L2TTLSeconds: 60}})
for _, key := range []string{"", strings.Repeat("x", MaxAPIKeyCredentialBytes+1)} {
_, err := svc.GetByKey(context.Background(), key)
require.ErrorIs(t, err, ErrAPIKeyNotFound)
}
require.Zero(t, cacheCalls.Load())
}
func TestAPIKeyService_GetByKeyAllowsMaximumLength(t *testing.T) {
key := strings.Repeat("x", MaxAPIKeyCredentialBytes)
var repoCalls atomic.Int32
repo := &authRepoStub{getByKeyForAuth: func(_ context.Context, got string) (*APIKey, error) {
repoCalls.Add(1)
require.Equal(t, key, got)
return nil, ErrAPIKeyNotFound
}}
svc := NewAPIKeyService(repo, nil, nil, nil, nil, nil, &config.Config{})
_, err := svc.GetByKey(context.Background(), key)
require.ErrorIs(t, err, ErrAPIKeyNotFound)
require.Equal(t, int32(1), repoCalls.Load())
}
func TestAPIKeyService_AuthLookupBulkheadRejectsExcessMisses(t *testing.T) {
entered := make(chan struct{})
release := make(chan struct{})
repo := &authRepoStub{getByKeyForAuth: func(context.Context, string) (*APIKey, error) {
close(entered)
<-release
return nil, ErrAPIKeyNotFound
}}
svc := NewAPIKeyService(repo, nil, nil, nil, nil, nil, &config.Config{APIKeyAuth: config.APIKeyAuthCacheConfig{LookupConcurrency: 1}})
done := make(chan error, 1)
go func() {
_, err := svc.GetByKey(context.Background(), "first")
done <- err
}()
<-entered
_, err := svc.GetByKey(context.Background(), "second")
require.ErrorIs(t, err, ErrAPIKeyAuthOverloaded)
metrics := svc.AuthLookupMetrics()
require.Equal(t, uint64(2), metrics.Total)
require.Equal(t, uint64(1), metrics.Rejected)
require.Equal(t, int64(1), metrics.InFlight)
require.Equal(t, 1, metrics.Capacity)
close(release)
require.ErrorIs(t, <-done, ErrAPIKeyNotFound)
}
func TestAPIKeyService_GetByKey_SingleflightCollapses(t *testing.T) {
@@ -0,0 +1,293 @@
package service
import (
"context"
"fmt"
"log/slog"
"math/rand/v2"
"sync"
"sync/atomic"
"time"
"github.com/google/uuid"
)
const (
authInvalidationBatchSize = 100
authInvalidationPollInterval = 500 * time.Millisecond
authInvalidationLease = 30 * time.Second
authInvalidationRedisTimeout = 2 * time.Second
authInvalidationSafetyDelay = 30 * time.Second
authInvalidationConcurrency = 16
)
type AuthCacheInvalidationEvent struct {
ID int64
CacheKey string
Attempts int
Stage int
CreatedAt time.Time
}
type AuthCacheInvalidationOutboxStats struct {
Pending int64
OldestCreatedAt *time.Time
MaxAttempts int
LastError string
}
type AuthCacheInvalidationOutboxRepository interface {
Claim(ctx context.Context, workerID string, limit int, lease time.Duration) ([]AuthCacheInvalidationEvent, error)
DeleteClaimed(ctx context.Context, id int64, workerID string) error
ScheduleSecondPass(ctx context.Context, id int64, workerID string, availableAt time.Time) error
RetryClaimed(ctx context.Context, id int64, workerID string, availableAt time.Time, lastError string) error
Stats(ctx context.Context) (AuthCacheInvalidationOutboxStats, error)
}
type AuthCacheInvalidationHealth struct {
Running bool `json:"running"`
Processed uint64 `json:"processed"`
Failures uint64 `json:"failures"`
Pending int64 `json:"pending"`
OldestLag time.Duration `json:"oldest_lag"`
LastError string `json:"last_error,omitempty"`
StatsError string `json:"stats_error,omitempty"`
// HealthySLA includes the delayed safety pass. RecoverySLA is the maximum
// convergence time after Redis becomes healthy, including capped backoff.
HealthySLA time.Duration `json:"healthy_sla"`
RecoverySLA time.Duration `json:"recovery_sla"`
MaxAttempts int `json:"max_attempts"`
}
type OpsAuthCacheInvalidationHealth struct {
Outbox AuthCacheInvalidationHealth `json:"outbox"`
Subscriber AuthCacheInvalidationSubscriberHealth `json:"subscriber"`
Lookup APIKeyAuthLookupMetrics `json:"lookup"`
InvalidAbuse InvalidAuthAbuseHealth `json:"invalid_abuse"`
}
func (s *OpsService) GetAuthCacheInvalidationHealth(ctx context.Context) OpsAuthCacheInvalidationHealth {
if s == nil {
return OpsAuthCacheInvalidationHealth{}
}
health := OpsAuthCacheInvalidationHealth{}
if s.authCacheInvalidationWorker != nil {
health.Outbox = s.authCacheInvalidationWorker.Health(ctx)
}
if s.apiKeyService != nil {
health.Subscriber = s.apiKeyService.AuthCacheInvalidationSubscriberHealth()
health.Lookup = s.apiKeyService.AuthLookupMetrics()
health.InvalidAbuse = s.apiKeyService.InvalidAuthAbuseHealth()
}
return health
}
type AuthCacheInvalidationWorker struct {
repo AuthCacheInvalidationOutboxRepository
cache APIKeyCache
local *APIKeyService
workerID string
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
start sync.Once
stop sync.Once
running atomic.Bool
processed atomic.Uint64
failures atomic.Uint64
lastError atomic.Value
}
func NewAuthCacheInvalidationWorker(repo AuthCacheInvalidationOutboxRepository, cache APIKeyCache, local ...*APIKeyService) *AuthCacheInvalidationWorker {
ctx, cancel := context.WithCancel(context.Background())
w := &AuthCacheInvalidationWorker{
repo: repo, cache: cache, workerID: uuid.NewString(), ctx: ctx, cancel: cancel,
}
if len(local) > 0 {
w.local = local[0]
}
w.lastError.Store("")
return w
}
func (w *AuthCacheInvalidationWorker) Start() {
if w == nil || w.repo == nil || w.cache == nil {
return
}
w.start.Do(func() {
w.running.Store(true)
w.wg.Add(1)
go w.run()
})
}
func (w *AuthCacheInvalidationWorker) Stop() {
if w == nil {
return
}
w.stop.Do(func() {
w.cancel()
w.wg.Wait()
w.running.Store(false)
})
}
func (w *AuthCacheInvalidationWorker) run() {
defer w.wg.Done()
defer w.running.Store(false)
ticker := time.NewTicker(authInvalidationPollInterval)
defer ticker.Stop()
for {
if err := w.processBatch(w.ctx); err != nil && w.ctx.Err() == nil {
w.recordFailure(err)
}
select {
case <-w.ctx.Done():
return
case <-ticker.C:
}
}
}
func (w *AuthCacheInvalidationWorker) processBatch(ctx context.Context) error {
events, err := w.repo.Claim(ctx, w.workerID, authInvalidationBatchSize, authInvalidationLease)
if err != nil {
return fmt.Errorf("claim auth cache invalidations: %w", err)
}
semaphore := make(chan struct{}, authInvalidationConcurrency)
var wg sync.WaitGroup
for i := range events {
select {
case <-ctx.Done():
wg.Wait()
return ctx.Err()
case semaphore <- struct{}{}:
}
wg.Add(1)
go func(event AuthCacheInvalidationEvent) {
defer wg.Done()
defer func() { <-semaphore }()
w.processEvent(ctx, event)
}(events[i])
}
wg.Wait()
return nil
}
func (w *AuthCacheInvalidationWorker) processEvent(parent context.Context, event AuthCacheInvalidationEvent) {
if w.local != nil {
w.local.invalidateLocalAuthCache(event.CacheKey)
}
ctx, cancel := context.WithTimeout(parent, authInvalidationRedisTimeout)
err := w.cache.DeleteAuthCache(ctx, event.CacheKey)
if err == nil {
err = w.cache.PublishAuthCacheInvalidation(ctx, event.CacheKey)
}
cancel()
if err != nil {
w.recordFailure(err)
retryAt := time.Now().UTC().Add(authInvalidationRetryDelay(event.Attempts + 1))
retryCtx, retryCancel := context.WithTimeout(context.Background(), 2*time.Second)
retryErr := w.repo.RetryClaimed(retryCtx, event.ID, w.workerID, retryAt, boundedAuthInvalidationError(err))
retryCancel()
if retryErr != nil {
w.recordFailure(fmt.Errorf("release failed auth invalidation %d: %w", event.ID, retryErr))
}
return
}
if event.Stage == 0 {
nextCtx, nextCancel := context.WithTimeout(context.Background(), 2*time.Second)
err = w.repo.ScheduleSecondPass(nextCtx, event.ID, w.workerID, time.Now().UTC().Add(authInvalidationSafetyDelay))
nextCancel()
if err != nil {
w.recordFailure(fmt.Errorf("schedule second auth invalidation pass %d: %w", event.ID, err))
return
}
w.processed.Add(1)
w.lastError.Store("")
return
}
ackCtx, ackCancel := context.WithTimeout(context.Background(), 2*time.Second)
err = w.repo.DeleteClaimed(ackCtx, event.ID, w.workerID)
ackCancel()
if err != nil {
w.recordFailure(fmt.Errorf("ack auth invalidation %d: %w", event.ID, err))
return
}
w.processed.Add(1)
w.lastError.Store("")
}
func authInvalidationRetryDelay(attempt int) time.Duration {
if attempt < 1 {
attempt = 1
}
if attempt > 9 {
attempt = 9
}
base := time.Second * time.Duration(1<<(attempt-1))
return time.Duration(float64(base) * (0.8 + rand.Float64()*0.4))
}
func boundedAuthInvalidationError(err error) string {
if err == nil {
return ""
}
message := err.Error()
if len(message) > 1024 {
return message[:1024]
}
return message
}
func (w *AuthCacheInvalidationWorker) recordFailure(err error) {
if err == nil {
return
}
w.failures.Add(1)
w.lastError.Store(boundedAuthInvalidationError(err))
slog.Warn("auth cache invalidation outbox processing failed", "error", err)
}
func (w *AuthCacheInvalidationWorker) Health(ctx context.Context) AuthCacheInvalidationHealth {
health := AuthCacheInvalidationHealth{
HealthySLA: authInvalidationSafetyDelay + 5*time.Second,
RecoverySLA: 6 * time.Minute,
}
if w == nil {
return health
}
health.Running = w.running.Load()
health.Processed = w.processed.Load()
health.Failures = w.failures.Load()
if value := w.lastError.Load(); value != nil {
health.LastError, _ = value.(string)
}
if w.repo == nil {
return health
}
stats, err := w.repo.Stats(ctx)
if err != nil {
health.StatsError = boundedAuthInvalidationError(err)
return health
}
health.Pending = stats.Pending
health.MaxAttempts = stats.MaxAttempts
if health.LastError == "" {
health.LastError = stats.LastError
}
if stats.OldestCreatedAt != nil {
health.OldestLag = time.Since(*stats.OldestCreatedAt)
if health.OldestLag < 0 {
health.OldestLag = 0
}
}
return health
}
func ProvideAuthCacheInvalidationWorker(repo AuthCacheInvalidationOutboxRepository, cache APIKeyCache, apiKeyService *APIKeyService) *AuthCacheInvalidationWorker {
worker := NewAuthCacheInvalidationWorker(repo, cache, apiKeyService)
worker.Start()
return worker
}
@@ -0,0 +1,281 @@
package service
import (
"context"
"errors"
"sync"
"testing"
"time"
"github.com/dgraph-io/ristretto"
"github.com/stretchr/testify/require"
)
type authInvalidationRepoStub struct {
mu sync.Mutex
events []AuthCacheInvalidationEvent
claimLimit int
scheduled []int64
deleted []int64
retried []int64
retryError string
stats AuthCacheInvalidationOutboxStats
statsErr error
}
func (r *authInvalidationRepoStub) Claim(_ context.Context, _ string, limit int, _ time.Duration) ([]AuthCacheInvalidationEvent, error) {
r.mu.Lock()
defer r.mu.Unlock()
r.claimLimit = limit
return append([]AuthCacheInvalidationEvent(nil), r.events...), nil
}
func (r *authInvalidationRepoStub) DeleteClaimed(_ context.Context, id int64, _ string) error {
r.mu.Lock()
defer r.mu.Unlock()
r.deleted = append(r.deleted, id)
return nil
}
func (r *authInvalidationRepoStub) ScheduleSecondPass(_ context.Context, id int64, _ string, _ time.Time) error {
r.mu.Lock()
defer r.mu.Unlock()
r.scheduled = append(r.scheduled, id)
return nil
}
func (r *authInvalidationRepoStub) RetryClaimed(_ context.Context, id int64, _ string, _ time.Time, lastError string) error {
r.mu.Lock()
defer r.mu.Unlock()
r.retried = append(r.retried, id)
r.retryError = lastError
return nil
}
func (r *authInvalidationRepoStub) Stats(context.Context) (AuthCacheInvalidationOutboxStats, error) {
return r.stats, r.statsErr
}
type authInvalidationCacheStub struct {
mu sync.Mutex
deleteFn func(context.Context, string) error
publishFn func(context.Context, string) error
subscribeFn func(context.Context, func(string)) error
deleted []string
published []string
}
func (*authInvalidationCacheStub) GetCreateAttemptCount(context.Context, int64) (int, error) {
return 0, nil
}
func (*authInvalidationCacheStub) IncrementCreateAttemptCount(context.Context, int64) error {
return nil
}
func (*authInvalidationCacheStub) DeleteCreateAttemptCount(context.Context, int64) error { return nil }
func (*authInvalidationCacheStub) IncrementDailyUsage(context.Context, string) error { return nil }
func (*authInvalidationCacheStub) SetDailyUsageExpiry(context.Context, string, time.Duration) error {
return nil
}
func (*authInvalidationCacheStub) GetAuthCache(context.Context, string) (*APIKeyAuthCacheEntry, error) {
return nil, errors.New("miss")
}
func (*authInvalidationCacheStub) SetAuthCache(context.Context, string, *APIKeyAuthCacheEntry, time.Duration) error {
return nil
}
func (c *authInvalidationCacheStub) DeleteAuthCache(ctx context.Context, key string) error {
c.mu.Lock()
c.deleted = append(c.deleted, key)
c.mu.Unlock()
if c.deleteFn != nil {
return c.deleteFn(ctx, key)
}
return nil
}
func (c *authInvalidationCacheStub) PublishAuthCacheInvalidation(ctx context.Context, key string) error {
c.mu.Lock()
c.published = append(c.published, key)
c.mu.Unlock()
if c.publishFn != nil {
return c.publishFn(ctx, key)
}
return nil
}
func (c *authInvalidationCacheStub) SubscribeAuthCacheInvalidation(ctx context.Context, handler func(string)) error {
if c.subscribeFn != nil {
return c.subscribeFn(ctx, handler)
}
return nil
}
func TestAuthCacheInvalidationWorker_FirstPassSchedulesSafetyPass(t *testing.T) {
repo := &authInvalidationRepoStub{}
cache := &authInvalidationCacheStub{}
worker := NewAuthCacheInvalidationWorker(repo, cache)
worker.processEvent(context.Background(), AuthCacheInvalidationEvent{ID: 7, CacheKey: "hash", Stage: 0})
require.Equal(t, []string{"hash"}, cache.deleted)
require.Equal(t, []string{"hash"}, cache.published)
require.Equal(t, []int64{7}, repo.scheduled)
require.Empty(t, repo.deleted)
}
func TestAuthCacheInvalidationWorker_SecondPassCleansEvent(t *testing.T) {
repo := &authInvalidationRepoStub{}
cache := &authInvalidationCacheStub{}
worker := NewAuthCacheInvalidationWorker(repo, cache)
worker.processEvent(context.Background(), AuthCacheInvalidationEvent{ID: 8, CacheKey: "hash", Stage: 1})
require.Equal(t, []int64{8}, repo.deleted)
require.Equal(t, uint64(1), worker.Health(context.Background()).Processed)
}
func TestAuthCacheInvalidationWorker_RetriesRedisAndPublishFailures(t *testing.T) {
for _, tc := range []struct {
name string
deleteErr error
publishErr error
published int
}{
{name: "redis down", deleteErr: errors.New("redis unavailable")},
{name: "publish failure after delete", publishErr: errors.New("publish failed"), published: 1},
} {
t.Run(tc.name, func(t *testing.T) {
repo := &authInvalidationRepoStub{}
cache := &authInvalidationCacheStub{
deleteFn: func(context.Context, string) error { return tc.deleteErr },
publishFn: func(context.Context, string) error { return tc.publishErr },
}
worker := NewAuthCacheInvalidationWorker(repo, cache)
worker.processEvent(context.Background(), AuthCacheInvalidationEvent{ID: 9, CacheKey: "hash"})
require.Equal(t, []int64{9}, repo.retried)
require.Len(t, cache.published, tc.published)
require.NotEmpty(t, repo.retryError)
require.Empty(t, repo.deleted)
require.Equal(t, uint64(1), worker.Health(context.Background()).Failures)
})
}
}
func TestAuthCacheInvalidationWorker_RedisSlowIsTimedOut(t *testing.T) {
repo := &authInvalidationRepoStub{}
cache := &authInvalidationCacheStub{deleteFn: func(ctx context.Context, _ string) error {
<-ctx.Done()
return ctx.Err()
}}
worker := NewAuthCacheInvalidationWorker(repo, cache)
started := time.Now()
worker.processEvent(context.Background(), AuthCacheInvalidationEvent{ID: 10, CacheKey: "hash"})
require.Less(t, time.Since(started), 3*time.Second)
require.Equal(t, []int64{10}, repo.retried)
require.Contains(t, repo.retryError, "deadline")
}
func TestAuthCacheInvalidationWorker_BoundedBatchAndHealth(t *testing.T) {
oldest := time.Now().Add(-time.Minute)
repo := &authInvalidationRepoStub{stats: AuthCacheInvalidationOutboxStats{
Pending: 12, OldestCreatedAt: &oldest, MaxAttempts: 4, LastError: "redis down",
}}
worker := NewAuthCacheInvalidationWorker(repo, &authInvalidationCacheStub{})
require.NoError(t, worker.processBatch(context.Background()))
require.Equal(t, authInvalidationBatchSize, repo.claimLimit)
health := worker.Health(context.Background())
require.Equal(t, int64(12), health.Pending)
require.Equal(t, 4, health.MaxAttempts)
require.Equal(t, "redis down", health.LastError)
require.GreaterOrEqual(t, health.OldestLag, time.Minute)
require.Equal(t, 35*time.Second, health.HealthySLA)
require.Equal(t, 6*time.Minute, health.RecoverySLA)
}
func TestAuthCacheInvalidationWorker_ProcessesClaimedBatchConcurrently(t *testing.T) {
events := make([]AuthCacheInvalidationEvent, 32)
for i := range events {
events[i] = AuthCacheInvalidationEvent{ID: int64(i + 1), CacheKey: "hash", Stage: 1}
}
repo := &authInvalidationRepoStub{events: events}
cache := &authInvalidationCacheStub{deleteFn: func(context.Context, string) error {
time.Sleep(100 * time.Millisecond)
return nil
}}
worker := NewAuthCacheInvalidationWorker(repo, cache)
started := time.Now()
require.NoError(t, worker.processBatch(context.Background()))
require.Less(t, time.Since(started), time.Second)
require.Len(t, repo.deleted, 32)
}
func TestAuthCacheInvalidationWorker_LifecycleIsManagedAndIdempotent(t *testing.T) {
worker := NewAuthCacheInvalidationWorker(&authInvalidationRepoStub{}, &authInvalidationCacheStub{})
worker.Start()
require.Eventually(t, func() bool { return worker.Health(context.Background()).Running }, time.Second, 10*time.Millisecond)
require.NotPanics(t, func() { worker.Stop(); worker.Stop() })
require.False(t, worker.Health(context.Background()).Running)
}
func TestAuthInvalidationRetryDelayIsBoundedAndJittered(t *testing.T) {
for attempt := 1; attempt <= 20; attempt++ {
delay := authInvalidationRetryDelay(attempt)
require.GreaterOrEqual(t, delay, 800*time.Millisecond)
require.LessOrEqual(t, delay, 308*time.Second)
}
}
func TestAuthCacheInvalidationSubscriber_RetriesInitialFailureAndStops(t *testing.T) {
ready := make(chan struct{})
var calls int
cache := &authInvalidationCacheStub{subscribeFn: func(ctx context.Context, _ func(string)) error {
calls++
if calls == 1 {
return errors.New("redis starting")
}
NotifyAuthCacheSubscriptionReady(ctx)
close(ready)
<-ctx.Done()
return ctx.Err()
}}
svc := NewAPIKeyService(nil, nil, nil, nil, nil, cache, nil)
localCache, err := ristretto.NewCache(&ristretto.Config{NumCounters: 10, MaxCost: 1, BufferItems: 64})
require.NoError(t, err)
defer localCache.Close()
svc.authNegativeCacheL1 = localCache
svc.StartAuthCacheInvalidationSubscriber(context.Background())
select {
case <-ready:
case <-time.After(2 * time.Second):
t.Fatal("subscriber did not retry")
}
require.Eventually(t, func() bool { return svc.AuthCacheInvalidationSubscriberHealth().Connected }, time.Second, 10*time.Millisecond)
require.Equal(t, uint64(1), svc.AuthCacheInvalidationSubscriberHealth().Failures)
require.NotPanics(t, func() { svc.StopAuthCacheInvalidationSubscriber(); svc.StopAuthCacheInvalidationSubscriber() })
}
func TestAuthCacheInvalidationSubscriber_ReconnectsAfterRuntimeDisconnect(t *testing.T) {
ready := make(chan int, 2)
var calls int
cache := &authInvalidationCacheStub{subscribeFn: func(ctx context.Context, _ func(string)) error {
calls++
NotifyAuthCacheSubscriptionReady(ctx)
ready <- calls
if calls == 1 {
return errors.New("connection dropped")
}
<-ctx.Done()
return ctx.Err()
}}
svc := NewAPIKeyService(nil, nil, nil, nil, nil, cache, nil)
localCache, err := ristretto.NewCache(&ristretto.Config{NumCounters: 10, MaxCost: 1, BufferItems: 64})
require.NoError(t, err)
defer localCache.Close()
svc.authNegativeCacheL1 = localCache
svc.StartAuthCacheInvalidationSubscriber(context.Background())
select {
case call := <-ready:
require.Equal(t, 1, call)
case <-time.After(time.Second):
t.Fatal("initial subscription did not start")
}
select {
case call := <-ready:
require.Equal(t, 2, call)
case <-time.After(2 * time.Second):
t.Fatal("subscriber did not reconnect after runtime disconnect")
}
require.Eventually(t, func() bool { return svc.AuthCacheInvalidationSubscriberHealth().Connected }, time.Second, 10*time.Millisecond)
require.Equal(t, uint64(1), svc.AuthCacheInvalidationSubscriberHealth().Failures)
svc.StopAuthCacheInvalidationSubscriber()
}
@@ -0,0 +1,276 @@
package service
import (
"sync"
"sync/atomic"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
)
const invalidAuthAbuseShardCount = 16
type invalidAuthAbuseEntry struct {
failures int
windowStart time.Time
blockedUntil time.Time
}
type invalidAuthAbuseShard struct {
mu sync.Mutex
entries map[string]*invalidAuthAbuseEntry
}
type invalidAuthOverflow struct {
mu sync.Mutex
failures int
windowStart time.Time
blockedUntil time.Time
}
type invalidAuthAbuseLimiter struct {
threshold int
window time.Duration
block time.Duration
capacity int64
shards [invalidAuthAbuseShardCount]invalidAuthAbuseShard
overflow invalidAuthOverflow
now func() time.Time
tracked atomic.Int64
recorded atomic.Uint64
blocked atomic.Uint64
rejected atomic.Uint64
expired atomic.Uint64
overflowed atomic.Uint64
globalBlocked atomic.Uint64
cleanupNext atomic.Int64
cleanupCursor atomic.Uint32
}
type InvalidAuthAbuseHealth struct {
Enabled bool `json:"enabled"`
Tracked int64 `json:"tracked"`
Capacity int64 `json:"capacity"`
Recorded uint64 `json:"recorded"`
Blocks uint64 `json:"blocks"`
Rejected uint64 `json:"rejected"`
Expired uint64 `json:"expired"`
Overflowed uint64 `json:"overflowed"`
GlobalBlocked uint64 `json:"global_blocked"`
}
func newInvalidAuthAbuseLimiter(cfg *config.Config) *invalidAuthAbuseLimiter {
if cfg == nil || !cfg.APIKeyAuth.InvalidAbuse.Enabled {
return nil
}
c := cfg.APIKeyAuth.InvalidAbuse
if c.Threshold <= 0 || c.WindowSeconds <= 0 || c.BlockSeconds <= 0 || c.Capacity <= 0 {
return nil
}
l := &invalidAuthAbuseLimiter{
threshold: c.Threshold,
window: time.Duration(c.WindowSeconds) * time.Second,
block: time.Duration(c.BlockSeconds) * time.Second,
capacity: int64(c.Capacity),
now: time.Now,
}
for i := range l.shards {
l.shards[i].entries = make(map[string]*invalidAuthAbuseEntry)
}
return l
}
func (s *APIKeyService) CheckInvalidAuthAbuse(clientKey string) (time.Duration, bool) {
if s == nil || s.invalidAuthAbuse == nil {
return 0, false
}
return s.invalidAuthAbuse.check(clientKey)
}
func (s *APIKeyService) RecordInvalidAuthFailure(clientKey string) {
if s == nil || s.invalidAuthAbuse == nil {
return
}
s.invalidAuthAbuse.record(clientKey)
}
func (s *APIKeyService) InvalidAuthAbuseHealth() InvalidAuthAbuseHealth {
if s == nil || s.invalidAuthAbuse == nil {
return InvalidAuthAbuseHealth{}
}
return s.invalidAuthAbuse.health()
}
func (l *invalidAuthAbuseLimiter) check(clientKey string) (time.Duration, bool) {
if l == nil || clientKey == "" {
return 0, false
}
now := l.now()
l.maybeCleanupAtCapacity(now)
shard := l.shard(clientKey)
shard.mu.Lock()
entry := shard.entries[clientKey]
if entry != nil && l.entryExpired(entry, now) {
delete(shard.entries, clientKey)
l.tracked.Add(-1)
l.expired.Add(1)
entry = nil
}
if entry != nil && entry.blockedUntil.After(now) {
retry := entry.blockedUntil.Sub(now)
shard.mu.Unlock()
l.rejected.Add(1)
return retry, true
}
shard.mu.Unlock()
if entry == nil && l.tracked.Load() >= l.capacity {
if retry, blocked := l.checkOverflow(now); blocked {
l.rejected.Add(1)
l.globalBlocked.Add(1)
return retry, true
}
}
return 0, false
}
func (l *invalidAuthAbuseLimiter) record(clientKey string) {
if l == nil || clientKey == "" {
return
}
l.recorded.Add(1)
now := l.now()
l.maybeCleanupAtCapacity(now)
shard := l.shard(clientKey)
shard.mu.Lock()
entry := shard.entries[clientKey]
if entry != nil && l.entryExpired(entry, now) {
delete(shard.entries, clientKey)
l.tracked.Add(-1)
l.expired.Add(1)
entry = nil
}
if entry == nil {
if !l.reserveEntry() {
shard.mu.Unlock()
l.recordOverflow(now)
return
}
entry = &invalidAuthAbuseEntry{windowStart: now}
shard.entries[clientKey] = entry
}
if entry.blockedUntil.After(now) {
shard.mu.Unlock()
return
}
if entry.windowStart.After(now) || !now.Before(entry.windowStart.Add(l.window)) {
entry.windowStart = now
entry.failures = 0
}
entry.failures++
if entry.failures >= l.threshold {
entry.failures = 0
entry.blockedUntil = now.Add(l.block)
entry.windowStart = entry.blockedUntil
l.blocked.Add(1)
}
shard.mu.Unlock()
}
func (l *invalidAuthAbuseLimiter) reserveEntry() bool {
for {
current := l.tracked.Load()
if current >= l.capacity {
return false
}
if l.tracked.CompareAndSwap(current, current+1) {
return true
}
}
}
func (l *invalidAuthAbuseLimiter) maybeCleanupAtCapacity(now time.Time) {
if l.tracked.Load() < l.capacity {
return
}
nowUnixNano := now.UnixNano()
for {
next := l.cleanupNext.Load()
if nowUnixNano < next {
return
}
if l.cleanupNext.CompareAndSwap(next, now.Add(100*time.Millisecond).UnixNano()) {
break
}
}
index := l.cleanupCursor.Add(1) - 1
shard := &l.shards[index%invalidAuthAbuseShardCount]
shard.mu.Lock()
for key, entry := range shard.entries {
if l.entryExpired(entry, now) {
delete(shard.entries, key)
l.tracked.Add(-1)
l.expired.Add(1)
}
}
shard.mu.Unlock()
}
func (l *invalidAuthAbuseLimiter) entryExpired(entry *invalidAuthAbuseEntry, now time.Time) bool {
return entry != nil && !entry.blockedUntil.After(now) && !entry.windowStart.After(now) && !now.Before(entry.windowStart.Add(l.window))
}
func (l *invalidAuthAbuseLimiter) shard(clientKey string) *invalidAuthAbuseShard {
const fnvOffset32 = uint32(2166136261)
const fnvPrime32 = uint32(16777619)
hash := fnvOffset32
for i := 0; i < len(clientKey); i++ {
hash ^= uint32(clientKey[i])
hash *= fnvPrime32
}
return &l.shards[hash%invalidAuthAbuseShardCount]
}
func (l *invalidAuthAbuseLimiter) recordOverflow(now time.Time) {
l.overflowed.Add(1)
l.overflow.mu.Lock()
defer l.overflow.mu.Unlock()
if l.overflow.blockedUntil.After(now) {
return
}
if l.overflow.windowStart.IsZero() || !now.Before(l.overflow.windowStart.Add(l.window)) {
l.overflow.windowStart = now
l.overflow.failures = 0
}
l.overflow.failures++
if l.overflow.failures >= l.threshold {
l.overflow.failures = 0
l.overflow.blockedUntil = now.Add(l.block)
l.overflow.windowStart = l.overflow.blockedUntil
l.blocked.Add(1)
}
}
func (l *invalidAuthAbuseLimiter) checkOverflow(now time.Time) (time.Duration, bool) {
l.overflow.mu.Lock()
defer l.overflow.mu.Unlock()
if l.overflow.blockedUntil.After(now) {
return l.overflow.blockedUntil.Sub(now), true
}
return 0, false
}
func (l *invalidAuthAbuseLimiter) health() InvalidAuthAbuseHealth {
return InvalidAuthAbuseHealth{
Enabled: true,
Tracked: l.tracked.Load(),
Capacity: l.capacity,
Recorded: l.recorded.Load(),
Blocks: l.blocked.Load(),
Rejected: l.rejected.Load(),
Expired: l.expired.Load(),
Overflowed: l.overflowed.Load(),
GlobalBlocked: l.globalBlocked.Load(),
}
}
@@ -0,0 +1,100 @@
//go:build unit
package service
import (
"fmt"
"sync"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
func newInvalidAuthLimiterForTest(threshold, capacity int) *invalidAuthAbuseLimiter {
cfg := &config.Config{APIKeyAuth: config.APIKeyAuthCacheConfig{
InvalidAbuse: config.InvalidAuthAbuseConfig{
Enabled: true, Threshold: threshold, WindowSeconds: 60, BlockSeconds: 10, Capacity: capacity,
},
}}
return newInvalidAuthAbuseLimiter(cfg)
}
func TestInvalidAuthAbuseLimiterBlocksAndExpires(t *testing.T) {
l := newInvalidAuthLimiterForTest(3, 16)
now := time.Date(2026, 7, 17, 0, 0, 0, 0, time.UTC)
l.now = func() time.Time { return now }
for range 3 {
l.record("203.0.113.1")
}
retry, blocked := l.check("203.0.113.1")
require.True(t, blocked)
require.Equal(t, 10*time.Second, retry)
now = now.Add(11 * time.Second)
_, blocked = l.check("203.0.113.1")
require.False(t, blocked)
now = now.Add(61 * time.Second)
_, blocked = l.check("203.0.113.1")
require.False(t, blocked)
require.Zero(t, l.health().Tracked)
}
func TestInvalidAuthAbuseLimiterCapacityUsesBoundedOverflowProtection(t *testing.T) {
l := newInvalidAuthLimiterForTest(2, 2)
now := time.Now()
l.now = func() time.Time { return now }
l.record("198.51.100.1")
l.record("198.51.100.2")
l.record("198.51.100.3")
l.record("198.51.100.4")
_, blocked := l.check("198.51.100.5")
require.True(t, blocked)
_, trackedBlocked := l.check("198.51.100.1")
require.False(t, trackedBlocked, "global overflow protection should spare existing tracked NATs")
health := l.health()
require.Equal(t, int64(2), health.Tracked)
require.Equal(t, uint64(2), health.Overflowed)
require.Equal(t, uint64(1), health.GlobalBlocked)
}
func TestInvalidAuthAbuseLimiterConcurrentCapacityIsBounded(t *testing.T) {
const capacity = 64
l := newInvalidAuthLimiterForTest(1000, capacity)
var wg sync.WaitGroup
for i := range 1000 {
wg.Add(1)
go func(i int) {
defer wg.Done()
l.record(fmt.Sprintf("198.51.100.%d", i))
}(i)
}
wg.Wait()
health := l.health()
require.LessOrEqual(t, health.Tracked, int64(capacity))
require.Equal(t, uint64(1000), health.Recorded)
require.Equal(t, uint64(1000-capacity), health.Overflowed)
}
func TestInvalidAuthAbuseLimiterReclaimsExpiredCapacity(t *testing.T) {
const capacity = 16
l := newInvalidAuthLimiterForTest(100, capacity)
now := time.Now()
l.now = func() time.Time { return now }
for i := range capacity {
l.record(fmt.Sprintf("source-%d", i))
}
require.Equal(t, int64(capacity), l.health().Tracked)
now = now.Add(61 * time.Second)
for i := range invalidAuthAbuseShardCount {
l.check(fmt.Sprintf("new-source-%d", i))
now = now.Add(101 * time.Millisecond)
}
require.Less(t, l.health().Tracked, int64(capacity))
l.record("fresh-source")
require.LessOrEqual(t, l.health().Tracked, int64(capacity))
}
@@ -25,19 +25,21 @@ type opsCleanupTarget struct {
}
type opsCleanupDeletedCounts struct {
errorLogs int64
alertEvents int64
systemLogs int64
logAudits int64
systemMetrics int64
hourlyPreagg int64
dailyPreagg int64
errorLogs int64
ingressRejects int64
alertEvents int64
systemLogs int64
logAudits int64
systemMetrics int64
hourlyPreagg int64
dailyPreagg int64
}
func (c opsCleanupDeletedCounts) String() string {
return fmt.Sprintf(
"error_logs=%d alert_events=%d system_logs=%d log_audits=%d system_metrics=%d hourly_preagg=%d daily_preagg=%d",
"error_logs=%d ingress_rejects=%d alert_events=%d system_logs=%d log_audits=%d system_metrics=%d hourly_preagg=%d daily_preagg=%d",
c.errorLogs,
c.ingressRejects,
c.alertEvents,
c.systemLogs,
c.logAudits,
@@ -299,6 +299,7 @@ func (s *OpsCleanupService) runCleanupOnce(ctx context.Context) (opsCleanupDelet
targets := []opsCleanupTarget{
{effective.ErrorLogRetentionDays, "ops_error_logs", "created_at", false, &out.errorLogs},
{effective.ErrorLogRetentionDays, "ops_ingress_reject_aggregates", "bucket_start", false, &out.ingressRejects},
{effective.ErrorLogRetentionDays, "ops_alert_events", "created_at", false, &out.alertEvents},
{effective.ErrorLogRetentionDays, "ops_system_logs", "created_at", false, &out.systemLogs},
{effective.ErrorLogRetentionDays, "ops_system_log_cleanup_audits", "created_at", false, &out.logAudits},
@@ -0,0 +1,451 @@
package service
import (
"context"
"log"
"strings"
"sync"
"sync/atomic"
"time"
)
const (
ingressRejectShardCount = 16
ingressRejectMaxEntries = 8192
ingressRejectMaxPendingBatches = 4
ingressRejectBucketSize = time.Minute
ingressRejectFlushInterval = 5 * time.Second
ingressRejectFlushTimeout = 5 * time.Second
)
type OpsIngressRejectAggregate struct {
ID int64 `json:"id"`
BucketStart time.Time `json:"bucket_start"`
RejectReason string `json:"reject_reason"`
RouteFamily string `json:"route_family"`
Protocol string `json:"protocol"`
ClientIP string `json:"client_ip"`
UserID *int64 `json:"user_id,omitempty"`
APIKeyID *int64 `json:"api_key_id,omitempty"`
RequestCount int64 `json:"request_count"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
}
type OpsIngressRejectFilter struct {
StartTime *time.Time
EndTime *time.Time
RejectReason string
RouteFamily string
Protocol string
ClientIP string
UserID *int64
APIKeyID *int64
Page int
PageSize int
}
type OpsIngressRejectList struct {
Items []*OpsIngressRejectAggregate `json:"items"`
Total int `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
}
type OpsIngressRejectHealth struct {
Cardinality int64 `json:"cardinality"`
Capacity int `json:"capacity"`
PendingBatches int `json:"pending_batches"`
PendingRows int `json:"pending_rows"`
Overflowed uint64 `json:"overflowed_count"`
Dropped uint64 `json:"dropped_count"`
Flushed uint64 `json:"flushed_request_count"`
FlushFailures uint64 `json:"flush_failure_count"`
Accepting bool `json:"accepting"`
LastError string `json:"last_error,omitempty"`
}
type OpsIngressRejectRepository interface {
BatchUpsertIngressRejects(ctx context.Context, items []*OpsIngressRejectAggregate) error
ListIngressRejects(ctx context.Context, filter *OpsIngressRejectFilter) (*OpsIngressRejectList, error)
}
type ingressRejectKey struct {
reason string
routeFamily string
protocol string
clientIP string
userID int64
apiKeyID int64
}
type ingressRejectShard struct {
mu sync.Mutex
items map[ingressRejectKey]*OpsIngressRejectAggregate
}
type OpsIngressRejectAggregator struct {
repo OpsIngressRejectRepository
shards [ingressRejectShardCount]ingressRejectShard
recordMu sync.RWMutex
snapshotMu sync.Mutex
bucket atomic.Int64
cardinality atomic.Int64
overflowMu sync.Mutex
overflow *OpsIngressRejectAggregate
pendingMu sync.Mutex
pending [][]*OpsIngressRejectAggregate
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
flushCh chan struct{}
started atomic.Bool
accepting atomic.Bool
stopOnce sync.Once
overflowed atomic.Uint64
dropped atomic.Uint64
flushed atomic.Uint64
flushFailures atomic.Uint64
lastError atomic.Value
}
func NewOpsIngressRejectAggregator(repo OpsIngressRejectRepository) *OpsIngressRejectAggregator {
ctx, cancel := context.WithCancel(context.Background())
a := &OpsIngressRejectAggregator{
repo: repo, ctx: ctx, cancel: cancel, flushCh: make(chan struct{}, 1),
}
for i := range a.shards {
a.shards[i].items = make(map[ingressRejectKey]*OpsIngressRejectAggregate)
}
a.lastError.Store("")
return a
}
func (a *OpsIngressRejectAggregator) Start() {
if a == nil || a.repo == nil || !a.started.CompareAndSwap(false, true) {
return
}
a.accepting.Store(true)
a.wg.Add(1)
go func() {
defer a.wg.Done()
ticker := time.NewTicker(ingressRejectFlushInterval)
defer ticker.Stop()
for {
select {
case <-a.ctx.Done():
return
case <-ticker.C:
a.snapshotAndEnqueue(false)
a.flushPending()
case <-a.flushCh:
a.flushPending()
}
}
}()
}
func (a *OpsIngressRejectAggregator) Stop() {
if a == nil {
return
}
a.stopOnce.Do(func() {
a.accepting.Store(false)
a.recordMu.Lock()
a.recordMu.Unlock()
a.cancel()
a.wg.Wait()
a.snapshotAndEnqueue(false)
a.flushPending()
})
}
func (a *OpsIngressRejectAggregator) RecordIngressReject(reason, routeFamily, protocol, clientIP string, userID, apiKeyID int64) {
if a == nil || a.repo == nil || !a.accepting.Load() {
return
}
now := time.Now().UTC()
bucket := now.Truncate(ingressRejectBucketSize)
key := ingressRejectKey{
reason: boundedDimension(reason, "unknown"), routeFamily: boundedDimension(routeFamily, "other"),
protocol: boundedDimension(protocol, "other"), clientIP: boundedDimension(clientIP, "0.0.0.0"),
userID: userID, apiKeyID: apiKeyID,
}
for {
a.ensureBucket(bucket)
a.recordMu.RLock()
if !a.accepting.Load() {
a.recordMu.RUnlock()
return
}
bucketUnix := a.bucket.Load()
if bucketUnix != bucket.Unix() {
a.recordMu.RUnlock()
continue
}
shard := &a.shards[ingressRejectHash(key)%ingressRejectShardCount]
shard.mu.Lock()
if item := shard.items[key]; item != nil {
item.RequestCount++
item.LastSeen = now
shard.mu.Unlock()
a.recordMu.RUnlock()
return
}
if a.reserveDimension() {
shard.items[key] = aggregateFromKey(key, bucket, now)
shard.mu.Unlock()
a.recordMu.RUnlock()
return
}
shard.mu.Unlock()
a.recordOverflow(bucket, now)
a.recordMu.RUnlock()
return
}
}
func (a *OpsIngressRejectAggregator) ensureBucket(bucket time.Time) {
if a.bucket.Load() == bucket.Unix() {
return
}
a.snapshotMu.Lock()
defer a.snapshotMu.Unlock()
a.recordMu.Lock()
defer a.recordMu.Unlock()
if a.bucket.Load() == bucket.Unix() {
return
}
items := a.snapshotLocked(true)
a.enqueue(items)
a.bucket.Store(bucket.Unix())
a.cardinality.Store(0)
select {
case a.flushCh <- struct{}{}:
default:
}
}
func (a *OpsIngressRejectAggregator) reserveDimension() bool {
for {
current := a.cardinality.Load()
if current >= ingressRejectMaxEntries-1 {
return false
}
if a.cardinality.CompareAndSwap(current, current+1) {
return true
}
}
}
func (a *OpsIngressRejectAggregator) recordOverflow(bucket, now time.Time) {
a.overflowed.Add(1)
a.overflowMu.Lock()
defer a.overflowMu.Unlock()
if a.overflow == nil || !a.overflow.BucketStart.Equal(bucket) {
a.overflow = &OpsIngressRejectAggregate{
BucketStart: bucket, RejectReason: "other", RouteFamily: "other", Protocol: "other",
ClientIP: "0.0.0.0", RequestCount: 1, FirstSeen: now, LastSeen: now,
}
return
}
a.overflow.RequestCount++
a.overflow.LastSeen = now
}
func (a *OpsIngressRejectAggregator) snapshotAndEnqueue(reset bool) {
if a == nil || a.repo == nil {
return
}
a.snapshotMu.Lock()
items := a.snapshotLocked(reset)
a.snapshotMu.Unlock()
a.enqueue(items)
}
// snapshotLocked captures counter deltas. When reset is false the dimension keys
// remain resident for the whole minute, so periodic flushes cannot reset the budget.
func (a *OpsIngressRejectAggregator) snapshotLocked(reset bool) []*OpsIngressRejectAggregate {
items := make([]*OpsIngressRejectAggregate, 0)
for i := range a.shards {
shard := &a.shards[i]
shard.mu.Lock()
for key, item := range shard.items {
if item.RequestCount > 0 {
copyItem := *item
items = append(items, &copyItem)
item.RequestCount = 0
}
if reset {
delete(shard.items, key)
}
}
shard.mu.Unlock()
}
a.overflowMu.Lock()
if a.overflow != nil && a.overflow.RequestCount > 0 {
copyItem := *a.overflow
items = append(items, &copyItem)
a.overflow.RequestCount = 0
}
if reset {
a.overflow = nil
}
a.overflowMu.Unlock()
return items
}
func (a *OpsIngressRejectAggregator) enqueue(items []*OpsIngressRejectAggregate) {
if len(items) == 0 {
return
}
a.pendingMu.Lock()
defer a.pendingMu.Unlock()
if len(a.pending) >= ingressRejectMaxPendingBatches {
a.dropped.Add(batchRequestCount(items))
return
}
a.pending = append(a.pending, items)
}
func (a *OpsIngressRejectAggregator) flushPending() {
for {
a.pendingMu.Lock()
if len(a.pending) == 0 {
a.pendingMu.Unlock()
return
}
batch := a.pending[0]
a.pendingMu.Unlock()
ctx, cancel := context.WithTimeout(context.Background(), ingressRejectFlushTimeout)
err := a.repo.BatchUpsertIngressRejects(ctx, batch)
cancel()
if err != nil {
a.flushFailures.Add(1)
a.lastError.Store(err.Error())
log.Printf("[IngressRejectAggregator] flush failed: %v", err)
return
}
a.pendingMu.Lock()
if len(a.pending) > 0 {
a.pending = a.pending[1:]
}
a.pendingMu.Unlock()
a.flushed.Add(batchRequestCount(batch))
a.lastError.Store("")
}
}
func batchRequestCount(items []*OpsIngressRejectAggregate) uint64 {
var count uint64
for _, item := range items {
if item != nil && item.RequestCount > 0 {
count += uint64(item.RequestCount)
}
}
return count
}
func (a *OpsIngressRejectAggregator) Health() OpsIngressRejectHealth {
h := OpsIngressRejectHealth{Capacity: ingressRejectMaxEntries}
if a == nil {
return h
}
h.Cardinality = a.cardinality.Load()
a.overflowMu.Lock()
if a.overflow != nil {
h.Cardinality++
}
a.overflowMu.Unlock()
a.pendingMu.Lock()
h.PendingBatches = len(a.pending)
for _, batch := range a.pending {
h.PendingRows += len(batch)
}
a.pendingMu.Unlock()
h.Overflowed = a.overflowed.Load()
h.Dropped = a.dropped.Load()
h.Flushed = a.flushed.Load()
h.FlushFailures = a.flushFailures.Load()
h.Accepting = a.accepting.Load()
if v := a.lastError.Load(); v != nil {
h.LastError, _ = v.(string)
}
return h
}
func boundedDimension(value, fallback string) string {
value = strings.TrimSpace(value)
if value == "" {
return fallback
}
if len(value) > 64 {
return value[:64]
}
return value
}
func ingressRejectHash(k ingressRejectKey) int {
h := uint64(1469598103934665603)
for _, value := range []string{k.reason, k.routeFamily, k.protocol, k.clientIP} {
for i := 0; i < len(value); i++ {
h ^= uint64(value[i])
h *= 1099511628211
}
}
h ^= uint64(k.userID)
h *= 1099511628211
h ^= uint64(k.apiKeyID)
return int(h & 0x7fffffff)
}
func aggregateFromKey(k ingressRejectKey, bucket, now time.Time) *OpsIngressRejectAggregate {
item := &OpsIngressRejectAggregate{
BucketStart: bucket, RejectReason: k.reason, RouteFamily: k.routeFamily, Protocol: k.protocol,
ClientIP: k.clientIP, RequestCount: 1, FirstSeen: now, LastSeen: now,
}
if k.userID > 0 {
value := k.userID
item.UserID = &value
}
if k.apiKeyID > 0 {
value := k.apiKeyID
item.APIKeyID = &value
}
return item
}
func (s *OpsService) SetIngressRejectAggregator(a *OpsIngressRejectAggregator) {
if s != nil {
s.ingressRejectAggregator = a
}
}
func (s *OpsService) RecordIngressReject(reason, routeFamily, protocol, clientIP string, userID, apiKeyID int64) {
if s != nil && s.ingressRejectAggregator != nil {
s.ingressRejectAggregator.RecordIngressReject(reason, routeFamily, protocol, clientIP, userID, apiKeyID)
}
}
func (s *OpsService) GetIngressRejectHealth() OpsIngressRejectHealth {
if s == nil || s.ingressRejectAggregator == nil {
return OpsIngressRejectHealth{Capacity: ingressRejectMaxEntries}
}
return s.ingressRejectAggregator.Health()
}
func (s *OpsService) ListIngressRejects(ctx context.Context, filter *OpsIngressRejectFilter) (*OpsIngressRejectList, error) {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return nil, err
}
repo, ok := s.opsRepo.(OpsIngressRejectRepository)
if !ok {
return &OpsIngressRejectList{Items: []*OpsIngressRejectAggregate{}, Page: 1, PageSize: 50}, nil
}
return repo.ListIngressRejects(ctx, filter)
}
@@ -0,0 +1,115 @@
package service
import (
"context"
"errors"
"fmt"
"sync"
"testing"
"github.com/stretchr/testify/require"
)
type ingressRejectRepoStub struct {
mu sync.Mutex
failCount int
calls int
requests int64
}
func (r *ingressRejectRepoStub) BatchUpsertIngressRejects(_ context.Context, items []*OpsIngressRejectAggregate) error {
r.mu.Lock()
defer r.mu.Unlock()
r.calls++
if r.failCount > 0 {
r.failCount--
return errors.New("temporary database failure")
}
for _, item := range items {
if item != nil {
r.requests += item.RequestCount
}
}
return nil
}
func (r *ingressRejectRepoStub) ListIngressRejects(context.Context, *OpsIngressRejectFilter) (*OpsIngressRejectList, error) {
return &OpsIngressRejectList{}, nil
}
func TestOpsIngressRejectAggregatorUsesGlobalBucketCapacity(t *testing.T) {
repo := &ingressRejectRepoStub{}
a := NewOpsIngressRejectAggregator(repo)
a.Start()
// Concentrating all dimensions in one shard must not waste capacity in the others.
inserted := 0
for i := 0; inserted < 600; i++ {
ip := fmt.Sprintf("target-%d", i)
key := ingressRejectKey{reason: "invalid_api_key", routeFamily: "messages", protocol: "anthropic", clientIP: ip}
if ingressRejectHash(key)%ingressRejectShardCount == 0 {
a.RecordIngressReject(key.reason, key.routeFamily, key.protocol, ip, 0, 0)
inserted++
}
}
require.Equal(t, int64(600), a.Health().Cardinality)
for i := 0; i < ingressRejectMaxEntries; i++ {
a.RecordIngressReject("invalid_api_key", "messages", "anthropic", fmt.Sprintf("rotating-%d", i), 0, 0)
}
health := a.Health()
require.Equal(t, int64(ingressRejectMaxEntries), health.Cardinality)
require.Greater(t, health.Overflowed, uint64(0))
a.snapshotAndEnqueue(false)
require.Equal(t, int64(ingressRejectMaxEntries), a.Health().Cardinality, "periodic flush must retain the minute budget")
a.Stop()
}
func TestOpsIngressRejectAggregatorConcurrentCountAndStopFlush(t *testing.T) {
repo := &ingressRejectRepoStub{}
a := NewOpsIngressRejectAggregator(repo)
a.Start()
const goroutines = 32
const perGoroutine = 200
var wg sync.WaitGroup
for i := 0; i < goroutines; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for j := 0; j < perGoroutine; j++ {
a.RecordIngressReject("invalid_api_key", "responses", "openai", "192.0.2.10", 0, 0)
}
}()
}
wg.Wait()
a.Stop()
repo.mu.Lock()
require.Equal(t, int64(goroutines*perGoroutine), repo.requests)
repo.mu.Unlock()
require.False(t, a.Health().Accepting)
a.RecordIngressReject("invalid_api_key", "responses", "openai", "192.0.2.10", 0, 0)
}
func TestOpsIngressRejectAggregatorRetriesBoundedPendingBatch(t *testing.T) {
repo := &ingressRejectRepoStub{failCount: 1}
a := NewOpsIngressRejectAggregator(repo)
a.accepting.Store(true)
a.RecordIngressReject("group_deleted", "messages", "anthropic", "192.0.2.20", 1, 2)
a.snapshotAndEnqueue(false)
a.flushPending()
health := a.Health()
require.Equal(t, 1, health.PendingBatches)
require.Equal(t, uint64(1), health.FlushFailures)
require.Equal(t, uint64(0), health.Dropped)
a.flushPending()
require.Equal(t, 0, a.Health().PendingBatches)
a.Stop()
repo.mu.Lock()
require.Equal(t, int64(1), repo.requests)
require.GreaterOrEqual(t, repo.calls, 2)
repo.mu.Unlock()
}
@@ -11,12 +11,14 @@ import (
)
type runtimeSettingRepoStub struct {
values map[string]string
deleted map[string]bool
setCalls int
getValueFn func(key string) (string, error)
setFn func(key, value string) error
deleteFn func(key string) error
values map[string]string
deleted map[string]bool
setCalls int
getValueCalls int
getMultipleCalls int
getValueFn func(key string) (string, error)
setFn func(key, value string) error
deleteFn func(key string) error
}
func newRuntimeSettingRepoStub() *runtimeSettingRepoStub {
@@ -35,6 +37,7 @@ func (s *runtimeSettingRepoStub) Get(ctx context.Context, key string) (*Setting,
}
func (s *runtimeSettingRepoStub) GetValue(_ context.Context, key string) (string, error) {
s.getValueCalls++
if s.getValueFn != nil {
return s.getValueFn(key)
}
@@ -57,6 +60,7 @@ func (s *runtimeSettingRepoStub) Set(_ context.Context, key, value string) error
}
func (s *runtimeSettingRepoStub) GetMultiple(_ context.Context, keys []string) (map[string]string, error) {
s.getMultipleCalls++
out := make(map[string]string, len(keys))
for _, key := range keys {
if value, ok := s.values[key]; ok {
+1 -16
View File
@@ -74,11 +74,6 @@ type OpsErrorLog struct {
// 关联 api_key 名称(LEFT JOIN api_keys 取得;软删只覆盖 key 列,name 保留,故已删 key 仍有原名)。
APIKeyName string `json:"api_key_name,omitempty"`
APIKeyDeleted bool `json:"api_key_deleted,omitempty"`
// 已删除 KEY 所有者(INVALID_API_KEY 且该 key 曾存在时的归因快照)。
// 认证失败行 user_id 为空,列表用户列以此回退显示所有者。
DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"`
DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"`
}
type OpsErrorLogDetail struct {
@@ -102,12 +97,7 @@ type OpsErrorLogDetail struct {
// vNext metric semantics
IsBusinessLimited bool `json:"is_business_limited"`
// Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted).
// OwnerUserID/OwnerEmail 已上移到 OpsErrorLog(列表用户列回退需要)。
AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"`
DeletedKeyName string `json:"deleted_key_name,omitempty"`
// Bound (non-deleted) key prefix, snapshotted at error time; mutually exclusive with AttemptedKeyPrefix.
// Bound (non-deleted) key prefix, snapshotted at error time.
APIKeyPrefix string `json:"api_key_prefix,omitempty"`
}
@@ -137,11 +127,6 @@ type OpsErrorLogFilter struct {
UserID *int64
APIKeyID *int64
// MatchDeletedKeyOwner: 用户侧专用。UserID 设置且为 true 时,归属从 user_id=UserID
// 放宽为 (user_id=UserID OR deleted_key_owner_user_id=UserID),使原所有者能看到
// 自己「已删除 key 认证失败」的记录。admin 路径不设此开关 → 行为不变。
MatchDeletedKeyOwner bool
// Model matches against requested_model first, then model.
Model string
// ModelFuzzy 为 true 时 Model 走 ILIKE 模糊匹配(仅用户端启用);false(默认)保持精确 =,管理端语义不变。
+1 -14
View File
@@ -10,8 +10,6 @@ type OpsRepository interface {
BatchInsertErrorLogs(ctx context.Context, inputs []*OpsInsertErrorLogInput) (int64, error)
ListErrorLogs(ctx context.Context, filter *OpsErrorLogFilter) (*OpsErrorLogList, error)
GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error)
// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计;未命中返回 (nil, nil)。
LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
ListRequestDetails(ctx context.Context, filter *OpsRequestDetailFilter) ([]*OpsRequestDetail, int64, error)
BatchInsertSystemLogs(ctx context.Context, inputs []*OpsInsertSystemLogInput) (int64, error)
ListSystemLogs(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
@@ -63,12 +61,6 @@ type OpsRepository interface {
GetLatestDailyBucketDate(ctx context.Context) (time.Time, bool, error)
}
// DeletedKeyAuditResult 是按明文 key 反查 deleted_api_key_audits 的结果。
type DeletedKeyAuditResult struct {
UserID int64
KeyName string
}
type OpsInsertErrorLogInput struct {
RequestID string
ClientRequestID string
@@ -127,12 +119,7 @@ type OpsInsertErrorLogInput struct {
CreatedAt time.Time
// 已删除 key 归因(仅 INVALID_API_KEY 认证失败时可能非空)
AttemptedKeyPrefix string // 提交 key 的脱敏前缀(前 8 位)
DeletedKeyOwnerUserID *int64 // 反查命中的原所有者 user_id
DeletedKeyName string // 反查命中的 key 名称
// 有效(未删除)key 报错时快照的 key 脱敏前缀(前 8 位);与 AttemptedKeyPrefix 互斥。
// 有效(未删除)key 报错时快照的 key 脱敏前缀(前 8 位)。
// 落库快照而非读时 JOIN:key 之后被删(key 列被 tombstone 覆盖)仍保留当时前缀。
APIKeyPrefix string
}
@@ -0,0 +1,49 @@
package service
import (
"strings"
"testing"
)
func TestSanitizeOpsUpstreamErrorsForQueueBoundsAndRedacts(t *testing.T) {
entry := &OpsInsertErrorLogInput{}
for i := 0; i < 20; i++ {
entry.UpstreamErrors = append(entry.UpstreamErrors, &OpsUpstreamErrorEvent{
Platform: strings.Repeat("p", 100),
AccountName: strings.Repeat("a", 300),
UpstreamStatusCode: 500,
UpstreamURL: strings.Repeat("u", 3000),
UpstreamResponseBody: `{"authorization":"Bearer secret","message":"` + strings.Repeat("x", 10_000) + `"}`,
Message: strings.Repeat("m", 3000),
Detail: `{"api_key":"secret","detail":"` + strings.Repeat("y", 10_000) + `"}`,
})
}
if err := SanitizeOpsUpstreamErrorsForQueue(entry); err != nil {
t.Fatal(err)
}
if entry.UpstreamErrors != nil {
t.Fatal("raw upstream event slice must be released before queueing")
}
if entry.UpstreamErrorsJSON == nil {
t.Fatal("sanitized upstream event JSON is missing")
}
events, err := ParseOpsUpstreamErrors(*entry.UpstreamErrorsJSON)
if err != nil {
t.Fatal(err)
}
if len(events) != 16 {
t.Fatalf("event count = %d, want 16", len(events))
}
for _, event := range events {
if len(event.Platform) > 32 || len(event.AccountName) > 128 || len(event.UpstreamURL) > 2048 || len(event.Message) > 2048 {
t.Fatalf("event fields were not bounded: %+v", event)
}
if len(event.UpstreamResponseBody) > OpsErrorLogQueueBodyMaxBytes || len(event.Detail) > OpsErrorLogQueueBodyMaxBytes {
t.Fatal("event body/detail exceeded queue limit")
}
if strings.Contains(event.UpstreamResponseBody, "Bearer secret") || strings.Contains(event.Detail, `"secret"`) {
t.Fatal("credential material was not redacted")
}
}
}
@@ -13,7 +13,6 @@ type opsRepoMock struct {
ListSystemLogsFn func(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
DeleteSystemLogsFn func(ctx context.Context, filter *OpsSystemLogCleanupFilter) (int64, error)
InsertSystemLogCleanupAuditFn func(ctx context.Context, input *OpsSystemLogCleanupAudit) error
LookupDeletedKeyAuditFn func(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
}
func (m *opsRepoMock) InsertErrorLog(ctx context.Context, input *OpsInsertErrorLogInput) (int64, error) {
@@ -190,11 +189,4 @@ func (m *opsRepoMock) GetLatestDailyBucketDate(ctx context.Context) (time.Time,
return time.Time{}, false, nil
}
func (m *opsRepoMock) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
if m.LookupDeletedKeyAuditFn != nil {
return m.LookupDeletedKeyAuditFn(ctx, key)
}
return nil, nil
}
var _ OpsRepository = (*opsRepoMock)(nil)
@@ -0,0 +1,186 @@
package service
import (
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"time"
)
var (
benchmarkOpsMonitoringEnabled bool
benchmarkOpsAdvancedSettings OpsAdvancedSettings
)
type opsRuntimeRefreshRepo struct {
SettingRepository
mu sync.RWMutex
values map[string]string
fail atomic.Bool
calls atomic.Int64
}
func (r *opsRuntimeRefreshRepo) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) {
r.calls.Add(1)
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
}
if r.fail.Load() {
return nil, errors.New("settings unavailable")
}
r.mu.RLock()
defer r.mu.RUnlock()
out := make(map[string]string, len(keys))
for _, key := range keys {
if value, ok := r.values[key]; ok {
out[key] = value
}
}
return out, nil
}
func (r *opsRuntimeRefreshRepo) set(key, value string) {
r.mu.Lock()
r.values[key] = value
r.mu.Unlock()
}
func waitForOpsRefresh(t *testing.T, timeout time.Duration, condition func() bool) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if condition() {
return
}
time.Sleep(time.Millisecond)
}
t.Fatal("condition not satisfied before timeout")
}
func TestOpsRuntimeSettingsSnapshotLoadsOnceAndServesHotPath(t *testing.T) {
repo := newRuntimeSettingRepoStub()
repo.values[SettingKeyOpsMonitoringEnabled] = "false"
repo.values[SettingKeyOpsAdvancedSettings] = `{"ignore_context_canceled":false,"auto_refresh_interval_seconds":45}`
svc := &OpsService{settingRepo: repo}
svc.initRuntimeSettings(context.Background())
if repo.getMultipleCalls != 1 {
t.Fatalf("startup GetMultiple calls = %d, want 1", repo.getMultipleCalls)
}
for range 1000 {
if svc.IsMonitoringEnabled(context.Background()) {
t.Fatal("monitoring enabled, want false")
}
cfg, err := svc.GetOpsAdvancedSettings(context.Background())
if err != nil {
t.Fatalf("GetOpsAdvancedSettings() error = %v", err)
}
if cfg.AutoRefreshIntervalSec != 45 {
t.Fatalf("AutoRefreshIntervalSec = %d, want 45", cfg.AutoRefreshIntervalSec)
}
}
if repo.getValueCalls != 0 || repo.getMultipleCalls != 1 {
t.Fatalf("hot path touched repository: get=%d get_multiple=%d", repo.getValueCalls, repo.getMultipleCalls)
}
}
func TestOpsRuntimeSettingsAdministrativeUpdatesAreImmediatelyVisible(t *testing.T) {
svc := &OpsService{}
svc.initRuntimeSettings(context.Background())
svc.SetMonitoringEnabled(false)
if svc.IsMonitoringEnabled(context.Background()) {
t.Fatal("monitoring update was not visible")
}
cfg := defaultOpsAdvancedSettings()
cfg.IgnoreNoAvailableAccounts = true
svc.storeAdvancedSettingsSnapshot(cfg)
got, err := svc.GetOpsAdvancedSettings(context.Background())
if err != nil {
t.Fatalf("GetOpsAdvancedSettings() error = %v", err)
}
if !got.IgnoreNoAvailableAccounts {
t.Fatal("advanced settings update was not visible")
}
if svc.IsMonitoringEnabled(context.Background()) {
t.Fatal("advanced update overwrote monitoring setting")
}
}
func TestOpsRuntimeSettingsBackgroundRefreshConverges(t *testing.T) {
repo := &opsRuntimeRefreshRepo{values: map[string]string{SettingKeyOpsMonitoringEnabled: "false"}}
svc := &OpsService{settingRepo: repo}
svc.initRuntimeSettings(context.Background())
if svc.IsMonitoringEnabled(context.Background()) {
t.Fatal("initial monitoring state = true, want false")
}
repo.set(SettingKeyOpsMonitoringEnabled, "true")
svc.startRuntimeSettingsRefresh(context.Background(), 5*time.Millisecond, 0, 50*time.Millisecond)
t.Cleanup(svc.StopRuntimeSettingsRefresh)
waitForOpsRefresh(t, time.Second, func() bool {
return svc.IsMonitoringEnabled(context.Background()) && svc.RuntimeSettingsRefreshHealth().SuccessTotal > 0
})
}
func TestOpsRuntimeSettingsRefreshFailuresKeepLastKnownGoodSnapshot(t *testing.T) {
repo := &opsRuntimeRefreshRepo{values: map[string]string{
SettingKeyOpsMonitoringEnabled: "false",
SettingKeyOpsAdvancedSettings: `{"ignore_no_available_accounts":true}`,
}}
svc := &OpsService{settingRepo: repo}
svc.initRuntimeSettings(context.Background())
repo.fail.Store(true)
svc.startRuntimeSettingsRefresh(context.Background(), 5*time.Millisecond, 0, 50*time.Millisecond)
t.Cleanup(svc.StopRuntimeSettingsRefresh)
waitForOpsRefresh(t, time.Second, func() bool {
return svc.RuntimeSettingsRefreshHealth().FailureTotal >= 3
})
if svc.IsMonitoringEnabled(context.Background()) {
t.Fatal("failed refresh overwrote last known monitoring state")
}
if !svc.OpsAdvancedSettingsSnapshot().IgnoreNoAvailableAccounts {
t.Fatal("failed refresh overwrote last known advanced settings")
}
}
func TestOpsRuntimeSettingsRefreshStopEndsLifecycle(t *testing.T) {
repo := &opsRuntimeRefreshRepo{values: map[string]string{}}
svc := &OpsService{settingRepo: repo}
svc.initRuntimeSettings(context.Background())
svc.startRuntimeSettingsRefresh(context.Background(), 5*time.Millisecond, 0, 50*time.Millisecond)
waitForOpsRefresh(t, time.Second, func() bool {
return svc.RuntimeSettingsRefreshHealth().SuccessTotal > 0
})
svc.StopRuntimeSettingsRefresh()
callsAfterStop := repo.calls.Load()
time.Sleep(20 * time.Millisecond)
if got := repo.calls.Load(); got != callsAfterStop {
t.Fatalf("refresh continued after Stop: before=%d after=%d", callsAfterStop, got)
}
if svc.RuntimeSettingsRefreshHealth().Running {
t.Fatal("refresh health still reports running after Stop")
}
// Idempotence is part of the cleanup contract.
svc.StopRuntimeSettingsRefresh()
}
func BenchmarkOpsRuntimeSettingsSnapshotRead(b *testing.B) {
svc := &OpsService{}
svc.initRuntimeSettings(context.Background())
ctx := context.Background()
b.ReportAllocs()
b.ResetTimer()
for range b.N {
benchmarkOpsMonitoringEnabled = svc.IsMonitoringEnabled(ctx)
benchmarkOpsAdvancedSettings = svc.OpsAdvancedSettingsSnapshot()
}
}
+291 -34
View File
@@ -6,7 +6,10 @@ import (
"encoding/json"
"errors"
"log"
"math/rand/v2"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -17,8 +20,27 @@ var ErrOpsDisabled = infraerrors.NotFound("OPS_DISABLED", "Ops monitoring is dis
const (
opsMaxStoredErrorBodyBytes = 20 * 1024
// OpsErrorLogQueueBodyMaxBytes bounds attacker-controlled response data while
// it waits in the asynchronous error-log queue.
OpsErrorLogQueueBodyMaxBytes = 8 * 1024
opsRuntimeSettingsRefreshInterval = 30 * time.Second
opsRuntimeSettingsRefreshJitter = 20
opsRuntimeSettingsRefreshTimeout = 3 * time.Second
opsRuntimeSettingsFailureLogEvery = time.Minute
)
type opsRuntimeSettingsSnapshot struct {
monitoringEnabled bool
advanced OpsAdvancedSettings
}
type OpsRuntimeSettingsRefreshHealth struct {
Running bool `json:"running"`
SuccessTotal uint64 `json:"success_total"`
FailureTotal uint64 `json:"failure_total"`
}
// OpsService provides ingestion and query APIs for the Ops monitoring module.
type OpsService struct {
opsRepo OpsRepository
@@ -31,12 +53,15 @@ type OpsService struct {
// getAccountAvailability is a unit-test hook for overriding account availability lookup.
getAccountAvailability func(ctx context.Context, platformFilter string, groupIDFilter *int64) (*OpsAccountAvailability, error)
concurrencyService *ConcurrencyService
gatewayService *GatewayService
openAIGatewayService *OpenAIGatewayService
geminiCompatService *GeminiMessagesCompatService
antigravityGatewayService *AntigravityGatewayService
systemLogSink *OpsSystemLogSink
concurrencyService *ConcurrencyService
gatewayService *GatewayService
openAIGatewayService *OpenAIGatewayService
geminiCompatService *GeminiMessagesCompatService
antigravityGatewayService *AntigravityGatewayService
systemLogSink *OpsSystemLogSink
ingressRejectAggregator *OpsIngressRejectAggregator
authCacheInvalidationWorker *AuthCacheInvalidationWorker
apiKeyService *APIKeyService
// cleanupReloader 由 wire 在 OpsCleanupService 构造完成后通过 SetCleanupReloader 注入。
// 解耦避免 OpsService -> OpsCleanupService 的硬依赖(cleanup 也读 settings,会循环)。
@@ -46,6 +71,19 @@ type OpsService struct {
// UpdateOpsAdvancedSettings 写入新配置后调用,把最新的 quota auto-pause 全局默认阈值
// 立即同步到调度热路径读取的内存缓存,避免下次请求才能感知新值。
quotaAutoPauseSink func(OpsOpenAIAccountQuotaAutoPauseSettings)
// Published snapshots are immutable. Gateway reads are lock-free; the mutex
// only serializes startup and administrative updates.
runtimeSettings atomic.Pointer[opsRuntimeSettingsSnapshot]
runtimeSettingsMu sync.Mutex
runtimeRefreshMu sync.Mutex
runtimeRefreshCancel context.CancelFunc
runtimeRefreshDone chan struct{}
runtimeRefreshRunning atomic.Bool
runtimeRefreshSuccess atomic.Uint64
runtimeRefreshFailure atomic.Uint64
runtimeRefreshLastFailureLog atomic.Int64
}
// CleanupReloader 由 OpsCleanupService 实现。
@@ -100,6 +138,7 @@ func NewOpsService(
antigravityGatewayService: antigravityGatewayService,
systemLogSink: systemLogSink,
}
svc.initRuntimeSettings(context.Background())
svc.applyRuntimeLogConfigOnStartup(context.Background())
return svc
}
@@ -112,22 +151,20 @@ func (s *OpsService) RequireMonitoringEnabled(ctx context.Context) error {
}
func (s *OpsService) IsMonitoringEnabled(ctx context.Context) bool {
_ = ctx
// Hard switch: disable ops entirely.
if s.cfg != nil && !s.cfg.Ops.Enabled {
return false
}
if s.settingRepo == nil {
return true
}
value, err := s.settingRepo.GetValue(ctx, SettingKeyOpsMonitoringEnabled)
if err != nil {
// Default enabled when key is missing, and fail-open on transient errors
// (ops should never block gateway traffic).
if errors.Is(err, ErrSettingNotFound) {
return true
}
return true
if snapshot := s.runtimeSettings.Load(); snapshot != nil {
return snapshot.monitoringEnabled
}
// Directly assembled test services and failed cold loads remain fail-open,
// without turning a request into a settings-table lookup.
return true
}
func parseOpsMonitoringEnabled(value string) bool {
switch strings.ToLower(strings.TrimSpace(value)) {
case "false", "0", "off", "disabled":
return false
@@ -136,6 +173,227 @@ func (s *OpsService) IsMonitoringEnabled(ctx context.Context) bool {
}
}
func (s *OpsService) initRuntimeSettings(ctx context.Context) {
if s == nil {
return
}
defaults := defaultOpsAdvancedSettings()
s.runtimeSettings.Store(&opsRuntimeSettingsSnapshot{monitoringEnabled: true, advanced: *defaults})
_ = s.RefreshRuntimeSettings(ctx)
}
// RefreshRuntimeSettings is the cold-path database load used at startup and by
// explicit administrative refreshes. Request processing only reads the atomic
// snapshot.
func (s *OpsService) RefreshRuntimeSettings(ctx context.Context) error {
if s == nil || s.settingRepo == nil {
return nil
}
if ctx == nil {
ctx = context.Background()
}
s.runtimeSettingsMu.Lock()
defer s.runtimeSettingsMu.Unlock()
values, err := s.settingRepo.GetMultiple(ctx, []string{
SettingKeyOpsMonitoringEnabled,
SettingKeyOpsAdvancedSettings,
})
if err != nil {
return err
}
monitoringEnabled := true
if raw, ok := values[SettingKeyOpsMonitoringEnabled]; ok {
monitoringEnabled = parseOpsMonitoringEnabled(raw)
}
advanced := defaultOpsAdvancedSettings()
if raw, ok := values[SettingKeyOpsAdvancedSettings]; ok {
if err := json.Unmarshal([]byte(raw), advanced); err != nil {
advanced = defaultOpsAdvancedSettings()
}
}
normalizeOpsAdvancedSettings(advanced)
s.runtimeSettings.Store(&opsRuntimeSettingsSnapshot{monitoringEnabled: monitoringEnabled, advanced: *advanced})
return nil
}
// StartRuntimeSettingsRefresh keeps DB-backed Ops settings converged across
// application instances without putting database I/O on request paths.
func (s *OpsService) StartRuntimeSettingsRefresh(ctx context.Context) {
s.startRuntimeSettingsRefresh(ctx, opsRuntimeSettingsRefreshInterval, opsRuntimeSettingsRefreshJitter, opsRuntimeSettingsRefreshTimeout)
}
func (s *OpsService) startRuntimeSettingsRefresh(ctx context.Context, interval time.Duration, jitterPercent int, timeout time.Duration) {
if s == nil || s.settingRepo == nil {
return
}
if ctx == nil {
ctx = context.Background()
}
if interval <= 0 {
interval = opsRuntimeSettingsRefreshInterval
}
if timeout <= 0 {
timeout = opsRuntimeSettingsRefreshTimeout
}
if jitterPercent < 0 {
jitterPercent = 0
}
if jitterPercent > 100 {
jitterPercent = 100
}
s.runtimeRefreshMu.Lock()
if s.runtimeRefreshCancel != nil {
s.runtimeRefreshMu.Unlock()
return
}
refreshCtx, cancel := context.WithCancel(ctx)
done := make(chan struct{})
s.runtimeRefreshCancel = cancel
s.runtimeRefreshDone = done
s.runtimeRefreshRunning.Store(true)
s.runtimeRefreshMu.Unlock()
go func() {
defer close(done)
defer s.runtimeRefreshRunning.Store(false)
for {
delay := jitterDuration(interval, jitterPercent)
timer := time.NewTimer(delay)
select {
case <-refreshCtx.Done():
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
return
case <-timer.C:
}
attemptCtx, attemptCancel := context.WithTimeout(refreshCtx, timeout)
err := s.RefreshRuntimeSettings(attemptCtx)
attemptCancel()
if err != nil {
s.runtimeRefreshFailure.Add(1)
s.logRuntimeSettingsRefreshFailure(err)
continue
}
s.runtimeRefreshSuccess.Add(1)
}
}()
}
func jitterDuration(base time.Duration, percent int) time.Duration {
if base <= 0 || percent <= 0 {
return base
}
delta := float64(percent) / 100
factor := 1 - delta + rand.Float64()*(2*delta)
if factor <= 0 {
return base
}
return time.Duration(float64(base) * factor)
}
func (s *OpsService) logRuntimeSettingsRefreshFailure(err error) {
if s == nil || err == nil {
return
}
now := time.Now().Unix()
for {
last := s.runtimeRefreshLastFailureLog.Load()
if last != 0 && now-last < int64(opsRuntimeSettingsFailureLogEvery/time.Second) {
return
}
if s.runtimeRefreshLastFailureLog.CompareAndSwap(last, now) {
log.Printf("[Ops] runtime settings refresh failed: %v", err)
return
}
}
}
// StopRuntimeSettingsRefresh is idempotent and waits for an in-flight refresh
// to observe cancellation before returning.
func (s *OpsService) StopRuntimeSettingsRefresh() {
if s == nil {
return
}
s.runtimeRefreshMu.Lock()
cancel := s.runtimeRefreshCancel
done := s.runtimeRefreshDone
s.runtimeRefreshMu.Unlock()
if cancel != nil {
cancel()
}
if done != nil {
<-done
}
s.runtimeRefreshMu.Lock()
if s.runtimeRefreshDone == done {
s.runtimeRefreshCancel = nil
s.runtimeRefreshDone = nil
}
s.runtimeRefreshMu.Unlock()
}
func (s *OpsService) RuntimeSettingsRefreshHealth() OpsRuntimeSettingsRefreshHealth {
if s == nil {
return OpsRuntimeSettingsRefreshHealth{}
}
return OpsRuntimeSettingsRefreshHealth{
Running: s.runtimeRefreshRunning.Load(),
SuccessTotal: s.runtimeRefreshSuccess.Load(),
FailureTotal: s.runtimeRefreshFailure.Load(),
}
}
// SetMonitoringEnabled publishes an already-persisted admin setting without a
// database round trip.
func (s *OpsService) SetMonitoringEnabled(enabled bool) {
if s == nil {
return
}
s.runtimeSettingsMu.Lock()
current := s.runtimeSettings.Load()
next := &opsRuntimeSettingsSnapshot{monitoringEnabled: enabled, advanced: *defaultOpsAdvancedSettings()}
if current != nil {
next.advanced = current.advanced
}
s.runtimeSettings.Store(next)
s.runtimeSettingsMu.Unlock()
}
func (s *OpsService) storeAdvancedSettingsSnapshot(cfg *OpsAdvancedSettings) {
if s == nil || cfg == nil {
return
}
s.runtimeSettingsMu.Lock()
current := s.runtimeSettings.Load()
next := &opsRuntimeSettingsSnapshot{monitoringEnabled: true, advanced: *cfg}
if current != nil {
next.monitoringEnabled = current.monitoringEnabled
}
s.runtimeSettings.Store(next)
s.runtimeSettingsMu.Unlock()
}
// SanitizeOpsErrorBodyForQueue removes credentials and truncates the body
// before it can consume capacity in the asynchronous queue.
func SanitizeOpsErrorBodyForQueue(raw string) (string, bool) {
return sanitizeErrorBodyForStorage(raw, OpsErrorLogQueueBodyMaxBytes)
}
// SanitizeOpsUpstreamErrorsForQueue bounds and serializes attempt-level data
// before the entry can consume asynchronous queue capacity.
func SanitizeOpsUpstreamErrorsForQueue(entry *OpsInsertErrorLogInput) error {
return sanitizeOpsUpstreamErrors(entry)
}
func (s *OpsService) RecordError(ctx context.Context, entry *OpsInsertErrorLogInput) error {
prepared, ok, err := s.prepareErrorLogInput(ctx, entry)
if err != nil {
@@ -295,7 +553,7 @@ func sanitizeOpsUpstreamErrors(entry *OpsInsertErrorLogInput) error {
return nil
}
const maxEvents = 32
const maxEvents = 16
events := entry.UpstreamErrors
if len(events) > maxEvents {
events = events[len(events)-maxEvents:]
@@ -308,9 +566,19 @@ func sanitizeOpsUpstreamErrors(entry *OpsInsertErrorLogInput) error {
}
out := *ev
out.Platform = strings.TrimSpace(out.Platform)
out.Platform = truncateString(strings.TrimSpace(out.Platform), 32)
out.AccountName = truncateString(strings.TrimSpace(out.AccountName), 128)
out.UpstreamRequestID = truncateString(strings.TrimSpace(out.UpstreamRequestID), 128)
out.UpstreamURL = truncateString(strings.TrimSpace(out.UpstreamURL), 2048)
if body := strings.TrimSpace(out.UpstreamResponseBody); body != "" {
out.UpstreamResponseBody, _ = sanitizeErrorBodyForStorage(body, OpsErrorLogQueueBodyMaxBytes)
} else {
out.UpstreamResponseBody = ""
}
out.Kind = truncateString(strings.TrimSpace(out.Kind), 64)
out.Stage = truncateString(strings.TrimSpace(out.Stage), 64)
out.Scope = truncateString(strings.TrimSpace(out.Scope), 64)
out.Reason = truncateString(strings.TrimSpace(out.Reason), 128)
if out.AccountID < 0 {
out.AccountID = 0
@@ -328,8 +596,8 @@ func sanitizeOpsUpstreamErrors(entry *OpsInsertErrorLogInput) error {
detail := strings.TrimSpace(out.Detail)
if detail != "" {
// Keep upstream detail small; request bodies are not stored here, only upstream error payloads.
sanitizedDetail, _ := sanitizeErrorBodyForStorage(detail, opsMaxStoredErrorBodyBytes)
// Keep upstream detail small while the event waits in the queue.
sanitizedDetail, _ := sanitizeErrorBodyForStorage(detail, OpsErrorLogQueueBodyMaxBytes)
out.Detail = sanitizedDetail
} else {
out.Detail = ""
@@ -375,8 +643,6 @@ func (s *OpsService) ListUserErrorRequests(ctx context.Context, userID int64, fi
filter = &f
uid := userID
filter.UserID = &uid
// 用户侧放宽归属:纳入「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录。
filter.MatchDeletedKeyOwner = true
// APIKeyID 透传:保留 handler 传入的值。安全由 buildOpsErrorLogsWhere 的
// "user_id = 自己 AND api_key_id = X" 双重约束保证——传入他人 key 只会得到空集,无泄露。
filter.View = "all"
@@ -444,23 +710,14 @@ func (s *OpsService) GetUserErrorRequestDetail(ctx context.Context, userID, id i
}
return nil, infraerrors.InternalServer("OPS_ERROR_LOAD_FAILED", "Failed to load ops error log").WithCause(err)
}
// 归属:直接归属(user_id)或经「已删除 key 归因」(deleted_key_owner_user_id)二者之一即可。
// 归属只能由通过鉴权时写入的 user_id 确定。
ownedDirectly := detail.UserID != nil && *detail.UserID == userID
ownedViaDeletedKey := detail.DeletedKeyOwnerUserID != nil && *detail.DeletedKeyOwnerUserID == userID
if !ownedDirectly && !ownedViaDeletedKey {
if !ownedDirectly {
return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
}
return ToUserErrorRequestDetail(detail), nil
}
// LookupDeletedKeyAudit 按明文 key 反查已删除 key 的原所有者;未命中或未启用返回 (nil, nil)。
func (s *OpsService) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
if s.opsRepo == nil {
return nil, nil
}
return s.opsRepo.LookupDeletedKeyAudit(ctx, key)
}
func (s *OpsService) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64) error {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return err
@@ -162,61 +162,3 @@ func TestGetUserErrorRequestDetail_InvalidID(t *testing.T) {
t.Fatal("expected error for id=-5")
}
}
func TestListUserErrorRequests_EnablesMatchDeletedKeyOwner(t *testing.T) {
stub := &stubOpsRepoForUserErr{}
svc := &OpsService{opsRepo: stub}
uid := int64(42)
if _, err := svc.ListUserErrorRequests(context.Background(), uid, &OpsErrorLogFilter{}); err != nil {
t.Fatal(err)
}
if stub.gotFilter == nil || !stub.gotFilter.MatchDeletedKeyOwner {
t.Fatal("ListUserErrorRequests should enable MatchDeletedKeyOwner for the user scope")
}
}
func TestGetUserErrorRequestDetail_DeletedKeyOwnerAccess(t *testing.T) {
ownerUID := int64(777)
otherUID := int64(2)
// 情况2:user_id=NULL,靠 deleted_key_owner_user_id 归因到 ownerUID
mk := func() *OpsErrorLogDetail {
return &OpsErrorLogDetail{
OpsErrorLog: OpsErrorLog{
ID: 55,
Phase: "auth",
Type: "api_error",
StatusCode: 401,
Message: "Invalid API key",
UserID: nil,
APIKeyName: "my-old-key",
APIKeyDeleted: true,
DeletedKeyOwnerUserID: &ownerUID,
},
}
}
// 原所有者(经 deleted_key 归因)→ 放行
svcOwner := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
got, err := svcOwner.GetUserErrorRequestDetail(context.Background(), ownerUID, 55)
if err != nil {
t.Fatalf("owner via deleted_key should be allowed, got err: %v", err)
}
if got == nil || got.ID != 55 {
t.Fatalf("expected detail ID=55, got %+v", got)
}
if !got.KeyDeleted || got.KeyName != "my-old-key" {
t.Fatalf("expected KeyDeleted=true KeyName=my-old-key, got %+v", got)
}
// 他人 → NotFound,不泄露存在性
svcOther := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
got2, err2 := svcOther.GetUserErrorRequestDetail(context.Background(), otherUID, 55)
if err2 == nil || got2 != nil {
t.Fatalf("non-owner should get (nil, NotFound), got detail=%+v err=%v", got2, err2)
}
if !infraerrors.IsNotFound(err2) {
t.Fatalf("expected NotFound, got %v", err2)
}
}
+17 -23
View File
@@ -373,6 +373,7 @@ func defaultOpsAdvancedSettings() *OpsAdvancedSettings {
IgnoreCountTokensErrors: true, // count_tokens 404 是预期行为,默认忽略
IgnoreContextCanceled: true, // Default to true - client disconnects are not errors
IgnoreNoAvailableAccounts: false, // Default to false - this is a real routing issue
IgnoreInvalidApiKeyErrors: true, // Legacy compatibility field; admission rejects are always excluded.
IgnoreInsufficientBalanceErrors: false, // 默认不忽略,余额不足可能需要关注
DisplayOpenAITokenStats: false,
DisplayAlertEvents: true,
@@ -385,6 +386,10 @@ func normalizeOpsAdvancedSettings(cfg *OpsAdvancedSettings) {
if cfg == nil {
return
}
// Admission rejects are a security/traffic concern, not an operational
// request-error category. Keep the legacy field true for old clients but do
// not allow it to re-enable those rows.
cfg.IgnoreInvalidApiKeyErrors = true
cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold5h = clampOpsQuotaAutoPauseThreshold(cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold5h)
cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold7d = clampOpsQuotaAutoPauseThreshold(cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold7d)
cfg.DataRetention.CleanupSchedule = strings.TrimSpace(cfg.DataRetention.CleanupSchedule)
@@ -439,32 +444,20 @@ func validateOpsAdvancedSettings(cfg *OpsAdvancedSettings) error {
}
func (s *OpsService) GetOpsAdvancedSettings(ctx context.Context) (*OpsAdvancedSettings, error) {
defaultCfg := defaultOpsAdvancedSettings()
if s == nil || s.settingRepo == nil {
return defaultCfg, nil
}
if ctx == nil {
ctx = context.Background()
}
_ = ctx
cfg := s.OpsAdvancedSettingsSnapshot()
return &cfg, nil
}
raw, err := s.settingRepo.GetValue(ctx, SettingKeyOpsAdvancedSettings)
if err != nil {
if errors.Is(err, ErrSettingNotFound) {
if b, mErr := json.Marshal(defaultCfg); mErr == nil {
_ = s.settingRepo.Set(ctx, SettingKeyOpsAdvancedSettings, string(b))
}
return defaultCfg, nil
// OpsAdvancedSettingsSnapshot returns a value copy for request hot paths. It
// avoids both repository I/O and pointer escape/allocation.
func (s *OpsService) OpsAdvancedSettingsSnapshot() OpsAdvancedSettings {
if s != nil {
if snapshot := s.runtimeSettings.Load(); snapshot != nil {
return snapshot.advanced
}
return nil, err
}
cfg := defaultOpsAdvancedSettings()
if err := json.Unmarshal([]byte(raw), cfg); err != nil {
return defaultCfg, nil
}
normalizeOpsAdvancedSettings(cfg)
return cfg, nil
return *defaultOpsAdvancedSettings()
}
func (s *OpsService) UpdateOpsAdvancedSettings(ctx context.Context, cfg *OpsAdvancedSettings) (*OpsAdvancedSettings, error) {
@@ -490,6 +483,7 @@ func (s *OpsService) UpdateOpsAdvancedSettings(ctx context.Context, cfg *OpsAdva
if err := s.settingRepo.Set(ctx, SettingKeyOpsAdvancedSettings, string(raw)); err != nil {
return nil, err
}
s.storeAdvancedSettingsSnapshot(cfg)
// Push the new quota auto-pause settings straight into the in-memory cache that
// the OpenAI scheduling hot path reads, so the next request observes the new value
// without waiting for the background refresher's TTL.
@@ -9,7 +9,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
)
func TestGetOpsAdvancedSettings_DefaultHidesOpenAITokenStats(t *testing.T) {
func TestGetOpsAdvancedSettings_DefaultSnapshotHidesOpenAITokenStats(t *testing.T) {
repo := newRuntimeSettingRepoStub()
svc := &OpsService{settingRepo: repo}
@@ -23,8 +23,8 @@ func TestGetOpsAdvancedSettings_DefaultHidesOpenAITokenStats(t *testing.T) {
if !cfg.DisplayAlertEvents {
t.Fatalf("DisplayAlertEvents = false, want true by default")
}
if repo.setCalls != 1 {
t.Fatalf("expected defaults to be persisted once, got %d", repo.setCalls)
if repo.getValueCalls != 0 || repo.getMultipleCalls != 0 {
t.Fatalf("hot-path snapshot read touched repository: get=%d get_multiple=%d", repo.getValueCalls, repo.getMultipleCalls)
}
}
@@ -46,6 +46,7 @@ func TestUpdateOpsAdvancedSettings_PersistsOpenAITokenStatsVisibility(t *testing
if updated.DisplayAlertEvents {
t.Fatalf("DisplayAlertEvents = true, want false")
}
readsAfterUpdate := repo.getValueCalls + repo.getMultipleCalls
reloaded, err := svc.GetOpsAdvancedSettings(context.Background())
if err != nil {
@@ -57,6 +58,9 @@ func TestUpdateOpsAdvancedSettings_PersistsOpenAITokenStatsVisibility(t *testing
if reloaded.DisplayAlertEvents {
t.Fatalf("reloaded DisplayAlertEvents = true, want false")
}
if got := repo.getValueCalls + repo.getMultipleCalls; got != readsAfterUpdate {
t.Fatalf("snapshot reload performed repository read: before=%d after=%d", readsAfterUpdate, got)
}
}
func TestGetOpsAdvancedSettings_BackfillsNewDisplayFlagsFromDefaults(t *testing.T) {
@@ -77,7 +81,7 @@ func TestGetOpsAdvancedSettings_BackfillsNewDisplayFlagsFromDefaults(t *testing.
"ignore_count_tokens_errors": true,
"ignore_context_canceled": true,
"ignore_no_available_accounts": false,
"ignore_invalid_api_key_errors": false,
"ignore_invalid_api_key_errors": true,
"auto_refresh_enabled": false,
"auto_refresh_interval_seconds": 30,
}
+13 -12
View File
@@ -92,18 +92,19 @@ type OpsAlertRuntimeSettings struct {
// OpsAdvancedSettings stores advanced ops configuration (data retention, aggregation).
type OpsAdvancedSettings struct {
DataRetention OpsDataRetentionSettings `json:"data_retention"`
Aggregation OpsAggregationSettings `json:"aggregation"`
OpenAIAccountQuotaAutoPause OpsOpenAIAccountQuotaAutoPauseSettings `json:"openai_account_quota_auto_pause"`
IgnoreCountTokensErrors bool `json:"ignore_count_tokens_errors"`
IgnoreContextCanceled bool `json:"ignore_context_canceled"`
IgnoreNoAvailableAccounts bool `json:"ignore_no_available_accounts"`
IgnoreInvalidApiKeyErrors bool `json:"ignore_invalid_api_key_errors"`
IgnoreInsufficientBalanceErrors bool `json:"ignore_insufficient_balance_errors"`
DisplayOpenAITokenStats bool `json:"display_openai_token_stats"`
DisplayAlertEvents bool `json:"display_alert_events"`
AutoRefreshEnabled bool `json:"auto_refresh_enabled"`
AutoRefreshIntervalSec int `json:"auto_refresh_interval_seconds"`
DataRetention OpsDataRetentionSettings `json:"data_retention"`
Aggregation OpsAggregationSettings `json:"aggregation"`
OpenAIAccountQuotaAutoPause OpsOpenAIAccountQuotaAutoPauseSettings `json:"openai_account_quota_auto_pause"`
IgnoreCountTokensErrors bool `json:"ignore_count_tokens_errors"`
IgnoreContextCanceled bool `json:"ignore_context_canceled"`
IgnoreNoAvailableAccounts bool `json:"ignore_no_available_accounts"`
// Deprecated compatibility field. It is always normalized to true.
IgnoreInvalidApiKeyErrors bool `json:"ignore_invalid_api_key_errors"`
IgnoreInsufficientBalanceErrors bool `json:"ignore_insufficient_balance_errors"`
DisplayOpenAITokenStats bool `json:"display_openai_token_stats"`
DisplayAlertEvents bool `json:"display_alert_events"`
AutoRefreshEnabled bool `json:"auto_refresh_enabled"`
AutoRefreshIntervalSec int `json:"auto_refresh_interval_seconds"`
}
type OpsOpenAIAccountQuotaAutoPauseSettings struct {
@@ -112,6 +112,11 @@ func (s *OpsSystemLogSink) WriteLogEvent(event *logger.LogEvent) {
}
func (s *OpsSystemLogSink) shouldIndex(event *logger.LogEvent) bool {
if event != nil && event.Fields != nil {
if skip, _ := event.Fields[logger.OpsSystemLogSkipField].(bool); skip {
return false
}
}
level := strings.ToLower(strings.TrimSpace(event.Level))
switch level {
case "warn", "warning", "error", "fatal", "panic", "dpanic":
@@ -36,6 +36,15 @@ func TestOpsSystemLogSink_ShouldIndex(t *testing.T) {
event: &logger.LogEvent{Level: "info", Component: "http.access"},
want: true,
},
{
name: "rejected access excluded from database sink",
event: &logger.LogEvent{
Level: "info",
Component: "http.access",
Fields: map[string]any{logger.OpsSystemLogSkipField: true},
},
want: false,
},
{
name: "access component from fields (real zap path)",
event: &logger.LogEvent{
+20
View File
@@ -567,6 +567,8 @@ func ProvideOpsService(
antigravityGatewayService *AntigravityGatewayService,
systemLogSink *OpsSystemLogSink,
settingService *SettingService,
authCacheInvalidationWorker *AuthCacheInvalidationWorker,
apiKeyService *APIKeyService,
) *OpsService {
svc := NewOpsService(
opsRepo,
@@ -587,9 +589,25 @@ func ProvideOpsService(
// a populated cache rather than zero defaults. Best-effort, sync-bounded.
settingService.WarmOpenAIQuotaAutoPauseSettings(context.Background())
}
svc.authCacheInvalidationWorker = authCacheInvalidationWorker
svc.apiKeyService = apiKeyService
svc.StartRuntimeSettingsRefresh(context.Background())
return svc
}
// ProvideOpsIngressRejectAggregator starts the bounded security aggregation
// runtime and attaches it to OpsService, which is the middleware recorder.
func ProvideOpsIngressRejectAggregator(opsRepo OpsRepository, opsService *OpsService) *OpsIngressRejectAggregator {
repo, ok := opsRepo.(OpsIngressRejectRepository)
if !ok {
return nil
}
aggregator := NewOpsIngressRejectAggregator(repo)
aggregator.Start()
opsService.SetIngressRejectAggregator(aggregator)
return aggregator
}
// ProvideSettingService wires SettingService with group reader and proxy repo.
func ProvideSettingService(settingRepo SettingRepository, groupRepo GroupRepository, proxyRepo ProxyRepository, cfg *config.Config) *SettingService {
svc := NewSettingService(settingRepo, cfg)
@@ -647,6 +665,7 @@ var ProviderSet = wire.NewSet(
NewUserService,
ProvideAPIKeyService,
ProvideAPIKeyAuthCacheInvalidator,
ProvideAuthCacheInvalidationWorker,
NewGroupService,
NewAccountService,
NewProxyService,
@@ -696,6 +715,7 @@ var ProviderSet = wire.NewSet(
ProvideBackupService,
ProvideOpsSystemLogSink,
ProvideOpsService,
ProvideOpsIngressRejectAggregator,
ProvideAuditLogService,
ProvideOpsMetricsCollector,
ProvideOpsAggregationService,
@@ -0,0 +1,27 @@
SET LOCAL lock_timeout = '5s';
SET LOCAL statement_timeout = '10min';
CREATE TABLE IF NOT EXISTS ops_ingress_reject_aggregates (
id BIGSERIAL PRIMARY KEY,
bucket_start TIMESTAMPTZ NOT NULL,
reject_reason VARCHAR(64) NOT NULL,
route_family VARCHAR(64) NOT NULL,
protocol VARCHAR(32) NOT NULL,
client_ip INET NOT NULL,
user_id BIGINT NOT NULL DEFAULT 0,
api_key_id BIGINT NOT NULL DEFAULT 0,
request_count BIGINT NOT NULL DEFAULT 0,
first_seen TIMESTAMPTZ NOT NULL,
last_seen TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
CONSTRAINT ops_ingress_reject_aggregates_dimensions_unique UNIQUE
(bucket_start, reject_reason, route_family, protocol, client_ip, user_id, api_key_id)
);
CREATE INDEX IF NOT EXISTS idx_ops_ingress_reject_aggregates_bucket
ON ops_ingress_reject_aggregates (bucket_start DESC);
CREATE INDEX IF NOT EXISTS idx_ops_ingress_reject_aggregates_reason_bucket
ON ops_ingress_reject_aggregates (reject_reason, bucket_start DESC);
CREATE INDEX IF NOT EXISTS idx_ops_ingress_reject_aggregates_ip_bucket
ON ops_ingress_reject_aggregates (client_ip, bucket_start DESC);
@@ -0,0 +1,198 @@
-- Durable, transactionally-enqueued API-key auth cache invalidation.
-- cache_key is always SHA-256 hex; plaintext credentials never leave api_keys.
CREATE TABLE IF NOT EXISTS auth_cache_invalidation_outbox (
id BIGSERIAL PRIMARY KEY,
cache_key CHAR(64) NOT NULL CHECK (cache_key ~ '^[0-9a-f]{64}$'),
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
available_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
delivery_stage SMALLINT NOT NULL DEFAULT 0 CHECK (delivery_stage IN (0, 1)),
attempts INTEGER NOT NULL DEFAULT 0 CHECK (attempts >= 0),
last_error TEXT,
claimed_at TIMESTAMPTZ,
claimed_by TEXT
);
CREATE INDEX IF NOT EXISTS idx_auth_cache_invalidation_outbox_available
ON auth_cache_invalidation_outbox (available_at, id)
WHERE claimed_at IS NULL;
CREATE INDEX IF NOT EXISTS idx_auth_cache_invalidation_outbox_lease
ON auth_cache_invalidation_outbox (claimed_at)
WHERE claimed_at IS NOT NULL;
CREATE INDEX IF NOT EXISTS idx_auth_cache_invalidation_outbox_cache_key
ON auth_cache_invalidation_outbox (cache_key);
CREATE INDEX IF NOT EXISTS idx_auth_cache_invalidation_outbox_created_at
ON auth_cache_invalidation_outbox (created_at);
CREATE OR REPLACE FUNCTION enqueue_auth_cache_invalidation(raw_key TEXT)
RETURNS VOID
LANGUAGE plpgsql
AS $$
BEGIN
IF raw_key IS NULL OR raw_key = '' THEN
RETURN;
END IF;
INSERT INTO auth_cache_invalidation_outbox (cache_key)
VALUES (encode(sha256(convert_to(raw_key, 'UTF8')), 'hex'));
END;
$$;
CREATE OR REPLACE FUNCTION enqueue_api_key_auth_cache_invalidation()
RETURNS TRIGGER
LANGUAGE plpgsql
AS $$
BEGIN
IF TG_OP = 'DELETE' THEN
PERFORM enqueue_auth_cache_invalidation(OLD.key);
RETURN OLD;
END IF;
IF OLD.key IS DISTINCT FROM NEW.key
OR OLD.status IS DISTINCT FROM NEW.status
OR OLD.deleted_at IS DISTINCT FROM NEW.deleted_at
OR OLD.user_id IS DISTINCT FROM NEW.user_id
OR OLD.group_id IS DISTINCT FROM NEW.group_id
OR OLD.ip_whitelist IS DISTINCT FROM NEW.ip_whitelist
OR OLD.ip_blacklist IS DISTINCT FROM NEW.ip_blacklist
OR OLD.expires_at IS DISTINCT FROM NEW.expires_at THEN
PERFORM enqueue_auth_cache_invalidation(OLD.key);
IF NEW.deleted_at IS NULL AND NEW.key IS DISTINCT FROM OLD.key THEN
PERFORM enqueue_auth_cache_invalidation(NEW.key);
END IF;
END IF;
RETURN NEW;
END;
$$;
DROP TRIGGER IF EXISTS trg_api_keys_auth_cache_invalidation ON api_keys;
CREATE TRIGGER trg_api_keys_auth_cache_invalidation
AFTER UPDATE OR DELETE ON api_keys
FOR EACH ROW EXECUTE FUNCTION enqueue_api_key_auth_cache_invalidation();
CREATE OR REPLACE FUNCTION enqueue_user_auth_cache_invalidation()
RETURNS TRIGGER
LANGUAGE plpgsql
AS $$
DECLARE
target_user_id BIGINT;
BEGIN
target_user_id := OLD.id;
IF TG_OP = 'UPDATE'
AND OLD.status IS NOT DISTINCT FROM NEW.status
AND OLD.role IS NOT DISTINCT FROM NEW.role
AND OLD.deleted_at IS NOT DISTINCT FROM NEW.deleted_at THEN
RETURN NEW;
END IF;
INSERT INTO auth_cache_invalidation_outbox (cache_key)
SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex')
FROM api_keys AS k
WHERE k.user_id = target_user_id
AND k.deleted_at IS NULL
AND k.key <> '';
IF TG_OP = 'DELETE' THEN
RETURN OLD;
END IF;
RETURN NEW;
END;
$$;
DROP TRIGGER IF EXISTS trg_users_auth_cache_invalidation ON users;
CREATE TRIGGER trg_users_auth_cache_invalidation
AFTER UPDATE OR DELETE ON users
FOR EACH ROW EXECUTE FUNCTION enqueue_user_auth_cache_invalidation();
CREATE OR REPLACE FUNCTION enqueue_group_auth_cache_invalidation()
RETURNS TRIGGER
LANGUAGE plpgsql
AS $$
DECLARE
target_group_id BIGINT;
BEGIN
target_group_id := OLD.id;
IF TG_OP = 'UPDATE'
AND OLD.status IS NOT DISTINCT FROM NEW.status
AND OLD.is_exclusive IS NOT DISTINCT FROM NEW.is_exclusive
AND OLD.deleted_at IS NOT DISTINCT FROM NEW.deleted_at THEN
RETURN NEW;
END IF;
INSERT INTO auth_cache_invalidation_outbox (cache_key)
SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex')
FROM api_keys AS k
WHERE k.group_id = target_group_id
AND k.deleted_at IS NULL
AND k.key <> '';
IF TG_OP = 'DELETE' THEN
RETURN OLD;
END IF;
RETURN NEW;
END;
$$;
DROP TRIGGER IF EXISTS trg_groups_auth_cache_invalidation ON groups;
CREATE TRIGGER trg_groups_auth_cache_invalidation
AFTER UPDATE OR DELETE ON groups
FOR EACH ROW EXECUTE FUNCTION enqueue_group_auth_cache_invalidation();
CREATE OR REPLACE FUNCTION enqueue_allowed_group_auth_cache_invalidation()
RETURNS TRIGGER
LANGUAGE plpgsql
AS $$
DECLARE
target_user_id BIGINT;
target_group_id BIGINT;
BEGIN
IF TG_OP = 'UPDATE'
AND (OLD.user_id IS DISTINCT FROM NEW.user_id
OR OLD.group_id IS DISTINCT FROM NEW.group_id) THEN
IF EXISTS (
SELECT 1 FROM groups g
WHERE g.id = OLD.group_id AND g.is_exclusive = TRUE
) THEN
INSERT INTO auth_cache_invalidation_outbox (cache_key)
SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex')
FROM api_keys AS k
WHERE k.user_id = OLD.user_id
AND k.group_id = OLD.group_id
AND k.deleted_at IS NULL
AND k.key <> '';
END IF;
target_user_id := NEW.user_id;
target_group_id := NEW.group_id;
ELSIF TG_OP = 'UPDATE' THEN
RETURN NEW;
ELSIF TG_OP = 'INSERT' THEN
target_user_id := NEW.user_id;
target_group_id := NEW.group_id;
ELSE
target_user_id := OLD.user_id;
target_group_id := OLD.group_id;
END IF;
IF EXISTS (
SELECT 1 FROM groups g
WHERE g.id = target_group_id AND g.is_exclusive = TRUE
) THEN
INSERT INTO auth_cache_invalidation_outbox (cache_key)
SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex')
FROM api_keys AS k
WHERE k.user_id = target_user_id
AND k.group_id = target_group_id
AND k.deleted_at IS NULL
AND k.key <> '';
END IF;
IF TG_OP = 'DELETE' THEN
RETURN OLD;
END IF;
RETURN NEW;
END;
$$;
DROP TRIGGER IF EXISTS trg_user_allowed_groups_auth_cache_invalidation ON user_allowed_groups;
CREATE TRIGGER trg_user_allowed_groups_auth_cache_invalidation
AFTER INSERT OR UPDATE OR DELETE ON user_allowed_groups
FOR EACH ROW EXECUTE FUNCTION enqueue_allowed_group_auth_cache_invalidation();
COMMENT ON TABLE auth_cache_invalidation_outbox IS
'Durable cross-instance auth cache invalidations; cache_key is SHA-256 hex, never plaintext API key';
@@ -0,0 +1,23 @@
-- Post-rollout finalizer for ingress-rejection log cleanup.
--
-- DO NOT run this during a rolling deployment. Run it only after:
-- 1. every application instance is on the release that no longer reads or
-- writes deleted_api_key_audits and the deprecated ops_error_logs columns;
-- 2. cleanup-ingress-reject-logs has been dry-run and, if desired, executed;
-- 3. a database backup or recovery point has been verified.
BEGIN;
SET LOCAL lock_timeout = '5s';
SET LOCAL statement_timeout = '10min';
DROP TABLE IF EXISTS deleted_api_key_audits;
ALTER TABLE IF EXISTS ops_error_logs
DROP COLUMN IF EXISTS attempted_key_prefix,
DROP COLUMN IF EXISTS deleted_key_owner_user_id,
DROP COLUMN IF EXISTS deleted_key_name;
COMMIT;
-- Run this separately during a normal maintenance window:
-- VACUUM (ANALYZE) ops_error_logs;
+19 -11
View File
@@ -1,12 +1,25 @@
{
servers {
max_header_size 64KB
timeouts {
read_header 10s
idle 2m
}
}
}
# 修改为你的域名
api.sub2api.com {
# This baseline assumes clients connect directly to Caddy. When Caddy is
# behind a CDN, configure explicit trusted proxy CIDRs and {client_ip} as
# documented in EDGE_SECURITY.md; {remote_host} would otherwise be the CDN.
# =========================================================================
# TLS 安全配置
# =========================================================================
tls {
# 仅使用 TLS 1.2 和 1.3
protocols tls1.2 tls1.3
# 优先使用的加密套件
ciphers TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384 TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384 TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305_SHA256 TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305_SHA256 TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256 TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256
}
@@ -20,22 +33,17 @@ api.sub2api.com {
health_interval 30s
health_timeout 10s
health_status 200
# 负载均衡策略(单节点可忽略,多节点时有用)
lb_policy round_robin
lb_try_duration 5s
lb_try_interval 250ms
# 传递真实客户端信息
# 兼容 Cloudflare 和直连:后端应优先读取 CF-Connecting-IP,其次 X-Real-IP
# 仅从实际 TCP 对端生成转发头,避免透传客户端伪造值。
header_up X-Real-IP {remote_host}
header_up X-Forwarded-For {remote_host}
header_up X-Forwarded-Proto {scheme}
header_up X-Forwarded-Host {host}
# 保留 Cloudflare 原始头(如果存在)
# 后端获取 IP 的优先级建议: CF-Connecting-IP → X-Real-IP → X-Forwarded-For
header_up CF-Connecting-IP {http.request.header.CF-Connecting-IP}
# 连接池优化
transport http {
keepalive 120s
@@ -44,7 +52,7 @@ api.sub2api.com {
write_buffer 16KB
compression off
}
# 故障转移
fail_duration 30s
max_fails 3
@@ -72,7 +80,7 @@ api.sub2api.com {
# 请求大小限制 (防止大文件攻击)
# =========================================================================
request_body {
max_size 100MB
max_size 256MB
}
# =========================================================================
+155
View File
@@ -0,0 +1,155 @@
# Edge and HTTP Ingress Security
Sub2API supports long-lived SSE and WebSocket requests. Protect the request
ingress without imposing a response `WriteTimeout`: a write deadline would
terminate healthy long generations and streams.
## Application defaults
- `server.max_header_bytes: 65536` limits HTTP/1 request headers to 64 KiB;
Go maps it to the corresponding HTTP/2 header-list limit.
- `server.read_header_timeout: 10` bounds slow-header attacks. It does not
limit request processing or response streaming.
- `server.max_request_body_size: 268435456` is the absolute 256 MiB safety net.
- `gateway.max_body_size: 268435456` remains available to multimodal, Gemini,
image, video, and batch-image endpoints.
- `gateway.text_max_body_size: 33554432` limits the known pure-text
`/embeddings` and `/alpha/search` endpoints to 32 MiB.
- H2C defaults to 50 concurrent streams per connection, a 2 MiB connection
upload window, and a 512 KiB stream upload window.
- Invalid credential abuse is limited in process by trusted client IP (IPv6
`/64`): 120 failures per 60 seconds followed by a 60-second block. This is a
per-instance safety net; multi-instance enforcement still belongs at the
load balancer, CDN, or WAF.
Do not add a single application-wide request semaphore: an SSE request may
legitimately occupy it for many minutes. Apply connection and unauthenticated
request controls at the edge; authenticated user/API-key concurrency remains
the application's responsibility.
## Trusted client IPs
`server.trusted_proxies` must contain only the CIDR/IP addresses that connect
directly to Sub2API, normally the local Nginx/Caddy address or the private load
balancer subnet. An empty list disables forwarded-IP trust.
Never trust `CF-Connecting-IP`, `X-Real-IP`, or `X-Forwarded-For` merely because
the header exists. A CDN deployment must firewall the origin so only the CDN or
load balancer can reach it, and the proxy must overwrite forwarded headers.
Example for a proxy on the same host:
```yaml
server:
trusted_proxies:
- 127.0.0.1/32
- ::1/128
```
## Nginx baseline
Define shared zones in the `http` block. Tune rates to measured legitimate
traffic; the values below are conservative starting points, not universal
capacity targets.
```nginx
limit_conn_zone $binary_remote_addr zone=sub2api_conn:20m;
limit_req_zone $binary_remote_addr zone=sub2api_auth:20m rate=5r/s;
limit_req_zone $binary_remote_addr zone=sub2api_api:40m rate=30r/s;
map $http_upgrade $connection_upgrade {
default upgrade;
'' close;
}
server {
listen 443 ssl http2;
server_name api.example.com;
client_header_timeout 10s;
client_max_body_size 256m;
large_client_header_buffers 4 16k;
limit_conn sub2api_conn 40;
location ~ ^/(auth|api/auth)/ {
limit_req zone=sub2api_auth burst=10 nodelay;
proxy_pass http://127.0.0.1:8080;
}
location ~ ^/(v1/)?(embeddings|alpha/search)$ {
client_max_body_size 32m;
limit_req zone=sub2api_api burst=60 nodelay;
proxy_pass http://127.0.0.1:8080;
}
location / {
limit_req zone=sub2api_api burst=60 nodelay;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $remote_addr;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection $connection_upgrade;
proxy_buffering off;
proxy_request_buffering off;
proxy_read_timeout 1800s;
proxy_send_timeout 1800s;
proxy_pass http://127.0.0.1:8080;
}
}
```
Do not use an incoming `$http_x_forwarded_for` value unless Nginx real-IP
processing is restricted to explicit trusted proxy CIDRs.
## Caddy and CDN
The bundled `deploy/Caddyfile` sets a 64 KiB header limit, a 10-second header
timeout, a 256 MiB absolute body limit, and overwrites forwarded addresses from
the TCP peer. It is therefore a direct-to-Caddy baseline. Do not use its
`{remote_host}` forwarding lines unchanged behind a CDN: all clients would be
attributed to a CDN egress address, collapsing rejection aggregation and the
invalid-auth limiter onto unrelated users.
For a CDN deployment, first firewall the origin so only current CDN egress
CIDRs can connect. Then configure those exact ranges as Caddy trusted proxies
and derive upstream headers from Caddy's parsed `{client_ip}`. For example:
```caddyfile
{
servers {
trusted_proxies static 192.0.2.0/24 2001:db8:1234::/48
trusted_proxies_strict
client_ip_headers CF-Connecting-IP X-Forwarded-For
}
}
api.example.com {
reverse_proxy 127.0.0.1:8080 {
header_up X-Real-IP {client_ip}
header_up X-Forwarded-For {client_ip}
}
}
```
Replace the documentation ranges with the CDN's published, automatically
maintained egress ranges. `CF-Connecting-IP` is safe here only because direct
origin access is blocked and Caddy trusts only those TCP peers. Configure
Sub2API `server.trusted_proxies` with the Caddy address/private subnet so the
application accepts only Caddy's rewritten headers.
Caddy core does not provide a general request-rate limiter; use a trusted
CDN/WAF, a supported rate-limit module, or host firewall controls.
At a CDN/WAF, configure connection limits, header/body limits, bot challenges,
and per-IP/ASN rates before traffic reaches the origin. Allow origin ingress
only from CDN egress CIDRs or a private load balancer. Keep the application port
off the public Internet.
## DDoS boundary
Application checks reduce amplification after a connection reaches Go. They
cannot absorb volumetric attacks, TLS floods, bandwidth saturation, or a large
distributed source set. Those require upstream network capacity, CDN/WAF
filtering, provider firewall rules, and origin isolation. Avoid high-cardinality
metrics or per-request database security logs during rejection storms.
+1
View File
@@ -27,6 +27,7 @@ This directory contains files for deploying Sub2API on Linux servers and Apple-s
| `sub2api-datamanagementd.service` | datamanagementd systemd service unit file |
| `DATAMANAGEMENTD_CN.md` | datamanagementd 部署与联动说明(中文) |
| `config.example.yaml` | Example configuration file |
| `EDGE_SECURITY.md` | Reverse proxy, CDN/WAF, trusted proxy, and ingress hardening guide |
---
+29
View File
@@ -27,6 +27,15 @@ server:
# 用于生成邮件中的外部链接(例如:重置密码链接)的前端基础地址
# Example: "https://example.com"
frontend_url: ""
# Maximum time to receive complete request headers. Does not limit response streams.
# 完整读取请求头的最大时间;不限制响应流持续时间。
read_header_timeout: 10
# Request header limit in bytes (64 KiB); also bounds the HTTP/2 header list.
# 请求头上限(字节,默认 64 KiB);同时约束 HTTP/2 header list。
max_header_bytes: 65536
# Keep-alive idle timeout in seconds.
# Keep-Alive 空闲连接超时(秒)。
idle_timeout: 120
# Trusted proxies for X-Forwarded-For parsing (CIDR/IP). Empty disables trusted proxies.
# 信任的代理地址(CIDR/IP 格式),用于解析 X-Forwarded-For 头。留空则禁用代理信任。
trusted_proxies: []
@@ -167,6 +176,9 @@ gateway:
# Max request body size in bytes (default: 256MB)
# 请求体最大字节数(默认 256MB)
max_body_size: 268435456
# Pure-text endpoint body limit (embeddings and alpha/search), default 32 MiB.
# 纯文本端点请求体上限(embeddings、alpha/search),默认 32 MiB。
text_max_body_size: 33554432
# Max bytes to read for non-stream upstream responses (default: 8MB)
# 非流式上游响应体读取上限(默认 8MB)
upstream_response_read_max_bytes: 8388608
@@ -708,6 +720,23 @@ api_key_auth_cache:
# Enable singleflight for cache misses
# 缓存未命中时启用 singleflight 合并回源
singleflight: true
# Maximum concurrent database lookups for authentication cache misses
# 认证缓存未命中时允许并发回源数据库的最大数量
lookup_concurrency: 64
# Process-local invalid-auth abuse protection. Counts only missing, malformed,
# deprecated-query, and confirmed invalid credentials; valid requests and
# Redis/DB failures do not consume the budget.
# 本机无效鉴权防护:仅统计缺失、格式错误、废弃 query 及确认无效的凭据。
invalid_abuse:
enabled: true
# Invalid attempts per trusted client IP (IPv6 grouped by /64) per window.
# 每个可信客户端 IP(IPv6 按 /64 聚合)在窗口内允许的无效次数。
threshold: 120
window_seconds: 60
block_seconds: 60
# Maximum tracked client identities per process; memory remains bounded.
# 每进程最多跟踪的客户端身份数量,确保内存有界。
capacity: 16384
# =============================================================================
# Dashboard Cache Configuration
-9
View File
@@ -935,10 +935,6 @@ export interface OpsErrorLog {
request_type?: number | null
user_agent?: string
// 已删除 KEY 所有者(INVALID_API_KEY 归因快照):认证失败行 user_id 为空,
// 用户列以此回退显示所有者
deleted_key_owner_user_id?: number | null
deleted_key_owner_email?: string | null
}
export interface OpsErrorDetail extends OpsErrorLog {
@@ -958,11 +954,6 @@ export interface OpsErrorDetail extends OpsErrorLog {
is_business_limited: boolean
// Deleted key owner info (INVALID_API_KEY attribution);
// owner user_id/email 已上移到 OpsErrorLog(列表用户列回退)
attempted_key_prefix?: string | null
deleted_key_name?: string | null
// Bound (non-deleted) key prefix, snapshotted at error time
api_key_prefix?: string | null
}
@@ -381,8 +381,6 @@ export default {
suggestPlatform: 'Platform error: prioritize investigation and fix',
suggestGeneric: 'See details for more context',
apiKeyPrefix: 'Key Prefix',
attemptedKeyPrefix: 'Attempted Key Prefix',
deletedKeyOwner: 'Deleted Key Owner',
keyDeletedBadge: 'Key Deleted'
},
requestDetails: {
@@ -712,8 +710,6 @@ export default {
ignoreContextCanceledHint: 'When enabled, client disconnect (context canceled) errors will not be written to the error log.',
ignoreNoAvailableAccounts: 'Ignore no available accounts errors',
ignoreNoAvailableAccountsHint: 'When enabled, "No available accounts" errors will not be written to the error log (not recommended; usually a config issue).',
ignoreInvalidApiKeyErrors: 'Ignore invalid API key errors',
ignoreInvalidApiKeyErrorsHint: 'When enabled, invalid or missing API key errors (INVALID_API_KEY, API_KEY_REQUIRED) will not be written to the error log.',
ignoreInsufficientBalanceErrors: 'Ignore Insufficient Balance Errors',
ignoreInsufficientBalanceErrorsHint: 'When enabled, insufficient account balance errors will not be written to the error log.',
autoRefresh: 'Auto Refresh',
@@ -381,8 +381,6 @@ export default {
suggestPlatform: '🚨 平台错误,建议立即排查修复',
suggestGeneric: '查看详情了解更多信息',
apiKeyPrefix: 'Key 前缀',
attemptedKeyPrefix: '尝试的 Key 前缀',
deletedKeyOwner: '已删除 Key 所有者',
keyDeletedBadge: 'Key 已删除'
},
requestDetails: {
@@ -713,8 +711,6 @@ export default {
'启用后,客户端主动断开连接(context canceled)的错误将不会写入错误日志。',
ignoreNoAvailableAccounts: '忽略无可用账号错误',
ignoreNoAvailableAccountsHint: '启用后,"No available accounts" 错误将不会写入错误日志(不推荐,这通常是配置问题)。',
ignoreInvalidApiKeyErrors: '忽略无效 API Key 错误',
ignoreInvalidApiKeyErrorsHint: '启用后,无效或缺失 API Key 的错误(INVALID_API_KEY、API_KEY_REQUIRED)将不会写入错误日志。',
ignoreInsufficientBalanceErrors: '忽略余额不足错误',
ignoreInsufficientBalanceErrorsHint: '启用后,账号余额不足(Insufficient balance)的错误将不会写入错误日志。',
autoRefresh: '自动刷新',
@@ -114,23 +114,6 @@
</div>
</div>
<div v-if="detail.attempted_key_prefix" class="rounded-xl bg-gray-50 p-4 dark:bg-dark-900">
<div class="text-xs font-bold uppercase tracking-wider text-gray-400">{{ t('admin.ops.errorDetail.attemptedKeyPrefix') }}</div>
<div class="mt-1 font-mono text-sm font-medium text-gray-900 dark:text-white">
{{ detail.attempted_key_prefix }}
</div>
</div>
<div v-if="detail.deleted_key_owner_email" class="rounded-xl bg-gray-50 p-4 dark:bg-dark-900">
<div class="text-xs font-bold uppercase tracking-wider text-gray-400">{{ t('admin.ops.errorDetail.deletedKeyOwner') }}</div>
<div class="mt-1 text-sm font-medium text-gray-900 dark:text-white">
{{ detail.deleted_key_owner_email }}
<span v-if="detail.deleted_key_name" class="ml-1 text-xs text-gray-500 dark:text-gray-400">({{ detail.deleted_key_name }})</span>
<span class="ml-2 inline-flex items-center rounded px-1.5 py-0.5 text-[10px] font-bold ring-1 ring-inset bg-red-50 text-red-700 ring-red-600/20 dark:bg-red-900/30 dark:text-red-400 dark:ring-red-500/30">
{{ t('admin.ops.errorDetail.keyDeletedBadge') }}
</span>
</div>
</div>
</div>
<!-- Response content (client request -> error_body; upstream -> upstream_error_detail/message) -->
@@ -77,19 +77,6 @@
<span v-else class="font-medium text-gray-900 dark:text-white">{{ row.user_email || '-' }}</span>
<span class="ml-1 text-gray-500 dark:text-gray-400">#{{ row.user_id }}</span>
</div>
<!-- 认证失败行 user_id 为空:回退显示已删除 KEY 所有者(归因快照,与详情弹窗一致) -->
<div v-else-if="row.deleted_key_owner_user_id" class="text-sm">
<button
v-if="userClickable && row.deleted_key_owner_email"
class="font-medium text-primary-600 underline decoration-dashed underline-offset-2 transition-colors hover:text-primary-700 dark:text-primary-400 dark:hover:text-primary-300"
:title="t('admin.usage.clickToViewBalance')"
@click.stop="emit('userClick', row.deleted_key_owner_user_id, row.deleted_key_owner_email ?? undefined)"
>
{{ row.deleted_key_owner_email }}
</button>
<span v-else class="font-medium text-gray-900 dark:text-white">{{ row.deleted_key_owner_email || '-' }}</span>
<span class="ml-1 text-gray-500 dark:text-gray-400">#{{ row.deleted_key_owner_user_id }}</span>
</div>
<span v-else class="text-sm text-gray-400 dark:text-gray-500">-</span>
</template>

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