mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 14:08:14 +08:00
Merge pull request #4515 from BenjaminAaron196/feat/filter-noise-rejected-requests
(fix) 过滤入口拒绝日志并强化鉴权安全边界
This commit is contained in:
@@ -0,0 +1,20 @@
|
||||
# Ingress rejection log cleanup
|
||||
|
||||
This maintenance command removes historical admission rejections from
|
||||
`ops_error_logs` without matching unrelated authentication or upstream errors.
|
||||
It is a dry run unless `--execute` is supplied, and always requires an explicit
|
||||
RFC3339 cutoff.
|
||||
|
||||
```sh
|
||||
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z
|
||||
go run ./cmd/cleanup-ingress-reject-logs --before 2026-07-17T00:00:00Z --execute
|
||||
```
|
||||
|
||||
Run the execute form only after every application instance has been upgraded so
|
||||
older instances cannot add new ingress rejection rows below the chosen cutoff.
|
||||
The classifier intentionally retains invariant failures such as
|
||||
`USER_NOT_FOUND`, database errors, quota/billing errors, and upstream failures.
|
||||
|
||||
After the rollout and cleanup are verified, run
|
||||
`backend/scripts/finalize-ingress-reject-cleanup.sql` in a maintenance window to
|
||||
remove the deprecated plaintext-key audit table and attribution columns.
|
||||
@@ -0,0 +1,218 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "github.com/Wei-Shaw/sub2api/ent/runtime"
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/repository"
|
||||
"github.com/lib/pq"
|
||||
)
|
||||
|
||||
const classifierVersion = "ingress-reject-v1"
|
||||
|
||||
type candidate struct {
|
||||
id int64
|
||||
statusCode int
|
||||
message string
|
||||
body string
|
||||
}
|
||||
|
||||
func main() {
|
||||
beforeRaw := flag.String("before", "", "required RFC3339 cutoff; only older rows are considered")
|
||||
execute := flag.Bool("execute", false, "delete matched rows (default is dry-run)")
|
||||
batchSize := flag.Int("batch-size", 5000, "scan/delete batch size (1-5000)")
|
||||
flag.Parse()
|
||||
|
||||
if *beforeRaw == "" {
|
||||
log.Fatal("--before is required")
|
||||
}
|
||||
before, err := time.Parse(time.RFC3339, *beforeRaw)
|
||||
if err != nil {
|
||||
log.Fatalf("invalid --before: %v", err)
|
||||
}
|
||||
if *batchSize < 1 || *batchSize > 5000 {
|
||||
log.Fatal("--batch-size must be between 1 and 5000")
|
||||
}
|
||||
|
||||
cfg, err := config.LoadForBootstrap()
|
||||
if err != nil {
|
||||
log.Fatalf("load config: %v", err)
|
||||
}
|
||||
client, db, err := repository.InitEnt(cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("initialize database: %v", err)
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
ctx := context.Background()
|
||||
counts, scanned, matched, deleted, err := cleanup(ctx, db, before, *batchSize, *execute)
|
||||
if err != nil {
|
||||
log.Fatalf("cleanup failed: %v", err)
|
||||
}
|
||||
|
||||
digest := sha256.Sum256([]byte(classifierVersion))
|
||||
mode := "dry-run"
|
||||
if *execute {
|
||||
mode = "execute"
|
||||
}
|
||||
fmt.Printf("mode=%s before=%s classifier=%s scanned=%d matched=%d deleted=%d\n",
|
||||
mode, before.UTC().Format(time.RFC3339), hex.EncodeToString(digest[:]), scanned, matched, deleted)
|
||||
reasons := make([]string, 0, len(counts))
|
||||
for reason := range counts {
|
||||
reasons = append(reasons, reason)
|
||||
}
|
||||
sort.Strings(reasons)
|
||||
for _, reason := range reasons {
|
||||
fmt.Printf("reason=%s count=%d\n", reason, counts[reason])
|
||||
}
|
||||
if *execute && deleted > 0 {
|
||||
fmt.Println("cleanup complete; schedule VACUUM (ANALYZE) ops_error_logs during normal maintenance")
|
||||
}
|
||||
}
|
||||
|
||||
func cleanup(ctx context.Context, db *sql.DB, before time.Time, batchSize int, execute bool) (map[string]int64, int64, int64, int64, error) {
|
||||
counts := make(map[string]int64)
|
||||
var cursor, scanned, matched, deleted int64
|
||||
for {
|
||||
rows, err := db.QueryContext(ctx, `
|
||||
SELECT id, COALESCE(status_code, 0), COALESCE(error_message, ''), COALESCE(error_body, '')
|
||||
FROM ops_error_logs
|
||||
WHERE id > $1
|
||||
AND created_at < $2
|
||||
AND error_phase = 'auth'
|
||||
AND account_id IS NULL
|
||||
AND upstream_status_code IS NULL
|
||||
AND COALESCE(upstream_error_message, '') = ''
|
||||
AND COALESCE(upstream_error_detail, '') = ''
|
||||
ORDER BY id ASC
|
||||
LIMIT $3`, cursor, before, batchSize)
|
||||
if err != nil {
|
||||
return nil, scanned, matched, deleted, err
|
||||
}
|
||||
|
||||
batch := make([]candidate, 0, batchSize)
|
||||
for rows.Next() {
|
||||
var item candidate
|
||||
if err := rows.Scan(&item.id, &item.statusCode, &item.message, &item.body); err != nil {
|
||||
_ = rows.Close()
|
||||
return nil, scanned, matched, deleted, err
|
||||
}
|
||||
batch = append(batch, item)
|
||||
cursor = item.id
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
_ = rows.Close()
|
||||
return nil, scanned, matched, deleted, err
|
||||
}
|
||||
_ = rows.Close()
|
||||
if len(batch) == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
ids := make([]int64, 0, len(batch))
|
||||
for _, item := range batch {
|
||||
scanned++
|
||||
if reason, ok := historicalIngressRejectReason(item); ok {
|
||||
matched++
|
||||
counts[reason]++
|
||||
ids = append(ids, item.id)
|
||||
}
|
||||
}
|
||||
if execute && len(ids) > 0 {
|
||||
result, err := db.ExecContext(ctx,
|
||||
`DELETE FROM ops_error_logs WHERE id = ANY($1) AND created_at < $2`, pq.Array(ids), before)
|
||||
if err != nil {
|
||||
return nil, scanned, matched, deleted, err
|
||||
}
|
||||
n, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, scanned, matched, deleted, err
|
||||
}
|
||||
deleted += n
|
||||
}
|
||||
}
|
||||
return counts, scanned, matched, deleted, nil
|
||||
}
|
||||
|
||||
func historicalIngressRejectReason(item candidate) (string, bool) {
|
||||
code, message := parseErrorIdentity(item.body, item.message)
|
||||
switch code {
|
||||
case "API_KEY_REQUIRED":
|
||||
return "missing_key", true
|
||||
case "INVALID_API_KEY":
|
||||
return "invalid_key", true
|
||||
case "API_KEY_DISABLED":
|
||||
return "key_disabled", true
|
||||
case "USER_INACTIVE":
|
||||
return "user_inactive", true
|
||||
case "GROUP_DELETED":
|
||||
return "group_deleted", true
|
||||
case "GROUP_DISABLED":
|
||||
return "group_disabled", true
|
||||
case "GROUP_NOT_ALLOWED":
|
||||
return "group_forbidden", true
|
||||
case "ACCESS_DENIED":
|
||||
return "ip_acl_denied", true
|
||||
case "api_key_in_query_deprecated":
|
||||
return "query_key_deprecated", true
|
||||
}
|
||||
|
||||
normalized := strings.TrimSpace(message)
|
||||
switch {
|
||||
case normalized == "API key is required":
|
||||
return "missing_key", true
|
||||
case normalized == "Invalid API key":
|
||||
return "invalid_key", true
|
||||
case normalized == "API key is disabled":
|
||||
return "key_disabled", true
|
||||
case normalized == "User account is not active":
|
||||
return "user_inactive", true
|
||||
case normalized == "API Key 所属分组已删除":
|
||||
return "group_deleted", true
|
||||
case normalized == "API Key 所属分组已停用":
|
||||
return "group_disabled", true
|
||||
case normalized == "API Key 所属专属分组不再允许当前用户使用":
|
||||
return "group_forbidden", true
|
||||
case normalized == "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.":
|
||||
return "group_unassigned", true
|
||||
case strings.HasPrefix(normalized, "Access denied. Your IP is "):
|
||||
return "ip_acl_denied", true
|
||||
case normalized == "Query parameter api_key is deprecated. Use Authorization header or key instead.":
|
||||
return "query_key_deprecated", true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
func parseErrorIdentity(body, fallbackMessage string) (string, string) {
|
||||
var payload struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Error struct {
|
||||
Code json.RawMessage `json:"code"`
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(body), &payload); err != nil {
|
||||
return "", fallbackMessage
|
||||
}
|
||||
message := payload.Message
|
||||
if message == "" {
|
||||
message = payload.Error.Message
|
||||
}
|
||||
if message == "" {
|
||||
message = fallbackMessage
|
||||
}
|
||||
return strings.TrimSpace(payload.Code), message
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package main
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestHistoricalIngressRejectReason(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
item candidate
|
||||
reason string
|
||||
match bool
|
||||
}{
|
||||
{name: "standard invalid key", item: candidate{body: `{"code":"INVALID_API_KEY","message":"Invalid API key"}`}, reason: "invalid_key", match: true},
|
||||
{name: "google missing key", item: candidate{body: `{"error":{"code":401,"message":"API key is required","status":"UNAUTHENTICATED"}}`}, reason: "missing_key", match: true},
|
||||
{name: "google group deleted", item: candidate{body: `{"error":{"code":403,"message":"API Key 所属分组已删除","status":"PERMISSION_DENIED"}}`}, reason: "group_deleted", match: true},
|
||||
{name: "ip acl", item: candidate{body: `{"code":"ACCESS_DENIED","message":"Access denied. Your IP is 192.0.2.1"}`}, reason: "ip_acl_denied", match: true},
|
||||
{name: "user not found remains", item: candidate{body: `{"code":"USER_NOT_FOUND","message":"User associated with API key not found"}`}, match: false},
|
||||
{name: "quota remains", item: candidate{body: `{"code":"API_KEY_QUOTA_EXHAUSTED","message":"quota"}`}, match: false},
|
||||
{name: "database failure remains", item: candidate{statusCode: 500, message: "Failed to validate API key", body: `{"code":"INTERNAL_ERROR","message":"Failed to validate API key"}`}, match: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
reason, ok := historicalIngressRejectReason(tt.item)
|
||||
if ok != tt.match || reason != tt.reason {
|
||||
t.Fatalf("got (%q, %v), want (%q, %v)", reason, ok, tt.reason, tt.match)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -81,6 +81,10 @@ func provideCleanup(
|
||||
opsCleanup *service.OpsCleanupService,
|
||||
opsScheduledReport *service.OpsScheduledReportService,
|
||||
opsSystemLogSink *service.OpsSystemLogSink,
|
||||
opsService *service.OpsService,
|
||||
opsIngressReject *service.OpsIngressRejectAggregator,
|
||||
apiKeyService *service.APIKeyService,
|
||||
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
|
||||
schedulerSnapshot *service.SchedulerSnapshotService,
|
||||
tokenRefresh *service.TokenRefreshService,
|
||||
accountExpiry *service.AccountExpiryService,
|
||||
@@ -121,6 +125,30 @@ func provideCleanup(
|
||||
|
||||
// 应用层清理步骤可并行执行,基础设施资源(Redis/Ent)最后按顺序关闭。
|
||||
parallelSteps := []cleanupStep{
|
||||
{"OpsIngressRejectAggregator", func() error {
|
||||
if opsIngressReject != nil {
|
||||
opsIngressReject.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"AuthCacheInvalidationWorker", func() error {
|
||||
if authCacheInvalidationWorker != nil {
|
||||
authCacheInvalidationWorker.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"AuthCacheInvalidationSubscriber", func() error {
|
||||
if apiKeyService != nil {
|
||||
apiKeyService.StopAuthCacheInvalidationSubscriber()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsRuntimeSettingsRefresh", func() error {
|
||||
if opsService != nil {
|
||||
opsService.StopRuntimeSettingsRefresh()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"PromptAuditService", func() error {
|
||||
if promptAudit != nil {
|
||||
return promptAudit.Shutdown(ctx)
|
||||
|
||||
@@ -157,7 +157,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
antigravityGatewayService := service.NewAntigravityGatewayService(accountRepository, gatewayCache, schedulerSnapshotService, antigravityTokenProvider, rateLimitService, httpUpstream, settingService, internal500CounterCache)
|
||||
geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
|
||||
opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
|
||||
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
|
||||
authCacheInvalidationOutboxRepository := repository.NewAuthCacheInvalidationOutboxRepository(db)
|
||||
authCacheInvalidationWorker := service.ProvideAuthCacheInvalidationWorker(authCacheInvalidationOutboxRepository, apiKeyCache, apiKeyService)
|
||||
opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService, authCacheInvalidationWorker, apiKeyService)
|
||||
usageHandler := handler.NewUsageHandler(usageService, apiKeyService, opsService, settingService)
|
||||
redeemHandler := handler.NewRedeemHandler(redeemService)
|
||||
subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
|
||||
@@ -306,6 +308,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
opsAlertEvaluatorService := service.ProvideOpsAlertEvaluatorService(opsService, opsRepository, emailService, redisClient, configConfig, proxyRepository)
|
||||
opsCleanupService := service.ProvideOpsCleanupService(opsRepository, db, redisClient, configConfig, channelMonitorService, settingRepository, opsService)
|
||||
opsScheduledReportService := service.ProvideOpsScheduledReportService(opsService, userService, emailService, redisClient, configConfig)
|
||||
opsIngressRejectAggregator := service.ProvideOpsIngressRejectAggregator(opsRepository, opsService)
|
||||
accountExpiryService := service.ProvideAccountExpiryService(accountRepository)
|
||||
proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository)
|
||||
subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db)
|
||||
@@ -314,7 +317,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db)
|
||||
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
|
||||
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
|
||||
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
|
||||
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
|
||||
application := &Application{
|
||||
Server: httpServer,
|
||||
PromptAudit: promptService,
|
||||
@@ -351,6 +354,10 @@ func provideCleanup(
|
||||
opsCleanup *service.OpsCleanupService,
|
||||
opsScheduledReport *service.OpsScheduledReportService,
|
||||
opsSystemLogSink *service.OpsSystemLogSink,
|
||||
opsService *service.OpsService,
|
||||
opsIngressReject *service.OpsIngressRejectAggregator,
|
||||
apiKeyService *service.APIKeyService,
|
||||
authCacheInvalidationWorker *service.AuthCacheInvalidationWorker,
|
||||
schedulerSnapshot *service.SchedulerSnapshotService,
|
||||
tokenRefresh *service.TokenRefreshService,
|
||||
accountExpiry *service.AccountExpiryService,
|
||||
@@ -390,6 +397,30 @@ func provideCleanup(
|
||||
}
|
||||
|
||||
parallelSteps := []cleanupStep{
|
||||
{"OpsIngressRejectAggregator", func() error {
|
||||
if opsIngressReject != nil {
|
||||
opsIngressReject.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"AuthCacheInvalidationWorker", func() error {
|
||||
if authCacheInvalidationWorker != nil {
|
||||
authCacheInvalidationWorker.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"AuthCacheInvalidationSubscriber", func() error {
|
||||
if apiKeyService != nil {
|
||||
apiKeyService.StopAuthCacheInvalidationSubscriber()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsRuntimeSettingsRefresh", func() error {
|
||||
if opsService != nil {
|
||||
opsService.StopRuntimeSettingsRefresh()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"PromptAuditService", func() error {
|
||||
if promptAudit != nil {
|
||||
return promptAudit.Shutdown(ctx)
|
||||
|
||||
@@ -58,6 +58,10 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
&service.OpsCleanupService{},
|
||||
&service.OpsScheduledReportService{},
|
||||
opsSystemLogSinkSvc,
|
||||
nil, // opsService
|
||||
nil, // opsIngressRejectAggregator
|
||||
nil, // apiKeyService
|
||||
nil, // authCacheInvalidationWorker
|
||||
schedulerSnapshotSvc,
|
||||
tokenRefreshSvc,
|
||||
accountExpirySvc,
|
||||
|
||||
@@ -645,6 +645,7 @@ type ServerConfig struct {
|
||||
EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header
|
||||
FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接
|
||||
ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒)
|
||||
MaxHeaderBytes int `mapstructure:"max_header_bytes"` // 请求头最大字节数(HTTP/2 映射为 header-list 上限)
|
||||
IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒)
|
||||
TrustedProxies []string `mapstructure:"trusted_proxies"` // 可信代理列表(CIDR/IP)
|
||||
MaxRequestBodySize int64 `mapstructure:"max_request_body_size"` // 全局最大请求体限制
|
||||
@@ -796,6 +797,8 @@ type GatewayConfig struct {
|
||||
OpenAIHighEffortFirstOutputTimeoutSeconds int `mapstructure:"openai_high_effort_first_output_timeout_seconds"`
|
||||
// 请求体最大字节数,用于网关请求体大小限制
|
||||
MaxBodySize int64 `mapstructure:"max_body_size"`
|
||||
// TextMaxBodySize limits endpoints that cannot carry inline image/video payloads.
|
||||
TextMaxBodySize int64 `mapstructure:"text_max_body_size"`
|
||||
// 非流式上游响应体读取上限(字节),用于防止无界读取导致内存放大
|
||||
UpstreamResponseReadMaxBytes int64 `mapstructure:"upstream_response_read_max_bytes"`
|
||||
// 代理探测响应体读取上限(字节)
|
||||
@@ -1419,12 +1422,22 @@ type RateLimitConfig struct {
|
||||
|
||||
// APIKeyAuthCacheConfig API Key 认证缓存配置
|
||||
type APIKeyAuthCacheConfig struct {
|
||||
L1Size int `mapstructure:"l1_size"`
|
||||
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
|
||||
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
|
||||
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
|
||||
JitterPercent int `mapstructure:"jitter_percent"`
|
||||
Singleflight bool `mapstructure:"singleflight"`
|
||||
L1Size int `mapstructure:"l1_size"`
|
||||
L1TTLSeconds int `mapstructure:"l1_ttl_seconds"`
|
||||
L2TTLSeconds int `mapstructure:"l2_ttl_seconds"`
|
||||
NegativeTTLSeconds int `mapstructure:"negative_ttl_seconds"`
|
||||
JitterPercent int `mapstructure:"jitter_percent"`
|
||||
Singleflight bool `mapstructure:"singleflight"`
|
||||
LookupConcurrency int `mapstructure:"lookup_concurrency"`
|
||||
InvalidAbuse InvalidAuthAbuseConfig `mapstructure:"invalid_abuse"`
|
||||
}
|
||||
|
||||
type InvalidAuthAbuseConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
Threshold int `mapstructure:"threshold"`
|
||||
WindowSeconds int `mapstructure:"window_seconds"`
|
||||
BlockSeconds int `mapstructure:"block_seconds"`
|
||||
Capacity int `mapstructure:"capacity"`
|
||||
}
|
||||
|
||||
// SubscriptionCacheConfig 订阅认证 L1 缓存配置
|
||||
@@ -1698,8 +1711,9 @@ func setDefaults() {
|
||||
viper.SetDefault("server.mode", "release")
|
||||
viper.SetDefault("server.enable_server_timing", false)
|
||||
viper.SetDefault("server.frontend_url", "")
|
||||
viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头
|
||||
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
|
||||
viper.SetDefault("server.read_header_timeout", 10) // 10秒读取请求头
|
||||
viper.SetDefault("server.max_header_bytes", 64*1024)
|
||||
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
|
||||
viper.SetDefault("server.trusted_proxies", []string{})
|
||||
viper.SetDefault("server.max_request_body_size", int64(256*1024*1024))
|
||||
// H2C 默认配置
|
||||
@@ -1983,6 +1997,12 @@ func setDefaults() {
|
||||
viper.SetDefault("api_key_auth_cache.negative_ttl_seconds", 30)
|
||||
viper.SetDefault("api_key_auth_cache.jitter_percent", 10)
|
||||
viper.SetDefault("api_key_auth_cache.singleflight", true)
|
||||
viper.SetDefault("api_key_auth_cache.lookup_concurrency", 64)
|
||||
viper.SetDefault("api_key_auth_cache.invalid_abuse.enabled", true)
|
||||
viper.SetDefault("api_key_auth_cache.invalid_abuse.threshold", 120)
|
||||
viper.SetDefault("api_key_auth_cache.invalid_abuse.window_seconds", 60)
|
||||
viper.SetDefault("api_key_auth_cache.invalid_abuse.block_seconds", 60)
|
||||
viper.SetDefault("api_key_auth_cache.invalid_abuse.capacity", 16384)
|
||||
|
||||
// Subscription auth L1 cache
|
||||
viper.SetDefault("subscription_cache.l1_size", 16384)
|
||||
@@ -2111,6 +2131,7 @@ func setDefaults() {
|
||||
viper.SetDefault("gateway.antigravity_fallback_cooldown_minutes", 1)
|
||||
viper.SetDefault("gateway.antigravity_extra_retries", 10)
|
||||
viper.SetDefault("gateway.max_body_size", int64(256*1024*1024))
|
||||
viper.SetDefault("gateway.text_max_body_size", int64(32*1024*1024))
|
||||
viper.SetDefault("gateway.upstream_response_read_max_bytes", DefaultUpstreamResponseReadMaxBytes)
|
||||
viper.SetDefault("gateway.proxy_probe_response_read_max_bytes", int64(1024*1024))
|
||||
viper.SetDefault("gateway.gemini_debug_response_headers", false)
|
||||
@@ -2208,6 +2229,49 @@ func setDefaults() {
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
if c.Server.ReadHeaderTimeout < 1 || c.Server.ReadHeaderTimeout > 60 {
|
||||
return fmt.Errorf("server.read_header_timeout must be between 1 and 60 seconds")
|
||||
}
|
||||
if c.Server.MaxHeaderBytes < 8*1024 || c.Server.MaxHeaderBytes > 1024*1024 {
|
||||
return fmt.Errorf("server.max_header_bytes must be between 8192 and 1048576 bytes")
|
||||
}
|
||||
if c.Server.IdleTimeout <= 0 {
|
||||
return fmt.Errorf("server.idle_timeout must be positive")
|
||||
}
|
||||
if c.Server.MaxRequestBodySize < 0 {
|
||||
return fmt.Errorf("server.max_request_body_size must be non-negative")
|
||||
}
|
||||
if c.Server.H2C.Enabled {
|
||||
if c.Server.H2C.MaxConcurrentStreams == 0 {
|
||||
return fmt.Errorf("server.h2c.max_concurrent_streams must be positive")
|
||||
}
|
||||
if c.Server.H2C.IdleTimeout <= 0 {
|
||||
return fmt.Errorf("server.h2c.idle_timeout must be positive")
|
||||
}
|
||||
if c.Server.H2C.MaxReadFrameSize < 16*1024 || c.Server.H2C.MaxReadFrameSize > 16*1024*1024-1 {
|
||||
return fmt.Errorf("server.h2c.max_read_frame_size must be between 16384 and 16777215 bytes")
|
||||
}
|
||||
if c.Server.H2C.MaxUploadBufferPerConnection < 65535 {
|
||||
return fmt.Errorf("server.h2c.max_upload_buffer_per_connection must be at least 65535 bytes")
|
||||
}
|
||||
if c.Server.H2C.MaxUploadBufferPerStream <= 0 {
|
||||
return fmt.Errorf("server.h2c.max_upload_buffer_per_stream must be positive")
|
||||
}
|
||||
}
|
||||
if c.APIKeyAuth.InvalidAbuse.Enabled {
|
||||
if c.APIKeyAuth.InvalidAbuse.Threshold < 10 {
|
||||
return fmt.Errorf("api_key_auth_cache.invalid_abuse.threshold must be at least 10")
|
||||
}
|
||||
if c.APIKeyAuth.InvalidAbuse.WindowSeconds < 1 || c.APIKeyAuth.InvalidAbuse.WindowSeconds > 3600 {
|
||||
return fmt.Errorf("api_key_auth_cache.invalid_abuse.window_seconds must be between 1 and 3600")
|
||||
}
|
||||
if c.APIKeyAuth.InvalidAbuse.BlockSeconds < 1 || c.APIKeyAuth.InvalidAbuse.BlockSeconds > 3600 {
|
||||
return fmt.Errorf("api_key_auth_cache.invalid_abuse.block_seconds must be between 1 and 3600")
|
||||
}
|
||||
if c.APIKeyAuth.InvalidAbuse.Capacity < 256 || c.APIKeyAuth.InvalidAbuse.Capacity > 1_000_000 {
|
||||
return fmt.Errorf("api_key_auth_cache.invalid_abuse.capacity must be between 256 and 1000000")
|
||||
}
|
||||
}
|
||||
jwtSecret := strings.TrimSpace(c.JWT.Secret)
|
||||
if jwtSecret == "" {
|
||||
return fmt.Errorf("jwt.secret is required")
|
||||
@@ -2733,6 +2797,9 @@ func (c *Config) Validate() error {
|
||||
if c.Gateway.MaxBodySize <= 0 {
|
||||
return fmt.Errorf("gateway.max_body_size must be positive")
|
||||
}
|
||||
if c.Gateway.TextMaxBodySize <= 0 || c.Gateway.TextMaxBodySize > c.Gateway.MaxBodySize {
|
||||
return fmt.Errorf("gateway.text_max_body_size must be positive and no greater than gateway.max_body_size")
|
||||
}
|
||||
if c.Gateway.UpstreamResponseReadMaxBytes <= 0 {
|
||||
return fmt.Errorf("gateway.upstream_response_read_max_bytes must be positive")
|
||||
}
|
||||
|
||||
@@ -35,6 +35,18 @@ func TestLoadServerTimingConfig(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoadHTTPIngressSafetyDefaults(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 10, cfg.Server.ReadHeaderTimeout)
|
||||
require.Equal(t, 64*1024, cfg.Server.MaxHeaderBytes)
|
||||
require.Equal(t, int64(32*1024*1024), cfg.Gateway.TextMaxBodySize)
|
||||
require.True(t, cfg.APIKeyAuth.InvalidAbuse.Enabled)
|
||||
require.Equal(t, 120, cfg.APIKeyAuth.InvalidAbuse.Threshold)
|
||||
require.Equal(t, 16384, cfg.APIKeyAuth.InvalidAbuse.Capacity)
|
||||
}
|
||||
|
||||
func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) {
|
||||
viper.Reset()
|
||||
t.Setenv("JWT_SECRET", "")
|
||||
@@ -1185,6 +1197,51 @@ func TestValidateConfigErrors(t *testing.T) {
|
||||
mutate func(*Config)
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "server read header timeout",
|
||||
mutate: func(c *Config) { c.Server.ReadHeaderTimeout = 0 },
|
||||
wantErr: "server.read_header_timeout",
|
||||
},
|
||||
{
|
||||
name: "server max header bytes too small",
|
||||
mutate: func(c *Config) { c.Server.MaxHeaderBytes = 4096 },
|
||||
wantErr: "server.max_header_bytes",
|
||||
},
|
||||
{
|
||||
name: "server max request body size",
|
||||
mutate: func(c *Config) { c.Server.MaxRequestBodySize = -1 },
|
||||
wantErr: "server.max_request_body_size",
|
||||
},
|
||||
{
|
||||
name: "h2c zero concurrent streams",
|
||||
mutate: func(c *Config) {
|
||||
c.Server.H2C.Enabled = true
|
||||
c.Server.H2C.MaxConcurrentStreams = 0
|
||||
},
|
||||
wantErr: "server.h2c.max_concurrent_streams",
|
||||
},
|
||||
{
|
||||
name: "h2c oversized read frame",
|
||||
mutate: func(c *Config) {
|
||||
c.Server.H2C.Enabled = true
|
||||
c.Server.H2C.MaxReadFrameSize = 16 * 1024 * 1024
|
||||
},
|
||||
wantErr: "server.h2c.max_read_frame_size",
|
||||
},
|
||||
{
|
||||
name: "invalid auth abuse threshold too small",
|
||||
mutate: func(c *Config) {
|
||||
c.APIKeyAuth.InvalidAbuse.Threshold = 9
|
||||
},
|
||||
wantErr: "api_key_auth_cache.invalid_abuse.threshold",
|
||||
},
|
||||
{
|
||||
name: "invalid auth abuse capacity too small",
|
||||
mutate: func(c *Config) {
|
||||
c.APIKeyAuth.InvalidAbuse.Capacity = 255
|
||||
},
|
||||
wantErr: "api_key_auth_cache.invalid_abuse.capacity",
|
||||
},
|
||||
{
|
||||
name: "jwt secret required",
|
||||
mutate: func(c *Config) { c.JWT.Secret = "" },
|
||||
@@ -1386,6 +1443,11 @@ func TestValidateConfigErrors(t *testing.T) {
|
||||
mutate: func(c *Config) { c.Gateway.MaxBodySize = 0 },
|
||||
wantErr: "gateway.max_body_size",
|
||||
},
|
||||
{
|
||||
name: "gateway text body exceeds media body",
|
||||
mutate: func(c *Config) { c.Gateway.TextMaxBodySize = c.Gateway.MaxBodySize + 1 },
|
||||
wantErr: "gateway.text_max_body_size",
|
||||
},
|
||||
{
|
||||
name: "gateway response header timeout",
|
||||
mutate: func(c *Config) { c.Gateway.ResponseHeaderTimeout = -1 },
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetAuthCacheInvalidationHealth exposes durable outbox lag and subscriber health.
|
||||
func (h *OpsHandler) GetAuthCacheInvalidationHealth(c *gin.Context) {
|
||||
if h.opsService == nil {
|
||||
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
|
||||
return
|
||||
}
|
||||
if err := h.opsService.RequireMonitoringEnabled(c.Request.Context()); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, h.opsService.GetAuthCacheInvalidationHealth(c.Request.Context()))
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var ingressRejectReasons = map[string]struct{}{
|
||||
"query_api_key_deprecated": {}, "api_key_required": {}, "invalid_api_key": {},
|
||||
"invalid_auth_rate_limited": {},
|
||||
"api_key_auth_overloaded": {},
|
||||
"api_key_disabled": {}, "ip_restricted": {}, "user_inactive": {}, "group_deleted": {},
|
||||
"group_disabled": {}, "group_not_allowed": {}, "group_unassigned": {}, "other": {},
|
||||
}
|
||||
|
||||
var ingressRejectRouteFamilies = map[string]struct{}{
|
||||
"antigravity": {}, "gemini": {}, "codex": {}, "messages": {}, "responses": {},
|
||||
"chat_completions": {}, "images": {}, "videos": {}, "embeddings": {}, "models": {}, "other": {},
|
||||
}
|
||||
|
||||
var ingressRejectProtocols = map[string]struct{}{
|
||||
"google": {}, "anthropic": {}, "openai": {}, "gateway": {}, "other": {},
|
||||
}
|
||||
|
||||
// ListIngressRejects returns bounded security aggregates, never raw credentials or request bodies.
|
||||
func (h *OpsHandler) ListIngressRejects(c *gin.Context) {
|
||||
if h.opsService == nil {
|
||||
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
|
||||
return
|
||||
}
|
||||
page, pageSize := response.ParsePagination(c)
|
||||
if pageSize > 200 {
|
||||
pageSize = 200
|
||||
}
|
||||
startTime, endTime, err := parseOpsTimeRange(c, "1h")
|
||||
if err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
filter := &service.OpsIngressRejectFilter{Page: page, PageSize: pageSize}
|
||||
if !startTime.IsZero() {
|
||||
filter.StartTime = &startTime
|
||||
}
|
||||
if !endTime.IsZero() {
|
||||
filter.EndTime = &endTime
|
||||
}
|
||||
if filter.RejectReason, err = parseIngressRejectEnum(c, "reason", ingressRejectReasons); err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if filter.RouteFamily, err = parseIngressRejectEnum(c, "route_family", ingressRejectRouteFamilies); err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if filter.Protocol, err = parseIngressRejectEnum(c, "protocol", ingressRejectProtocols); err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if raw := strings.TrimSpace(c.Query("client_ip")); raw != "" {
|
||||
addr, parseErr := netip.ParseAddr(raw)
|
||||
if parseErr != nil {
|
||||
response.BadRequest(c, "Invalid client_ip")
|
||||
return
|
||||
}
|
||||
addr = addr.Unmap()
|
||||
if addr.Is6() {
|
||||
addr = netip.PrefixFrom(addr, 64).Masked().Addr()
|
||||
}
|
||||
filter.ClientIP = addr.String()
|
||||
}
|
||||
if filter.UserID, err = parseOptionalPositiveID(c, "user_id"); err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if filter.APIKeyID, err = parseOptionalPositiveID(c, "api_key_id"); err != nil {
|
||||
response.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.opsService.ListIngressRejects(c.Request.Context(), filter)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *OpsHandler) GetIngressRejectHealth(c *gin.Context) {
|
||||
if h.opsService == nil {
|
||||
response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
|
||||
return
|
||||
}
|
||||
if err := h.opsService.RequireMonitoringEnabled(c.Request.Context()); err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, h.opsService.GetIngressRejectHealth())
|
||||
}
|
||||
|
||||
func parseIngressRejectEnum(c *gin.Context, name string, allowed map[string]struct{}) (string, error) {
|
||||
value := strings.TrimSpace(c.Query(name))
|
||||
if value == "" {
|
||||
return "", nil
|
||||
}
|
||||
if _, ok := allowed[value]; !ok {
|
||||
return "", &ingressRejectQueryError{message: "Invalid " + name}
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func parseOptionalPositiveID(c *gin.Context, name string) (*int64, error) {
|
||||
raw := strings.TrimSpace(c.Query(name))
|
||||
if raw == "" {
|
||||
return nil, nil
|
||||
}
|
||||
value, err := strconv.ParseInt(raw, 10, 64)
|
||||
if err != nil || value <= 0 {
|
||||
return nil, &ingressRejectQueryError{message: "Invalid " + name}
|
||||
}
|
||||
return &value, nil
|
||||
}
|
||||
|
||||
type ingressRejectQueryError struct{ message string }
|
||||
|
||||
func (e *ingressRejectQueryError) Error() string { return e.message }
|
||||
@@ -0,0 +1,41 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newIngressRejectHandlerForTest() *OpsHandler {
|
||||
return NewOpsHandler(service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil))
|
||||
}
|
||||
|
||||
func TestListIngressRejectsValidatesFilters(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, path := range []string{
|
||||
"/api/v1/admin/ops/ingress-rejections?reason=not-valid",
|
||||
"/api/v1/admin/ops/ingress-rejections?client_ip=not-an-ip",
|
||||
"/api/v1/admin/ops/ingress-rejections?user_id=0",
|
||||
"/api/v1/admin/ops/ingress-rejections?api_key_id=-1",
|
||||
} {
|
||||
t.Run(path, func(t *testing.T) {
|
||||
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
context.Request = httptest.NewRequest(http.MethodGet, path, nil)
|
||||
newIngressRejectHandlerForTest().ListIngressRejects(context)
|
||||
require.Equal(t, http.StatusBadRequest, context.Writer.Status())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetIngressRejectHealthAvailable(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
context, _ := gin.CreateTestContext(recorder)
|
||||
context.Request = httptest.NewRequest(http.MethodGet, "/api/v1/admin/ops/ingress-rejections/health", nil)
|
||||
newIngressRejectHandlerForTest().GetIngressRejectHealth(context)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), `"capacity":8192`)
|
||||
}
|
||||
@@ -1688,6 +1688,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
if h.opsService != nil {
|
||||
h.opsService.SetMonitoringEnabled(settings.OpsMonitoringEnabled)
|
||||
}
|
||||
|
||||
// Update OpenAI fast policy (stored under dedicated key, only when provided).
|
||||
if req.OpenAIFastPolicySettings != nil {
|
||||
|
||||
@@ -206,6 +206,10 @@ func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T)
|
||||
require.NotContains(t, job.entry.ErrorMessage, "earlier inference failure")
|
||||
require.NotNil(t, job.entry.UpstreamStatusCode)
|
||||
require.Zero(t, *job.entry.UpstreamStatusCode)
|
||||
require.Len(t, job.entry.UpstreamErrors, 2)
|
||||
require.Equal(t, http.StatusForbidden, job.entry.UpstreamErrors[0].UpstreamStatusCode)
|
||||
require.Nil(t, job.entry.UpstreamErrors)
|
||||
require.NotNil(t, job.entry.UpstreamErrorsJSON)
|
||||
events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 2)
|
||||
require.Equal(t, http.StatusForbidden, events[0].UpstreamStatusCode)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -1706,11 +1764,8 @@ func shouldSkipOpsErrorLog(ctx context.Context, ops *service.OpsService, message
|
||||
}
|
||||
|
||||
// Get advanced settings to check filter configuration
|
||||
settings, err := ops.GetOpsAdvancedSettings(ctx)
|
||||
if err != nil || settings == nil {
|
||||
// If we can't get settings, don't skip (fail open)
|
||||
return false
|
||||
}
|
||||
_ = ctx
|
||||
settings := ops.OpsAdvancedSettingsSnapshot()
|
||||
|
||||
msgLower := strings.ToLower(message)
|
||||
bodyLower := strings.ToLower(body)
|
||||
|
||||
@@ -1,39 +1,9 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func TestLooksLikeSystemKey(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
want bool
|
||||
}{
|
||||
{"sk-abcdef0123456789", true},
|
||||
{"ABCdef_-0123456789", true},
|
||||
{"short", false},
|
||||
{"with space xxxxxxxxxx", false},
|
||||
{"汉字key1234567890", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := looksLikeSystemKey(c.in); got != c.want {
|
||||
t.Errorf("looksLikeSystemKey(%q)=%v want %v", c.in, got, c.want)
|
||||
}
|
||||
}
|
||||
long := make([]byte, 129)
|
||||
for i := range long {
|
||||
long[i] = 'a'
|
||||
}
|
||||
if looksLikeSystemKey(string(long)) {
|
||||
t.Errorf("129-char key should be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyPrefix(t *testing.T) {
|
||||
if got := keyPrefix("sk-3f2a9c7e", 8); got != "sk-3f2a9" {
|
||||
t.Errorf("keyPrefix=%q want %q", got, "sk-3f2a9")
|
||||
@@ -42,77 +12,3 @@ func TestKeyPrefix(t *testing.T) {
|
||||
t.Errorf("short key should be returned as-is, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractAttemptedKey(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
headers map[string]string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "Bearer in Authorization",
|
||||
headers: map[string]string{"Authorization": "Bearer sk-testkey0123456789"},
|
||||
want: "sk-testkey0123456789",
|
||||
},
|
||||
{
|
||||
name: "Bearer case-insensitive",
|
||||
headers: map[string]string{"Authorization": "BEARER sk-testkey0123456789"},
|
||||
want: "sk-testkey0123456789",
|
||||
},
|
||||
{
|
||||
name: "x-api-key header",
|
||||
headers: map[string]string{"x-api-key": "sk-xapikey0123456789"},
|
||||
want: "sk-xapikey0123456789",
|
||||
},
|
||||
{
|
||||
name: "x-goog-api-key header",
|
||||
headers: map[string]string{"x-goog-api-key": "sk-goog0123456789"},
|
||||
want: "sk-goog0123456789",
|
||||
},
|
||||
{
|
||||
name: "Authorization takes priority over x-api-key",
|
||||
headers: map[string]string{"Authorization": "Bearer sk-auth0123456789", "x-api-key": "sk-xapi0123456789"},
|
||||
want: "sk-auth0123456789",
|
||||
},
|
||||
{
|
||||
name: "x-api-key takes priority over x-goog-api-key",
|
||||
headers: map[string]string{"x-api-key": "sk-xapi0123456789", "x-goog-api-key": "sk-goog0123456789"},
|
||||
want: "sk-xapi0123456789",
|
||||
},
|
||||
{
|
||||
name: "no key headers",
|
||||
headers: map[string]string{},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "Bearer with leading/trailing spaces trimmed",
|
||||
headers: map[string]string{"Authorization": "Bearer sk-trimmed0123456789 "},
|
||||
want: "sk-trimmed0123456789",
|
||||
},
|
||||
{
|
||||
// 非 Bearer Authorization 应被忽略,继续 fall-through 到 x-api-key(与认证中间件一致)
|
||||
name: "non-Bearer Authorization falls through to x-api-key",
|
||||
headers: map[string]string{"Authorization": "junk-not-bearer", "x-api-key": "sk-realkey1234567"},
|
||||
want: "sk-realkey1234567",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
|
||||
for k, v := range tc.headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
c.Request = req
|
||||
|
||||
got := extractAttemptedKey(c)
|
||||
if got != tc.want {
|
||||
t.Errorf("extractAttemptedKey(%v) = %q, want %q", tc.headers, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
@@ -12,6 +15,82 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type ingressRejectSettingRepo struct {
|
||||
service.SettingRepository
|
||||
getValueCalls int
|
||||
}
|
||||
|
||||
func (r *ingressRejectSettingRepo) GetValue(context.Context, string) (string, error) {
|
||||
r.getValueCalls++
|
||||
return "", service.ErrSettingNotFound
|
||||
}
|
||||
|
||||
func (r *ingressRejectSettingRepo) GetMultiple(context.Context, []string) (map[string]string, error) {
|
||||
r.getValueCalls++
|
||||
return map[string]string{}, nil
|
||||
}
|
||||
|
||||
func (r *ingressRejectSettingRepo) Set(context.Context, string, string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
type ingressRejectOpsRepo struct {
|
||||
service.OpsRepository
|
||||
insertCalls int
|
||||
}
|
||||
|
||||
func (r *ingressRejectOpsRepo) InsertErrorLog(context.Context, *service.OpsInsertErrorLogInput) (int64, error) {
|
||||
r.insertCalls++
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(context.Context, []*service.OpsInsertErrorLogInput) (int64, error) {
|
||||
r.insertCalls++
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func TestOpsErrorLogQueueByteBudget(t *testing.T) {
|
||||
previousBytes := opsErrorLogQueueBytes.Load()
|
||||
previousLen := opsErrorLogQueueLen.Load()
|
||||
opsErrorLogQueueBytes.Store(0)
|
||||
opsErrorLogQueueLen.Store(0)
|
||||
t.Cleanup(func() {
|
||||
opsErrorLogQueueBytes.Store(previousBytes)
|
||||
opsErrorLogQueueLen.Store(previousLen)
|
||||
})
|
||||
|
||||
if !reserveOpsErrorLogQueueBytes(opsErrorLogMaxQueueBytes - 1) {
|
||||
t.Fatal("first reservation within byte budget should succeed")
|
||||
}
|
||||
if reserveOpsErrorLogQueueBytes(2) {
|
||||
t.Fatal("reservation beyond byte budget should be rejected")
|
||||
}
|
||||
if got := OpsErrorLogQueueBytes(); got != opsErrorLogMaxQueueBytes-1 {
|
||||
t.Fatalf("queued bytes = %d, want %d", got, opsErrorLogMaxQueueBytes-1)
|
||||
}
|
||||
if got := OpsErrorLogQueueLength(); got != 1 {
|
||||
t.Fatalf("queue length = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEstimateOpsErrorLogJobBytesIncludesVariablePayloads(t *testing.T) {
|
||||
base := estimateOpsErrorLogJobBytes(&service.OpsInsertErrorLogInput{})
|
||||
message := "upstream message"
|
||||
detail := "upstream detail"
|
||||
events := `[{"error":"x"}]`
|
||||
entry := &service.OpsInsertErrorLogInput{
|
||||
ErrorBody: strings.Repeat("x", 1024),
|
||||
ErrorMessage: "client error",
|
||||
UserAgent: "test-agent",
|
||||
UpstreamErrorMessage: &message,
|
||||
UpstreamErrorDetail: &detail,
|
||||
UpstreamErrorsJSON: &events,
|
||||
}
|
||||
if got := estimateOpsErrorLogJobBytes(entry); got <= base+1024 {
|
||||
t.Fatalf("estimated bytes = %d, expected variable payloads above %d", got, base+1024)
|
||||
}
|
||||
}
|
||||
|
||||
func resetOpsErrorLoggerStateForTest(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
@@ -119,6 +198,29 @@ func TestOpsCaptureWriterPool_ResetOnRelease(t *testing.T) {
|
||||
require.Zero(t, reused.buf.Len(), "writer should be reset before reuse")
|
||||
}
|
||||
|
||||
func TestOpsCaptureWriterPool_DropsLargeBuffers(t *testing.T) {
|
||||
w := &opsCaptureWriter{}
|
||||
w.buf.Grow(opsCaptureWriterPoolMaxRetainedCapacity + 1)
|
||||
require.False(t, shouldPoolOpsCaptureWriter(w))
|
||||
}
|
||||
|
||||
func TestEnqueueOpsErrorLog_SanitizesAndBoundsBodyBeforeQueue(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 1)
|
||||
ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
secret := strings.Repeat("s", service.OpsErrorLogQueueBodyMaxBytes)
|
||||
entry := &service.OpsInsertErrorLogInput{
|
||||
ErrorPhase: "request",
|
||||
ErrorType: "api_error",
|
||||
ErrorBody: `{"authorization":"Bearer ` + secret + `","message":"failed"}`,
|
||||
}
|
||||
|
||||
enqueueOpsErrorLog(ops, entry)
|
||||
job := <-opsErrorLogQueue
|
||||
require.LessOrEqual(t, len(job.entry.ErrorBody), service.OpsErrorLogQueueBodyMaxBytes)
|
||||
require.NotContains(t, job.entry.ErrorBody, secret)
|
||||
require.Equal(t, int64(1), OpsErrorLogSanitizedTotal())
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_DoesNotBreakOuterMiddlewares(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -149,6 +251,42 @@ func setupOpsErrorLogTestQueue(t *testing.T, size int) {
|
||||
opsErrorLogMu.Unlock()
|
||||
}
|
||||
|
||||
func TestOpsErrorLoggerMiddleware_HardSkipsIngressRejection(t *testing.T) {
|
||||
setupOpsErrorLogTestQueue(t, 4)
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
settings := &ingressRejectSettingRepo{}
|
||||
repo := &ingressRejectOpsRepo{}
|
||||
ops := service.NewOpsService(repo, settings, nil, nil, nil, nil, nil, nil, nil, nil, nil)
|
||||
// Construction may read unrelated runtime settings; only request-path reads matter here.
|
||||
settings.getValueCalls = 0
|
||||
|
||||
router := gin.New()
|
||||
router.Use(OpsErrorLoggerMiddleware(ops))
|
||||
router.GET("/v1/messages", func(c *gin.Context) {
|
||||
middleware2.MarkIngressRejected(c, middleware2.IngressRejectInvalidAPIKey)
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"code": "INVALID_API_KEY", "message": "Invalid API key"})
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
require.JSONEq(t, `{"code":"INVALID_API_KEY","message":"Invalid API key"}`, w.Body.String())
|
||||
require.Zero(t, settings.getValueCalls, "ingress rejection must bypass monitoring settings reads")
|
||||
require.Zero(t, repo.insertCalls, "ingress rejection must bypass inserts")
|
||||
require.Zero(t, OpsErrorLogEnqueuedTotal(), "ingress rejection must not enter the error queue")
|
||||
}
|
||||
|
||||
func TestNormalizeOpsPersistentUserAgentBoundsAndPreservesUTF8(t *testing.T) {
|
||||
value := strings.Repeat("a", opsErrorLogMaxUserAgentBytes-1) + "你" + strings.Repeat("b", 32)
|
||||
got := normalizeOpsPersistentUserAgent(" " + value + " ")
|
||||
require.LessOrEqual(t, len(got), opsErrorLogMaxUserAgentBytes)
|
||||
require.True(t, utf8.ValidString(got))
|
||||
require.NotContains(t, got, "b")
|
||||
}
|
||||
|
||||
// 就地(in-band) SSE 错误挂在已固化的 HTTP 200 流上:wire 状态码为 200,
|
||||
// 常规 status>=400 采集路径不会触发。logOpsStreamError 必须据 MarkOpsStreamError
|
||||
// 补记一条错误日志,且用 IntendedStatus(429) 分级、StatusCode 仍记 wire 的 200。
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -8,40 +8,10 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GetClientIP 从 Gin Context 中提取客户端真实 IP 地址。
|
||||
// 按以下优先级检查 Header:
|
||||
// 1. CF-Connecting-IP (Cloudflare)
|
||||
// 2. X-Real-IP (Nginx)
|
||||
// 3. X-Forwarded-For (取第一个非私有 IP)
|
||||
// 4. c.ClientIP() (Gin 内置方法)
|
||||
// GetClientIP resolves a client address only through Gin's configured trusted
|
||||
// proxy chain. Forwarding headers from a direct or untrusted peer are ignored.
|
||||
func GetClientIP(c *gin.Context) string {
|
||||
// 1. Cloudflare
|
||||
if ip := c.GetHeader("CF-Connecting-IP"); ip != "" {
|
||||
return normalizeIP(ip)
|
||||
}
|
||||
|
||||
// 2. Nginx X-Real-IP
|
||||
if ip := c.GetHeader("X-Real-IP"); ip != "" {
|
||||
return normalizeIP(ip)
|
||||
}
|
||||
|
||||
// 3. X-Forwarded-For (多个 IP 时取第一个公网 IP)
|
||||
if xff := c.GetHeader("X-Forwarded-For"); xff != "" {
|
||||
ips := strings.Split(xff, ",")
|
||||
for _, ip := range ips {
|
||||
ip = strings.TrimSpace(ip)
|
||||
if ip != "" && !isPrivateIP(ip) {
|
||||
return normalizeIP(ip)
|
||||
}
|
||||
}
|
||||
// 如果都是私有 IP,返回第一个
|
||||
if len(ips) > 0 {
|
||||
return normalizeIP(strings.TrimSpace(ips[0]))
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Gin 内置方法
|
||||
return normalizeIP(c.ClientIP())
|
||||
return GetTrustedClientIP(c)
|
||||
}
|
||||
|
||||
// GetTrustedClientIP 从 Gin 的可信代理解析链提取客户端 IP。
|
||||
@@ -54,14 +24,10 @@ func GetTrustedClientIP(c *gin.Context) string {
|
||||
return normalizeIP(c.ClientIP())
|
||||
}
|
||||
|
||||
// GetSecurityClientIP 返回安全敏感场景(API Key IP 限制、审计日志、会话 IP/UA 绑定)
|
||||
// 使用的客户端 IP。trustForwarded 对应系统设置「信任反代传递的客户端 IP」:
|
||||
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
|
||||
// 关闭时走 Gin trusted_proxies 解析链。
|
||||
func GetSecurityClientIP(c *gin.Context, trustForwarded bool) string {
|
||||
if trustForwarded {
|
||||
return GetClientIP(c)
|
||||
}
|
||||
// GetSecurityClientIP returns the address resolved through Gin's configured
|
||||
// trusted-proxy chain. The legacy toggle is retained for configuration/API
|
||||
// compatibility, but never makes raw forwarding headers trustworthy by itself.
|
||||
func GetSecurityClientIP(c *gin.Context, _ bool) string {
|
||||
return GetTrustedClientIP(c)
|
||||
}
|
||||
|
||||
@@ -75,9 +41,6 @@ func normalizeIP(ip string) string {
|
||||
return ip
|
||||
}
|
||||
|
||||
// privateNets 预编译私有 IP CIDR 块,避免每次调用 isPrivateIP 时重复解析
|
||||
var privateNets []*net.IPNet
|
||||
|
||||
// CompiledIPRules 表示预编译的 IP 匹配规则。
|
||||
// PatternCount 记录原始规则数量,用于保留“规则存在但全无效”时的行为语义。
|
||||
type CompiledIPRules struct {
|
||||
@@ -86,23 +49,6 @@ type CompiledIPRules struct {
|
||||
PatternCount int
|
||||
}
|
||||
|
||||
func init() {
|
||||
for _, cidr := range []string{
|
||||
"10.0.0.0/8",
|
||||
"172.16.0.0/12",
|
||||
"192.168.0.0/16",
|
||||
"127.0.0.0/8",
|
||||
"::1/128",
|
||||
"fc00::/7",
|
||||
} {
|
||||
_, block, err := net.ParseCIDR(cidr)
|
||||
if err != nil {
|
||||
panic("invalid CIDR: " + cidr)
|
||||
}
|
||||
privateNets = append(privateNets, block)
|
||||
}
|
||||
}
|
||||
|
||||
// CompileIPRules 将 IP/CIDR 字符串规则预编译为可复用结构。
|
||||
// 非法规则会被忽略,但 PatternCount 会保留原始规则条数。
|
||||
func CompileIPRules(patterns []string) *CompiledIPRules {
|
||||
@@ -150,20 +96,6 @@ func matchesCompiledRules(parsedIP net.IP, rules *CompiledIPRules) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// isPrivateIP 检查 IP 是否为私有地址。
|
||||
func isPrivateIP(ipStr string) bool {
|
||||
ip := net.ParseIP(ipStr)
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
for _, block := range privateNets {
|
||||
if block.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// MatchesPattern 检查 IP 是否匹配指定的模式(支持单个 IP 或 CIDR)。
|
||||
// pattern 可以是:
|
||||
// - 单个 IP: "192.168.1.100"
|
||||
|
||||
@@ -10,48 +10,6 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIsPrivateIP(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ip string
|
||||
expected bool
|
||||
}{
|
||||
// 私有 IPv4
|
||||
{"10.x 私有地址", "10.0.0.1", true},
|
||||
{"10.x 私有地址段末", "10.255.255.255", true},
|
||||
{"172.16.x 私有地址", "172.16.0.1", true},
|
||||
{"172.31.x 私有地址", "172.31.255.255", true},
|
||||
{"192.168.x 私有地址", "192.168.1.1", true},
|
||||
{"127.0.0.1 本地回环", "127.0.0.1", true},
|
||||
{"127.x 回环段", "127.255.255.255", true},
|
||||
|
||||
// 公网 IPv4
|
||||
{"8.8.8.8 公网 DNS", "8.8.8.8", false},
|
||||
{"1.1.1.1 公网", "1.1.1.1", false},
|
||||
{"172.15.255.255 非私有", "172.15.255.255", false},
|
||||
{"172.32.0.0 非私有", "172.32.0.0", false},
|
||||
{"11.0.0.1 公网", "11.0.0.1", false},
|
||||
|
||||
// IPv6
|
||||
{"::1 IPv6 回环", "::1", true},
|
||||
{"fc00:: IPv6 私有", "fc00::1", true},
|
||||
{"fd00:: IPv6 私有", "fd00::1", true},
|
||||
{"2001:db8::1 IPv6 公网", "2001:db8::1", false},
|
||||
|
||||
// 无效输入
|
||||
{"空字符串", "", false},
|
||||
{"非法字符串", "not-an-ip", false},
|
||||
{"不完整 IP", "192.168", false},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := isPrivateIP(tc.ip)
|
||||
require.Equal(t, tc.expected, got, "isPrivateIP(%q)", tc.ip)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetTrustedClientIPUsesGinClientIP(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -95,7 +53,7 @@ func TestCheckIPRestrictionWithCompiledRules_InvalidWhitelistStillDenies(t *test
|
||||
require.Equal(t, "access denied", reason)
|
||||
}
|
||||
|
||||
func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
|
||||
func TestGetSecurityClientIPNeverTrustsHeadersFromUntrustedPeer(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
for _, tc := range []struct {
|
||||
@@ -103,8 +61,8 @@ func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
|
||||
trustForwarded bool
|
||||
want string
|
||||
}{
|
||||
{name: "trust disabled uses trusted proxy chain", trustForwarded: false, want: "9.9.9.9"},
|
||||
{name: "trust enabled uses forwarded header", trustForwarded: true, want: "1.2.3.4"},
|
||||
{name: "legacy toggle disabled", trustForwarded: false, want: "9.9.9.9"},
|
||||
{name: "legacy toggle enabled", trustForwarded: true, want: "9.9.9.9"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r := gin.New()
|
||||
@@ -124,3 +82,18 @@ func TestGetSecurityClientIPHonorsTrustToggle(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSecurityClientIPUsesConfiguredTrustedProxy(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
require.NoError(t, r.SetTrustedProxies([]string{"9.9.9.9"}))
|
||||
r.GET("/t", func(c *gin.Context) { c.String(200, GetSecurityClientIP(c, true)) })
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest("GET", "/t", nil)
|
||||
req.RemoteAddr = "9.9.9.9:12345"
|
||||
req.Header.Set("X-Forwarded-For", "1.2.3.4")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
require.Equal(t, "1.2.3.4", w.Body.String())
|
||||
}
|
||||
|
||||
@@ -26,6 +26,10 @@ const (
|
||||
LevelWarn = zapcore.WarnLevel
|
||||
LevelError = zapcore.ErrorLevel
|
||||
LevelFatal = zapcore.FatalLevel
|
||||
|
||||
// OpsSystemLogSkipField keeps an event in the standard logger while
|
||||
// preventing the database-backed Ops system-log sink from indexing it.
|
||||
OpsSystemLogSkipField = "ops_system_log_skip"
|
||||
)
|
||||
|
||||
type Sink interface {
|
||||
|
||||
@@ -110,28 +110,25 @@ func (c *apiKeyCache) SubscribeAuthCacheInvalidation(ctx context.Context, handle
|
||||
return fmt.Errorf("subscribe to auth cache invalidation: %w", err)
|
||||
}
|
||||
|
||||
go func() {
|
||||
defer func() {
|
||||
if err := pubsub.Close(); err != nil {
|
||||
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
ch := pubsub.Channel()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case msg, ok := <-ch:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if msg != nil {
|
||||
handler(msg.Payload)
|
||||
}
|
||||
}
|
||||
defer func() {
|
||||
if err := pubsub.Close(); err != nil {
|
||||
log.Printf("Warning: failed to close auth cache invalidation pubsub: %v", err)
|
||||
}
|
||||
}()
|
||||
service.NotifyAuthCacheSubscriptionReady(ctx)
|
||||
|
||||
return nil
|
||||
ch := pubsub.Channel()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case msg, ok := <-ch:
|
||||
if !ok {
|
||||
return errors.New("auth cache invalidation pubsub channel closed")
|
||||
}
|
||||
if msg != nil {
|
||||
handler(msg.Payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAPIKeyCacheSubscriber_BlocksUntilContextCancellation(t *testing.T) {
|
||||
server := miniredis.RunT(t)
|
||||
client := redis.NewClient(&redis.Options{Addr: server.Addr()})
|
||||
defer func() { _ = client.Close() }()
|
||||
cache := NewAPIKeyCache(client)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
received := make(chan string, 1)
|
||||
returned := make(chan error, 1)
|
||||
go func() {
|
||||
returned <- cache.SubscribeAuthCacheInvalidation(ctx, func(value string) { received <- value })
|
||||
}()
|
||||
|
||||
var value string
|
||||
require.Eventually(t, func() bool {
|
||||
require.NoError(t, client.Publish(context.Background(), authCacheInvalidateChannel, "hash").Err())
|
||||
select {
|
||||
case value = <-received:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}, time.Second, 10*time.Millisecond)
|
||||
require.Equal(t, "hash", value)
|
||||
select {
|
||||
case err := <-returned:
|
||||
t.Fatalf("subscriber returned while connection was active: %v", err)
|
||||
default:
|
||||
}
|
||||
cancel()
|
||||
select {
|
||||
case err := <-returned:
|
||||
require.True(t, errors.Is(err, context.Canceled))
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("subscriber did not stop after context cancellation")
|
||||
}
|
||||
}
|
||||
@@ -326,16 +326,14 @@ func (r *apiKeyRepository) Delete(ctx context.Context, id int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteWithAudit 在同一事务内:
|
||||
// 1. 把(明文 key、所有者、key 名称)写入 deleted_api_key_audits;
|
||||
// 2. 软删除该 key(tombstone 覆盖 key 列以释放唯一约束)。
|
||||
//
|
||||
// 保证"被删除的 key 一定能反查到所有者"。事务模式与 group_repo.DeleteCascade 一致。
|
||||
// DeleteWithAudit keeps the legacy method name for rolling-upgrade compatibility.
|
||||
// It atomically tombstones and soft-deletes the key without retaining credential
|
||||
// material. Tombstoning releases the unique key value for safe reuse.
|
||||
func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error {
|
||||
tombstoneKey := fmt.Sprintf("__deleted__%d__%d", id, time.Now().UnixNano())
|
||||
|
||||
if existingTx := dbent.TxFromContext(ctx); existingTx != nil {
|
||||
return r.deleteWithAudit(ctx, existingTx.Client(), id, tombstoneKey)
|
||||
return r.deleteWithTombstone(ctx, existingTx.Client(), id, tombstoneKey)
|
||||
}
|
||||
|
||||
tx, err := r.client.Tx(ctx)
|
||||
@@ -348,7 +346,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
|
||||
exec = tx.Client()
|
||||
}
|
||||
|
||||
if err := r.deleteWithAudit(ctx, exec, id, tombstoneKey); err != nil {
|
||||
if err := r.deleteWithTombstone(ctx, exec, id, tombstoneKey); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -358,17 +356,7 @@ func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *apiKeyRepository) deleteWithAudit(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
|
||||
// 1. 审计:数据源即 api_keys 当前行;WHERE deleted_at IS NULL 保证只对未删除行写一次。
|
||||
if _, err := exec.ExecContext(ctx, `
|
||||
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
|
||||
SELECT key, id, user_id, name, NOW()
|
||||
FROM api_keys
|
||||
WHERE id = $1 AND deleted_at IS NULL`, id); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. 软删除(tombstone 覆盖 key)。
|
||||
func (r *apiKeyRepository) deleteWithTombstone(ctx context.Context, exec *dbent.Client, id int64, tombstoneKey string) error {
|
||||
res, err := exec.ExecContext(ctx, `
|
||||
UPDATE api_keys
|
||||
SET key = $1, deleted_at = NOW(), updated_at = NOW()
|
||||
|
||||
@@ -556,7 +556,7 @@ func TestIncrementQuotaUsed_Concurrent(t *testing.T) {
|
||||
"并发递增后总和应为 %v,实际为 %v", float64(goroutines)*increment, got.QuotaUsed)
|
||||
}
|
||||
|
||||
func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
|
||||
func (s *APIKeyRepoSuite) TestDeleteWithAudit_TombstonesWithoutRetainingCredential() {
|
||||
user := s.mustCreateUser("delwithaudit@test.com")
|
||||
key := &service.APIKey{
|
||||
UserID: user.ID,
|
||||
@@ -571,18 +571,24 @@ func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
|
||||
_, err := s.repo.GetByID(s.ctx, key.ID)
|
||||
s.Require().Error(err)
|
||||
|
||||
rows, qErr := s.client.QueryContext(s.ctx,
|
||||
`SELECT key, key_name, user_id, api_key_id FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
|
||||
s.Require().NoError(qErr)
|
||||
defer rows.Close()
|
||||
s.Require().True(rows.Next(), "expected one audit row")
|
||||
var auditKey, auditName string
|
||||
var auditUserID, auditAPIKeyID int64
|
||||
s.Require().NoError(rows.Scan(&auditKey, &auditName, &auditUserID, &auditAPIKeyID))
|
||||
s.Require().Equal("sk-del-audit-1", auditKey)
|
||||
s.Require().Equal("Audit Me", auditName)
|
||||
s.Require().Equal(user.ID, auditUserID)
|
||||
s.Require().Equal(key.ID, auditAPIKeyID)
|
||||
var tombstone string
|
||||
var deletedAt time.Time
|
||||
rows, err := s.repo.sql.QueryContext(s.ctx, `SELECT key, deleted_at FROM api_keys WHERE id = $1`, key.ID)
|
||||
s.Require().NoError(err)
|
||||
s.Require().True(rows.Next())
|
||||
s.Require().NoError(rows.Scan(&tombstone, &deletedAt))
|
||||
s.Require().NoError(rows.Close())
|
||||
s.Require().NotEqual("sk-del-audit-1", tombstone)
|
||||
s.Require().Contains(tombstone, "__deleted__")
|
||||
|
||||
var auditCount int
|
||||
auditRows, err := s.repo.sql.QueryContext(s.ctx,
|
||||
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
|
||||
s.Require().NoError(err)
|
||||
s.Require().True(auditRows.Next())
|
||||
s.Require().NoError(auditRows.Scan(&auditCount))
|
||||
s.Require().NoError(auditRows.Close())
|
||||
s.Require().Zero(auditCount, "deleted credentials must not be retained")
|
||||
}
|
||||
|
||||
func (s *APIKeyRepoSuite) TestDeleteWithAudit_RepeatIsIdempotent() {
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build integration
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAuthCacheInvalidationTriggers_CoverSecurityMutationsOnly(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
suffix := time.Now().UnixNano()
|
||||
group := mustCreateGroup(t, integrationEntClient, &service.Group{
|
||||
Name: fmt.Sprintf("auth-outbox-group-%d", suffix), RateMultiplier: 1, IsExclusive: true,
|
||||
})
|
||||
user := mustCreateUser(t, integrationEntClient, &service.User{
|
||||
Email: fmt.Sprintf("auth-outbox-%d@example.com", suffix), Concurrency: 5,
|
||||
})
|
||||
groupID := group.ID
|
||||
keyValue := fmt.Sprintf("sk-auth-outbox-%d", suffix)
|
||||
apiKeyRepo := NewAPIKeyRepository(integrationEntClient, integrationDB)
|
||||
key := &service.APIKey{UserID: user.ID, GroupID: &groupID, Key: keyValue, Name: "outbox", Status: service.StatusActive}
|
||||
require.NoError(t, apiKeyRepo.Create(ctx, key))
|
||||
|
||||
sum := sha256.Sum256([]byte(keyValue))
|
||||
cacheKey := hex.EncodeToString(sum[:])
|
||||
clear := func() {
|
||||
_, err := integrationDB.ExecContext(ctx, "DELETE FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
count := func() int {
|
||||
var value int
|
||||
require.NoError(t, integrationDB.QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey).Scan(&value))
|
||||
return value
|
||||
}
|
||||
clear()
|
||||
t.Cleanup(clear)
|
||||
t.Cleanup(func() {
|
||||
// Keep the shared integration database isolated for suites that assert
|
||||
// platform-wide group counts. The final clear cleanup runs after this one
|
||||
// and removes invalidations emitted by these hard deletes.
|
||||
_, err := integrationDB.ExecContext(ctx, "DELETE FROM user_allowed_groups WHERE user_id = $1 OR group_id = $2", user.ID, group.ID)
|
||||
require.NoError(t, err)
|
||||
_, err = integrationDB.ExecContext(ctx, "DELETE FROM api_keys WHERE id = $1", key.ID)
|
||||
require.NoError(t, err)
|
||||
_, err = integrationDB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", user.ID)
|
||||
require.NoError(t, err)
|
||||
_, err = integrationDB.ExecContext(ctx, "DELETE FROM groups WHERE id = $1", group.ID)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
_, err := integrationDB.ExecContext(ctx, `
|
||||
UPDATE api_keys
|
||||
SET quota_used = quota_used + 1,
|
||||
usage_5h = usage_5h + 1,
|
||||
last_used_at = NOW()
|
||||
WHERE id = $1`, key.ID)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, count(), "usage-only key updates must not enqueue")
|
||||
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'disabled' WHERE id = $1", key.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, count(), "key disable must enqueue")
|
||||
clear()
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE api_keys SET status = 'active' WHERE id = $1", key.ID)
|
||||
require.NoError(t, err)
|
||||
clear()
|
||||
|
||||
userRepo := NewUserRepository(integrationEntClient, integrationDB)
|
||||
loadedUser, err := userRepo.GetByID(ctx, user.ID)
|
||||
require.NoError(t, err)
|
||||
loadedUser.Balance += 10
|
||||
require.NoError(t, userRepo.Update(ctx, loadedUser))
|
||||
require.Zero(t, count(), "balance update with unchanged allowed groups must not enqueue")
|
||||
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'disabled' WHERE id = $1", user.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, count(), "user disable must enqueue all active keys")
|
||||
clear()
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE users SET status = 'active' WHERE id = $1", user.ID)
|
||||
require.NoError(t, err)
|
||||
clear()
|
||||
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET name = name || '-cosmetic' WHERE id = $1", group.ID)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, count(), "cosmetic group update must not enqueue")
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'disabled' WHERE id = $1", group.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, count(), "group disable must enqueue bound keys")
|
||||
clear()
|
||||
_, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'active' WHERE id = $1", group.ID)
|
||||
require.NoError(t, err)
|
||||
clear()
|
||||
|
||||
_, err = integrationDB.ExecContext(ctx,
|
||||
"INSERT INTO user_allowed_groups (user_id, group_id) VALUES ($1, $2)", user.ID, group.ID)
|
||||
require.NoError(t, err)
|
||||
clear()
|
||||
_, err = integrationDB.ExecContext(ctx,
|
||||
"DELETE FROM user_allowed_groups WHERE user_id = $1 AND group_id = $2", user.ID, group.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, count(), "exclusive-group revocation must enqueue")
|
||||
clear()
|
||||
|
||||
require.NoError(t, apiKeyRepo.DeleteWithAudit(ctx, key.ID))
|
||||
require.Equal(t, 1, count(), "tombstone delete must hash OLD.key exactly once")
|
||||
var stored string
|
||||
require.NoError(t, integrationDB.QueryRowContext(ctx,
|
||||
"SELECT cache_key FROM auth_cache_invalidation_outbox WHERE cache_key = $1 LIMIT 1", cacheKey).Scan(&stored))
|
||||
require.Equal(t, cacheKey, stored)
|
||||
require.NotContains(t, stored, keyValue)
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
)
|
||||
|
||||
type authCacheInvalidationOutboxRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewAuthCacheInvalidationOutboxRepository(db *sql.DB) service.AuthCacheInvalidationOutboxRepository {
|
||||
return &authCacheInvalidationOutboxRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *authCacheInvalidationOutboxRepository) Claim(ctx context.Context, workerID string, limit int, lease time.Duration) ([]service.AuthCacheInvalidationEvent, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("nil auth cache invalidation outbox database")
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
leaseSeconds := int64(lease / time.Second)
|
||||
if leaseSeconds < 1 {
|
||||
leaseSeconds = 30
|
||||
}
|
||||
rows, err := r.db.QueryContext(ctx, `
|
||||
WITH candidates AS (
|
||||
SELECT id
|
||||
FROM auth_cache_invalidation_outbox
|
||||
WHERE available_at <= NOW()
|
||||
AND (claimed_at IS NULL OR claimed_at < NOW() - ($3 * INTERVAL '1 second'))
|
||||
ORDER BY id ASC
|
||||
LIMIT $2
|
||||
FOR UPDATE SKIP LOCKED
|
||||
)
|
||||
UPDATE auth_cache_invalidation_outbox AS o
|
||||
SET claimed_at = NOW(), claimed_by = $1
|
||||
FROM candidates AS c
|
||||
WHERE o.id = c.id
|
||||
RETURNING o.id, o.cache_key, o.attempts, o.delivery_stage, o.created_at
|
||||
`, workerID, limit, leaseSeconds)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
events := make([]service.AuthCacheInvalidationEvent, 0, limit)
|
||||
for rows.Next() {
|
||||
var event service.AuthCacheInvalidationEvent
|
||||
if err := rows.Scan(&event.ID, &event.CacheKey, &event.Attempts, &event.Stage, &event.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
event.CacheKey = strings.TrimSpace(event.CacheKey)
|
||||
events = append(events, event)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return events, nil
|
||||
}
|
||||
|
||||
func (r *authCacheInvalidationOutboxRepository) ScheduleSecondPass(ctx context.Context, id int64, workerID string, availableAt time.Time) error {
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE auth_cache_invalidation_outbox
|
||||
SET delivery_stage = 1,
|
||||
available_at = $3,
|
||||
last_error = NULL,
|
||||
claimed_at = NULL,
|
||||
claimed_by = NULL
|
||||
WHERE id = $1 AND claimed_by = $2 AND delivery_stage = 0
|
||||
`, id, workerID, availableAt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
affected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if affected != 1 {
|
||||
return fmt.Errorf("auth cache invalidation claim %d cannot schedule second pass", id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *authCacheInvalidationOutboxRepository) DeleteClaimed(ctx context.Context, id int64, workerID string) error {
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
DELETE FROM auth_cache_invalidation_outbox
|
||||
WHERE id = $1 AND claimed_by = $2
|
||||
`, id, workerID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
affected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if affected != 1 {
|
||||
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *authCacheInvalidationOutboxRepository) RetryClaimed(ctx context.Context, id int64, workerID string, availableAt time.Time, lastError string) error {
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE auth_cache_invalidation_outbox
|
||||
SET attempts = attempts + 1,
|
||||
available_at = $3,
|
||||
last_error = $4,
|
||||
claimed_at = NULL,
|
||||
claimed_by = NULL
|
||||
WHERE id = $1 AND claimed_by = $2
|
||||
`, id, workerID, availableAt, lastError)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
affected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if affected != 1 {
|
||||
return fmt.Errorf("auth cache invalidation claim %d is no longer owned by %s", id, workerID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *authCacheInvalidationOutboxRepository) Stats(ctx context.Context) (service.AuthCacheInvalidationOutboxStats, error) {
|
||||
var (
|
||||
stats service.AuthCacheInvalidationOutboxStats
|
||||
oldest sql.NullTime
|
||||
lastError sql.NullString
|
||||
)
|
||||
err := r.db.QueryRowContext(ctx, `
|
||||
SELECT COUNT(*), MIN(created_at), COALESCE(MAX(attempts), 0),
|
||||
(SELECT last_error
|
||||
FROM auth_cache_invalidation_outbox
|
||||
WHERE last_error IS NOT NULL
|
||||
ORDER BY available_at DESC, id DESC
|
||||
LIMIT 1)
|
||||
FROM auth_cache_invalidation_outbox
|
||||
`).Scan(&stats.Pending, &oldest, &stats.MaxAttempts, &lastError)
|
||||
if err != nil {
|
||||
return stats, err
|
||||
}
|
||||
if oldest.Valid {
|
||||
value := oldest.Time
|
||||
stats.OldestCreatedAt = &value
|
||||
}
|
||||
if lastError.Valid {
|
||||
stats.LastError = lastError.String
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
sqlmock "github.com/DATA-DOG/go-sqlmock"
|
||||
"github.com/Wei-Shaw/sub2api/migrations"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAuthCacheInvalidationOutboxRepository_ClaimUsesLeaseAndSkipLocked(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
created := time.Now().UTC()
|
||||
mock.ExpectQuery("(?s)claimed_at < NOW\\(\\) - .*FOR UPDATE SKIP LOCKED.*RETURNING").
|
||||
WithArgs("worker-a", 100, int64(30)).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}).
|
||||
AddRow(int64(4), strings.Repeat("a", 64), 2, 1, created))
|
||||
|
||||
repo := NewAuthCacheInvalidationOutboxRepository(db)
|
||||
events, err := repo.Claim(context.Background(), "worker-a", 100, 30*time.Second)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1)
|
||||
require.Equal(t, int64(4), events[0].ID)
|
||||
require.Equal(t, 1, events[0].Stage)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestAuthCacheInvalidationOutboxRepository_ClaimIsBoundedByDefault(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
mock.ExpectQuery("(?s)FROM auth_cache_invalidation_outbox.*LIMIT \\$2.*SKIP LOCKED").
|
||||
WithArgs("worker", 100, int64(30)).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "cache_key", "attempts", "delivery_stage", "created_at"}))
|
||||
repo := NewAuthCacheInvalidationOutboxRepository(db)
|
||||
_, err = repo.Claim(context.Background(), "worker", 0, 0)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestAuthCacheInvalidationOutboxRepository_ClaimOwnershipTransitions(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
repo := NewAuthCacheInvalidationOutboxRepository(db)
|
||||
|
||||
next := time.Now().UTC().Add(time.Minute)
|
||||
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
|
||||
WithArgs(int64(1), "worker", next).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
require.NoError(t, repo.ScheduleSecondPass(context.Background(), 1, "worker", next))
|
||||
|
||||
retryAt := next.Add(time.Minute)
|
||||
mock.ExpectExec("UPDATE auth_cache_invalidation_outbox").
|
||||
WithArgs(int64(2), "worker", retryAt, "publish failed").
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
require.NoError(t, repo.RetryClaimed(context.Background(), 2, "worker", retryAt, "publish failed"))
|
||||
|
||||
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
|
||||
WithArgs(int64(3), "worker").
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
require.NoError(t, repo.DeleteClaimed(context.Background(), 3, "worker"))
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestAuthCacheInvalidationOutboxRepository_RejectsLostClaim(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
mock.ExpectExec("DELETE FROM auth_cache_invalidation_outbox").
|
||||
WithArgs(int64(3), "old-worker").
|
||||
WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
repo := NewAuthCacheInvalidationOutboxRepository(db)
|
||||
err = repo.DeleteClaimed(context.Background(), 3, "old-worker")
|
||||
require.ErrorContains(t, err, "no longer owned")
|
||||
}
|
||||
|
||||
func TestAuthCacheInvalidationOutboxRepository_StatsExposeDurableLagAndFailures(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
oldest := time.Now().UTC().Add(-time.Minute)
|
||||
mock.ExpectQuery("(?s)SELECT COUNT\\(\\*\\), MIN\\(created_at\\), COALESCE\\(MAX\\(attempts\\), 0\\)").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"count", "min", "max", "last_error"}).AddRow(5, oldest, 7, "redis down"))
|
||||
repo := NewAuthCacheInvalidationOutboxRepository(db)
|
||||
stats, err := repo.Stats(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(5), stats.Pending)
|
||||
require.Equal(t, 7, stats.MaxAttempts)
|
||||
require.Equal(t, "redis down", stats.LastError)
|
||||
require.NotNil(t, stats.OldestCreatedAt)
|
||||
}
|
||||
|
||||
func TestAuthCacheInvalidationMigration_SecurityCoverageAndNoPlaintextPayload(t *testing.T) {
|
||||
content, err := migrations.FS.ReadFile("184_auth_cache_invalidation_outbox.sql")
|
||||
require.NoError(t, err)
|
||||
sqlText := string(content)
|
||||
for _, required := range []string{
|
||||
"encode(sha256(convert_to(raw_key, 'UTF8')), 'hex')",
|
||||
"OLD.key", "OLD.status", "OLD.deleted_at", "OLD.user_id", "OLD.group_id",
|
||||
"OLD.ip_whitelist", "OLD.ip_blacklist", "OLD.expires_at",
|
||||
"trg_users_auth_cache_invalidation", "trg_groups_auth_cache_invalidation",
|
||||
"trg_user_allowed_groups_auth_cache_invalidation", "FOR EACH ROW",
|
||||
"delivery_stage", "claimed_at", "available_at",
|
||||
} {
|
||||
require.Contains(t, sqlText, required)
|
||||
}
|
||||
require.NotContains(t, sqlText, "quota_used IS DISTINCT")
|
||||
require.NotContains(t, sqlText, "last_used_at IS DISTINCT")
|
||||
|
||||
plaintext := "sk-plaintext-must-not-be-stored"
|
||||
sum := sha256.Sum256([]byte(plaintext))
|
||||
require.Len(t, hex.EncodeToString(sum[:]), 64)
|
||||
require.NotContains(t, sqlText, plaintext)
|
||||
}
|
||||
@@ -128,23 +128,31 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
|
||||
// 获取分布式锁,确保多实例部署时只有一个实例执行迁移。
|
||||
// 这是 PostgreSQL 特有的 Advisory Lock 机制。
|
||||
if err := pgAdvisoryLock(ctx, db); err != nil {
|
||||
lockConn, err := db.Conn(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("acquire migrations lock connection: %w", err)
|
||||
}
|
||||
defer func() { _ = lockConn.Close() }()
|
||||
if err := pgAdvisoryLock(ctx, lockConn); err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
// 无论迁移是否成功,都要释放锁。
|
||||
// 使用 context.Background() 确保即使原 ctx 已取消也能释放锁。
|
||||
_ = pgAdvisoryUnlock(context.Background(), db)
|
||||
// 独立超时确保原 ctx 取消后仍会尝试释放,但数据库链路异常不会
|
||||
// 无限阻塞进程退出。
|
||||
unlockCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
_ = pgAdvisoryUnlock(unlockCtx, lockConn)
|
||||
}()
|
||||
|
||||
// 创建迁移记录表(如果不存在)。
|
||||
// 该表记录所有已应用的迁移及其校验和。
|
||||
if _, err := db.ExecContext(ctx, schemaMigrationsTableDDL); err != nil {
|
||||
if _, err := lockConn.ExecContext(ctx, schemaMigrationsTableDDL); err != nil {
|
||||
return fmt.Errorf("create schema_migrations: %w", err)
|
||||
}
|
||||
|
||||
// 自动对齐 Atlas 基线(如果检测到 legacy schema_migrations 且缺失 atlas_schema_revisions)。
|
||||
if err := ensureAtlasBaselineAligned(ctx, db, fsys); err != nil {
|
||||
if err := ensureAtlasBaselineAligned(ctx, lockConn, fsys); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -175,7 +183,7 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
|
||||
// 检查该迁移是否已经应用
|
||||
var existing string
|
||||
rowErr := db.QueryRowContext(ctx, "SELECT checksum FROM schema_migrations WHERE filename = $1", name).Scan(&existing)
|
||||
rowErr := lockConn.QueryRowContext(ctx, "SELECT checksum FROM schema_migrations WHERE filename = $1", name).Scan(&existing)
|
||||
if rowErr == nil {
|
||||
// 迁移已应用,验证校验和是否匹配
|
||||
if existing != checksum {
|
||||
@@ -207,7 +215,7 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
}
|
||||
|
||||
if nonTx {
|
||||
if err := prepareNonTransactionalMigration(ctx, db, name); err != nil {
|
||||
if err := prepareNonTransactionalMigration(ctx, lockConn, name); err != nil {
|
||||
return fmt.Errorf("prepare migration %s: %w", name, err)
|
||||
}
|
||||
|
||||
@@ -222,18 +230,18 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
if stripSQLLineComment(trimmed) == "" {
|
||||
continue
|
||||
}
|
||||
if _, err := db.ExecContext(ctx, trimmed); err != nil {
|
||||
if _, err := lockConn.ExecContext(ctx, trimmed); err != nil {
|
||||
return fmt.Errorf("apply migration %s (non-tx statement %d): %w", name, i+1, err)
|
||||
}
|
||||
}
|
||||
if _, err := db.ExecContext(ctx, "INSERT INTO schema_migrations (filename, checksum) VALUES ($1, $2)", name, checksum); err != nil {
|
||||
if _, err := lockConn.ExecContext(ctx, "INSERT INTO schema_migrations (filename, checksum) VALUES ($1, $2)", name, checksum); err != nil {
|
||||
return fmt.Errorf("record migration %s (non-tx): %w", name, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// 默认迁移在事务中执行,确保原子性:要么完全成功,要么完全回滚。
|
||||
tx, err := db.BeginTx(ctx, nil)
|
||||
tx, err := lockConn.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("begin migration %s: %w", name, err)
|
||||
}
|
||||
@@ -260,7 +268,14 @@ func applyMigrationsFS(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name string) error {
|
||||
type migrationConnection interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error)
|
||||
}
|
||||
|
||||
func prepareNonTransactionalMigration(ctx context.Context, db migrationConnection, name string) error {
|
||||
switch name {
|
||||
case paymentOrdersOutTradeNoUniqueMigration:
|
||||
return preparePaymentOrdersOutTradeNoUniqueMigration(ctx, db)
|
||||
@@ -273,7 +288,7 @@ func prepareNonTransactionalMigration(ctx context.Context, db *sql.DB, name stri
|
||||
}
|
||||
}
|
||||
|
||||
func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db *sql.DB) error {
|
||||
func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db migrationConnection) error {
|
||||
duplicates, err := findDuplicatePaymentOrderOutTradeNos(ctx, db)
|
||||
if err != nil {
|
||||
return fmt.Errorf("precheck duplicate out_trade_no: %w", err)
|
||||
@@ -289,7 +304,7 @@ func preparePaymentOrdersOutTradeNoUniqueMigration(ctx context.Context, db *sql.
|
||||
return dropInvalidIndexIfPresent(ctx, db, paymentOrdersOutTradeNoUniqueIndex)
|
||||
}
|
||||
|
||||
func dropInvalidIndexIfPresent(ctx context.Context, db *sql.DB, indexName string) error {
|
||||
func dropInvalidIndexIfPresent(ctx context.Context, db migrationConnection, indexName string) error {
|
||||
invalid, err := indexIsInvalid(ctx, db, indexName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check invalid index %s: %w", indexName, err)
|
||||
@@ -304,7 +319,7 @@ func dropInvalidIndexIfPresent(ctx context.Context, db *sql.DB, indexName string
|
||||
return nil
|
||||
}
|
||||
|
||||
func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db *sql.DB) ([]string, error) {
|
||||
func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db migrationConnection) ([]string, error) {
|
||||
rows, err := db.QueryContext(ctx, `
|
||||
SELECT out_trade_no, COUNT(*) AS duplicate_count
|
||||
FROM payment_orders
|
||||
@@ -336,7 +351,7 @@ func findDuplicatePaymentOrderOutTradeNos(ctx context.Context, db *sql.DB) ([]st
|
||||
return duplicates, nil
|
||||
}
|
||||
|
||||
func indexIsInvalid(ctx context.Context, db *sql.DB, indexName string) (bool, error) {
|
||||
func indexIsInvalid(ctx context.Context, db migrationConnection, indexName string) (bool, error) {
|
||||
var invalid bool
|
||||
err := db.QueryRowContext(ctx, `
|
||||
SELECT EXISTS (
|
||||
@@ -352,7 +367,7 @@ func indexIsInvalid(ctx context.Context, db *sql.DB, indexName string) (bool, er
|
||||
return invalid, err
|
||||
}
|
||||
|
||||
func ensureAtlasBaselineAligned(ctx context.Context, db *sql.DB, fsys fs.FS) error {
|
||||
func ensureAtlasBaselineAligned(ctx context.Context, db migrationConnection, fsys fs.FS) error {
|
||||
hasLegacy, err := tableExists(ctx, db, "schema_migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("check schema_migrations: %w", err)
|
||||
@@ -393,7 +408,7 @@ func ensureAtlasBaselineAligned(ctx context.Context, db *sql.DB, fsys fs.FS) err
|
||||
return nil
|
||||
}
|
||||
|
||||
func tableExists(ctx context.Context, db *sql.DB, tableName string) (bool, error) {
|
||||
func tableExists(ctx context.Context, db migrationConnection, tableName string) (bool, error) {
|
||||
var exists bool
|
||||
err := db.QueryRowContext(ctx, `
|
||||
SELECT EXISTS (
|
||||
@@ -524,7 +539,12 @@ func stripSQLLineComment(s string) string {
|
||||
// pgAdvisoryLock 获取 PostgreSQL Advisory Lock。
|
||||
// Advisory Lock 是一种轻量级的锁机制,不与任何特定的数据库对象关联。
|
||||
// 它非常适合用于应用层面的分布式锁场景,如迁移序列化。
|
||||
func pgAdvisoryLock(ctx context.Context, db *sql.DB) error {
|
||||
type advisoryLockConnection interface {
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
}
|
||||
|
||||
func pgAdvisoryLock(ctx context.Context, db advisoryLockConnection) error {
|
||||
ticker := time.NewTicker(migrationsLockRetryInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -546,7 +566,7 @@ func pgAdvisoryLock(ctx context.Context, db *sql.DB) error {
|
||||
|
||||
// pgAdvisoryUnlock 释放 PostgreSQL Advisory Lock。
|
||||
// 必须在获取锁后确保释放,否则会阻塞其他实例的迁移操作。
|
||||
func pgAdvisoryUnlock(ctx context.Context, db *sql.DB) error {
|
||||
func pgAdvisoryUnlock(ctx context.Context, db advisoryLockConnection) error {
|
||||
_, err := db.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", migrationsAdvisoryLockID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("release migrations lock: %w", err)
|
||||
|
||||
@@ -275,6 +275,9 @@ func TestApplyMigrationsFS_TransactionalMigration(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = db.Close() }()
|
||||
// The advisory lock and all migration work must share one session. This also
|
||||
// proves startup cannot self-deadlock when deployments cap the pool at one.
|
||||
db.SetMaxOpenConns(1)
|
||||
|
||||
prepareMigrationsBootstrapExpectations(mock)
|
||||
mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1").
|
||||
|
||||
@@ -5,11 +5,32 @@ package repository
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestMigrationsRunner_ConcurrentInstancesSerializeOnSessionLock(t *testing.T) {
|
||||
const instances = 2
|
||||
errorsByInstance := make([]error, instances)
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < instances; i++ {
|
||||
wg.Add(1)
|
||||
go func(index int) {
|
||||
defer wg.Done()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
errorsByInstance[index] = ApplyMigrations(ctx, integrationDB)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
for i, err := range errorsByInstance {
|
||||
require.NoErrorf(t, err, "migration instance %d", i)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) {
|
||||
tx := testTx(t)
|
||||
|
||||
@@ -111,6 +132,13 @@ func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) {
|
||||
requireColumn(t, tx, "ops_system_logs", "api_key_id", "bigint", 0, true)
|
||||
requireIndex(t, tx, "ops_system_logs", "idx_ops_system_logs_api_key_id_created_at")
|
||||
|
||||
// Bounded ingress rejection security aggregates.
|
||||
requireColumn(t, tx, "ops_ingress_reject_aggregates", "bucket_start", "timestamp with time zone", 0, false)
|
||||
requireColumn(t, tx, "ops_ingress_reject_aggregates", "client_ip", "inet", 0, false)
|
||||
requireColumn(t, tx, "ops_ingress_reject_aggregates", "request_count", "bigint", 0, false)
|
||||
requireIndex(t, tx, "ops_ingress_reject_aggregates", "idx_ops_ingress_reject_aggregates_bucket")
|
||||
requireIndex(t, tx, "ops_ingress_reject_aggregates", "idx_ops_ingress_reject_aggregates_ip_bucket")
|
||||
|
||||
// user_allowed_groups table should exist
|
||||
var uagRegclass sql.NullString
|
||||
require.NoError(t, tx.QueryRowContext(context.Background(), "SELECT to_regclass('public.user_allowed_groups')").Scan(&uagRegclass))
|
||||
|
||||
@@ -137,26 +137,17 @@ func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
|
||||
func TestBuildOpsErrorLogsWhere_UserOwnershipIsDirectOnly(t *testing.T) {
|
||||
uid := int64(42)
|
||||
|
||||
// 开关开启 → 归属放宽为 OR(user_id 或 deleted_key_owner_user_id),且共用同一占位符
|
||||
on := &service.OpsErrorLogFilter{UserID: &uid, MatchDeletedKeyOwner: true}
|
||||
whereOn, argsOn := buildOpsErrorLogsWhere(on)
|
||||
if !strings.Contains(whereOn, "(e.user_id = $1 OR e.deleted_key_owner_user_id = $1)") {
|
||||
t.Fatalf("MatchDeletedKeyOwner=true should widen to OR, got: %s", whereOn)
|
||||
filter := &service.OpsErrorLogFilter{UserID: &uid}
|
||||
where, args := buildOpsErrorLogsWhere(filter)
|
||||
if !strings.Contains(where, "e.user_id = $1") {
|
||||
t.Fatalf("user scope should match user_id exactly, got: %s", where)
|
||||
}
|
||||
if len(argsOn) != 1 || argsOn[0] != uid {
|
||||
t.Fatalf("expected single reused arg %d, got %v", uid, argsOn)
|
||||
if len(args) != 1 || args[0] != uid {
|
||||
t.Fatalf("expected user id arg %d, got %v", uid, args)
|
||||
}
|
||||
|
||||
// 开关关闭(默认)→ 仅精确 user_id,绝不出现 deleted_key_owner_user_id(admin 回归)
|
||||
off := &service.OpsErrorLogFilter{UserID: &uid}
|
||||
whereOff, _ := buildOpsErrorLogsWhere(off)
|
||||
if !strings.Contains(whereOff, "e.user_id = $1") {
|
||||
t.Fatalf("default should match user_id exactly, got: %s", whereOff)
|
||||
}
|
||||
if strings.Contains(whereOff, "deleted_key_owner_user_id") {
|
||||
t.Fatalf("default must NOT include deleted_key_owner_user_id, got: %s", whereOff)
|
||||
if strings.Contains(where, "deleted_key_owner_user_id") {
|
||||
t.Fatalf("user ownership must not depend on deleted-key attribution: %s", where)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
)
|
||||
|
||||
const ingressRejectUpsertChunkSize = 500
|
||||
|
||||
func (r *opsRepository) BatchUpsertIngressRejects(ctx context.Context, items []*service.OpsIngressRejectAggregate) error {
|
||||
if r == nil || r.db == nil || len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
|
||||
for start := 0; start < len(items); start += ingressRejectUpsertChunkSize {
|
||||
end := start + ingressRejectUpsertChunkSize
|
||||
if end > len(items) {
|
||||
end = len(items)
|
||||
}
|
||||
valid := make([]*service.OpsIngressRejectAggregate, 0, end-start)
|
||||
for _, item := range items[start:end] {
|
||||
if item != nil && item.RequestCount > 0 {
|
||||
valid = append(valid, item)
|
||||
}
|
||||
}
|
||||
if len(valid) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
var query strings.Builder
|
||||
_, _ = query.WriteString(`INSERT INTO ops_ingress_reject_aggregates
|
||||
(bucket_start, reject_reason, route_family, protocol, client_ip, user_id, api_key_id, request_count, first_seen, last_seen)
|
||||
VALUES `)
|
||||
args := make([]any, 0, len(valid)*10)
|
||||
for i, item := range valid {
|
||||
if i > 0 {
|
||||
_ = query.WriteByte(',')
|
||||
}
|
||||
base := len(args)
|
||||
fmt.Fprintf(&query, "($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d)",
|
||||
base+1, base+2, base+3, base+4, base+5, base+6, base+7, base+8, base+9, base+10)
|
||||
var userID, apiKeyID int64
|
||||
if item.UserID != nil {
|
||||
userID = *item.UserID
|
||||
}
|
||||
if item.APIKeyID != nil {
|
||||
apiKeyID = *item.APIKeyID
|
||||
}
|
||||
args = append(args, item.BucketStart.UTC(), item.RejectReason, item.RouteFamily, item.Protocol,
|
||||
item.ClientIP, userID, apiKeyID, item.RequestCount, item.FirstSeen.UTC(), item.LastSeen.UTC())
|
||||
}
|
||||
_, _ = query.WriteString(`
|
||||
ON CONFLICT (bucket_start, reject_reason, route_family, protocol, client_ip, user_id, api_key_id)
|
||||
DO UPDATE SET request_count = ops_ingress_reject_aggregates.request_count + EXCLUDED.request_count,
|
||||
first_seen = LEAST(ops_ingress_reject_aggregates.first_seen, EXCLUDED.first_seen),
|
||||
last_seen = GREATEST(ops_ingress_reject_aggregates.last_seen, EXCLUDED.last_seen),
|
||||
updated_at = NOW()`)
|
||||
if _, err := tx.ExecContext(ctx, query.String(), args...); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (r *opsRepository) ListIngressRejects(ctx context.Context, filter *service.OpsIngressRejectFilter) (*service.OpsIngressRejectList, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, fmt.Errorf("nil ops repository")
|
||||
}
|
||||
if filter == nil {
|
||||
filter = &service.OpsIngressRejectFilter{}
|
||||
}
|
||||
page, pageSize := filter.Page, filter.PageSize
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize <= 0 {
|
||||
pageSize = 50
|
||||
}
|
||||
if pageSize > 200 {
|
||||
pageSize = 200
|
||||
}
|
||||
|
||||
clauses := []string{"1=1"}
|
||||
args := make([]any, 0)
|
||||
add := func(expr string, value any) {
|
||||
args = append(args, value)
|
||||
clauses = append(clauses, fmt.Sprintf(expr, len(args)))
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
add("bucket_start >= $%d", filter.StartTime.UTC())
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
add("bucket_start < $%d", filter.EndTime.UTC())
|
||||
}
|
||||
if value := strings.TrimSpace(filter.RejectReason); value != "" {
|
||||
add("reject_reason = $%d", value)
|
||||
}
|
||||
if value := strings.TrimSpace(filter.RouteFamily); value != "" {
|
||||
add("route_family = $%d", value)
|
||||
}
|
||||
if value := strings.TrimSpace(filter.Protocol); value != "" {
|
||||
add("protocol = $%d", value)
|
||||
}
|
||||
if value := strings.TrimSpace(filter.ClientIP); value != "" {
|
||||
add("client_ip = $%d::inet", value)
|
||||
}
|
||||
if filter.UserID != nil {
|
||||
add("user_id = $%d", *filter.UserID)
|
||||
}
|
||||
if filter.APIKeyID != nil {
|
||||
add("api_key_id = $%d", *filter.APIKeyID)
|
||||
}
|
||||
where := "WHERE " + strings.Join(clauses, " AND ")
|
||||
|
||||
var total int
|
||||
if err := r.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM ops_ingress_reject_aggregates "+where, args...).Scan(&total); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
args = append(args, pageSize, (page-1)*pageSize)
|
||||
query := fmt.Sprintf(`SELECT id,bucket_start,reject_reason,route_family,protocol,host(client_ip),user_id,api_key_id,request_count,first_seen,last_seen
|
||||
FROM ops_ingress_reject_aggregates %s ORDER BY bucket_start DESC,id DESC LIMIT $%d OFFSET $%d`, where, len(args)-1, len(args))
|
||||
rows, err := r.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
result := &service.OpsIngressRejectList{
|
||||
Items: make([]*service.OpsIngressRejectAggregate, 0, pageSize), Total: total, Page: page, PageSize: pageSize,
|
||||
}
|
||||
for rows.Next() {
|
||||
item := &service.OpsIngressRejectAggregate{}
|
||||
var userID, apiKeyID int64
|
||||
if err := rows.Scan(&item.ID, &item.BucketStart, &item.RejectReason, &item.RouteFamily, &item.Protocol,
|
||||
&item.ClientIP, &userID, &apiKeyID, &item.RequestCount, &item.FirstSeen, &item.LastSeen); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userID > 0 {
|
||||
item.UserID = &userID
|
||||
}
|
||||
if apiKeyID > 0 {
|
||||
item.APIKeyID = &apiKeyID
|
||||
}
|
||||
result.Items = append(result.Items, item)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBatchUpsertIngressRejectsUsesFixedMultiRowChunks(t *testing.T) {
|
||||
db, mock, err := sqlmock.New()
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
repo := &opsRepository{db: db}
|
||||
now := time.Now().UTC().Truncate(time.Minute)
|
||||
items := make([]*service.OpsIngressRejectAggregate, ingressRejectUpsertChunkSize+1)
|
||||
for i := range items {
|
||||
items[i] = &service.OpsIngressRejectAggregate{
|
||||
BucketStart: now, RejectReason: "invalid_api_key", RouteFamily: "messages",
|
||||
Protocol: "anthropic", ClientIP: "192.0.2.1", RequestCount: 1, FirstSeen: now, LastSeen: now,
|
||||
}
|
||||
}
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, int64(ingressRejectUpsertChunkSize)))
|
||||
mock.ExpectExec("INSERT INTO ops_ingress_reject_aggregates").WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
mock.ExpectCommit()
|
||||
require.NoError(t, repo.BatchUpsertIngressRejects(context.Background(), items))
|
||||
require.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -56,12 +55,9 @@ INSERT INTO ops_error_logs (
|
||||
response_latency_ms,
|
||||
time_to_first_token_ms,
|
||||
created_at,
|
||||
attempted_key_prefix,
|
||||
deleted_key_owner_user_id,
|
||||
deleted_key_name,
|
||||
api_key_prefix
|
||||
) VALUES (
|
||||
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38,$39,$40,$41
|
||||
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38
|
||||
)`
|
||||
|
||||
func NewOpsRepository(db *sql.DB) service.OpsRepository {
|
||||
@@ -170,9 +166,6 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
|
||||
opsNullInt64(input.ResponseLatencyMs),
|
||||
opsNullInt64(input.TimeToFirstTokenMs),
|
||||
input.CreatedAt,
|
||||
opsNullString(input.AttemptedKeyPrefix),
|
||||
opsNullInt64(input.DeletedKeyOwnerUserID),
|
||||
opsNullString(input.DeletedKeyName),
|
||||
opsNullString(input.APIKeyPrefix),
|
||||
}
|
||||
}
|
||||
@@ -274,16 +267,12 @@ SELECT
|
||||
COALESCE(e.user_agent, ''),
|
||||
e.request_type,
|
||||
COALESCE(ak.name, ''),
|
||||
ak.deleted_at,
|
||||
COALESCE(e.deleted_key_name, ''),
|
||||
e.deleted_key_owner_user_id,
|
||||
COALESCE(du.email, '')
|
||||
ak.deleted_at
|
||||
FROM ops_error_logs e
|
||||
LEFT JOIN accounts a ON e.account_id = a.id
|
||||
LEFT JOIN groups g ON e.group_id = g.id
|
||||
LEFT JOIN users u ON e.user_id = u.id
|
||||
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
|
||||
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
|
||||
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
|
||||
` + where + `
|
||||
ORDER BY ` + opsErrorLogsOrderBy(filter) + `
|
||||
@@ -313,9 +302,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
|
||||
var requestType sql.NullInt64
|
||||
var apiKeyName string
|
||||
var apiKeyDeletedAt sql.NullTime
|
||||
var deletedKeyName string
|
||||
var deletedKeyOwnerID sql.NullInt64
|
||||
var deletedKeyOwnerEmail string
|
||||
if err := rows.Scan(
|
||||
&item.ID,
|
||||
&item.CreatedAt,
|
||||
@@ -352,9 +338,6 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
|
||||
&requestType,
|
||||
&apiKeyName,
|
||||
&apiKeyDeletedAt,
|
||||
&deletedKeyName,
|
||||
&deletedKeyOwnerID,
|
||||
&deletedKeyOwnerEmail,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -395,21 +378,8 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
|
||||
v := int16(requestType.Int64)
|
||||
item.RequestType = &v
|
||||
}
|
||||
// Key 名称:优先关联到的 ak.name(已软删的 key name 仍保留);
|
||||
// 关联不到(api_key_id 为空 / 历史硬删)时回退错误记录里快照的 deleted_key_name。
|
||||
if apiKeyName != "" {
|
||||
item.APIKeyName = apiKeyName
|
||||
} else {
|
||||
item.APIKeyName = deletedKeyName
|
||||
}
|
||||
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
|
||||
item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
|
||||
// 已删除 KEY 所有者快照:认证失败行 user_id 为空,列表用户列以此回退。
|
||||
if deletedKeyOwnerID.Valid {
|
||||
v := deletedKeyOwnerID.Int64
|
||||
item.DeletedKeyOwnerUserID = &v
|
||||
item.DeletedKeyOwnerEmail = deletedKeyOwnerEmail
|
||||
}
|
||||
item.APIKeyName = apiKeyName
|
||||
item.APIKeyDeleted = apiKeyDeletedAt.Valid
|
||||
out = append(out, &item)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
@@ -477,10 +447,6 @@ SELECT
|
||||
e.upstream_latency_ms,
|
||||
e.response_latency_ms,
|
||||
e.time_to_first_token_ms,
|
||||
COALESCE(e.attempted_key_prefix, ''),
|
||||
e.deleted_key_owner_user_id,
|
||||
COALESCE(du.email, ''),
|
||||
COALESCE(e.deleted_key_name, ''),
|
||||
COALESCE(e.api_key_prefix, ''),
|
||||
COALESCE(ak.name, ''),
|
||||
ak.deleted_at
|
||||
@@ -488,7 +454,6 @@ FROM ops_error_logs e
|
||||
LEFT JOIN users u ON e.user_id = u.id
|
||||
LEFT JOIN accounts a ON e.account_id = a.id
|
||||
LEFT JOIN groups g ON e.group_id = g.id
|
||||
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
|
||||
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
|
||||
WHERE e.id = $1
|
||||
LIMIT 1`
|
||||
@@ -509,7 +474,6 @@ LIMIT 1`
|
||||
var responseLatency sql.NullInt64
|
||||
var ttft sql.NullInt64
|
||||
var requestType sql.NullInt64
|
||||
var deletedKeyOwnerUserID sql.NullInt64
|
||||
var detailAPIKeyName string
|
||||
var detailAPIKeyDeletedAt sql.NullTime
|
||||
|
||||
@@ -557,10 +521,6 @@ LIMIT 1`
|
||||
&upstreamLatency,
|
||||
&responseLatency,
|
||||
&ttft,
|
||||
&out.AttemptedKeyPrefix,
|
||||
&deletedKeyOwnerUserID,
|
||||
&out.DeletedKeyOwnerEmail,
|
||||
&out.DeletedKeyName,
|
||||
&out.APIKeyPrefix,
|
||||
&detailAPIKeyName,
|
||||
&detailAPIKeyDeletedAt,
|
||||
@@ -626,18 +586,8 @@ LIMIT 1`
|
||||
v := int16(requestType.Int64)
|
||||
out.RequestType = &v
|
||||
}
|
||||
if deletedKeyOwnerUserID.Valid {
|
||||
v := deletedKeyOwnerUserID.Int64
|
||||
out.DeletedKeyOwnerUserID = &v
|
||||
}
|
||||
// Key 名称:优先关联到的 ak.name;关联不到时回退快照的 deleted_key_name。
|
||||
if detailAPIKeyName != "" {
|
||||
out.APIKeyName = detailAPIKeyName
|
||||
} else {
|
||||
out.APIKeyName = out.DeletedKeyName
|
||||
}
|
||||
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
|
||||
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid || (detailAPIKeyName == "" && out.DeletedKeyName != "")
|
||||
out.APIKeyName = detailAPIKeyName
|
||||
out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid
|
||||
|
||||
// Normalize upstream_errors to empty string when stored as JSON null.
|
||||
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
|
||||
@@ -648,26 +598,6 @@ LIMIT 1`
|
||||
return &out, nil
|
||||
}
|
||||
|
||||
// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计。
|
||||
// 同一 key 可能有多条历史(反复创建/删除),取 deleted_at 最近一条(id 作同毫秒 tiebreaker)。
|
||||
// 未命中返回 (nil, nil)。
|
||||
func (r *opsRepository) LookupDeletedKeyAudit(ctx context.Context, key string) (*service.DeletedKeyAuditResult, error) {
|
||||
var res service.DeletedKeyAuditResult
|
||||
err := r.db.QueryRowContext(ctx, `
|
||||
SELECT user_id, key_name
|
||||
FROM deleted_api_key_audits
|
||||
WHERE key = $1
|
||||
ORDER BY deleted_at DESC, id DESC
|
||||
LIMIT 1`, key).Scan(&res.UserID, &res.KeyName)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &res, nil
|
||||
}
|
||||
|
||||
func (r *opsRepository) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64, resolvedAt *time.Time) error {
|
||||
if r == nil || r.db == nil {
|
||||
return fmt.Errorf("nil ops repository")
|
||||
@@ -1082,12 +1012,7 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
|
||||
if filter.UserID != nil && *filter.UserID > 0 {
|
||||
args = append(args, *filter.UserID)
|
||||
n := itoa(len(args))
|
||||
if filter.MatchDeletedKeyOwner {
|
||||
// 用户侧:把「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录也纳入。
|
||||
clauses = append(clauses, "(e.user_id = $"+n+" OR e.deleted_key_owner_user_id = $"+n+")")
|
||||
} else {
|
||||
clauses = append(clauses, "e.user_id = $"+n)
|
||||
}
|
||||
clauses = append(clauses, "e.user_id = $"+n)
|
||||
}
|
||||
if filter.APIKeyID != nil && *filter.APIKeyID > 0 {
|
||||
args = append(args, *filter.APIKeyID)
|
||||
|
||||
@@ -14,7 +14,7 @@ func TestOpsInsertErrorLogArgsPreservesExplicitZeroUpstreamStatus(t *testing.T)
|
||||
zero := 0
|
||||
args := opsInsertErrorLogArgs(&service.OpsInsertErrorLogInput{UpstreamStatusCode: &zero})
|
||||
|
||||
require.Len(t, args, 41)
|
||||
require.Len(t, args, 38)
|
||||
encoded, ok := args[27].(sql.NullInt64)
|
||||
require.True(t, ok)
|
||||
require.True(t, encoded.Valid)
|
||||
|
||||
@@ -11,47 +11,13 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestGetErrorLogByID_DeletedKeyOwner 验证:
|
||||
// 1. 带 deleted_key_owner_user_id 的记录能正确 JOIN users 返回 DeletedKeyOwnerEmail
|
||||
// 2. 新列全为 NULL 的普通记录 Scan 不报错,这些字段为空/nil
|
||||
func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
|
||||
func TestGetErrorLogByID_APIKeyPrefixAndUpstreamStatus(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE ops_error_logs RESTART IDENTITY CASCADE")
|
||||
|
||||
repo := NewOpsRepository(integrationDB).(*opsRepository)
|
||||
|
||||
// ── Case 1: 带 deleted_key_owner 信息的记录 ──────────────────────────────
|
||||
owner := mustCreateUser(t, integrationEntClient, &service.User{
|
||||
Email: "deleted-key-owner-" + time.Now().Format("150405.000000000") + "@example.com",
|
||||
})
|
||||
|
||||
var insertedID int64
|
||||
err := integrationDB.QueryRowContext(ctx, `
|
||||
INSERT INTO ops_error_logs (
|
||||
error_phase, error_type, severity, status_code, created_at,
|
||||
attempted_key_prefix, deleted_key_owner_user_id, deleted_key_name
|
||||
) VALUES (
|
||||
'auth', 'INVALID_API_KEY', 'error', 401, NOW(),
|
||||
'sk-test-abc', $1, 'my-deleted-key'
|
||||
) RETURNING id`,
|
||||
owner.ID,
|
||||
).Scan(&insertedID)
|
||||
require.NoError(t, err)
|
||||
require.Positive(t, insertedID)
|
||||
|
||||
detail, err := repo.GetErrorLogByID(ctx, insertedID)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, detail)
|
||||
|
||||
require.Equal(t, "sk-test-abc", detail.AttemptedKeyPrefix)
|
||||
require.NotNil(t, detail.DeletedKeyOwnerUserID)
|
||||
require.Equal(t, owner.ID, *detail.DeletedKeyOwnerUserID)
|
||||
require.Equal(t, owner.Email, detail.DeletedKeyOwnerEmail)
|
||||
require.Equal(t, "my-deleted-key", detail.DeletedKeyName)
|
||||
|
||||
// ── Case 2: 新列全为 NULL 的普通错误记录 ──────────────────────────────────
|
||||
var plainID int64
|
||||
err = integrationDB.QueryRowContext(ctx, `
|
||||
err := integrationDB.QueryRowContext(ctx, `
|
||||
INSERT INTO ops_error_logs (
|
||||
error_phase, error_type, severity, status_code, created_at
|
||||
) VALUES (
|
||||
@@ -59,20 +25,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
|
||||
) RETURNING id`,
|
||||
).Scan(&plainID)
|
||||
require.NoError(t, err)
|
||||
require.Positive(t, plainID)
|
||||
|
||||
plain, err := repo.GetErrorLogByID(ctx, plainID)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, plain)
|
||||
require.Empty(t, plain.APIKeyPrefix)
|
||||
|
||||
require.Empty(t, plain.AttemptedKeyPrefix, "no prefix for plain error")
|
||||
require.Nil(t, plain.DeletedKeyOwnerUserID, "no owner for plain error")
|
||||
require.Empty(t, plain.DeletedKeyOwnerEmail, "no owner email for plain error")
|
||||
require.Empty(t, plain.DeletedKeyName, "no key name for plain error")
|
||||
require.Empty(t, plain.APIKeyPrefix, "no api key prefix for plain error")
|
||||
|
||||
// ── Case 3: 有效(未删除)key 报错,经 InsertErrorLog 快照 api_key_prefix ──────
|
||||
// 走真实 InsertErrorLog 写入路径(覆盖新列 + $41 占位符),再 GetErrorLogByID 读回。
|
||||
validID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
|
||||
ErrorPhase: "request",
|
||||
ErrorType: "api_error",
|
||||
@@ -82,17 +39,11 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
|
||||
APIKeyPrefix: "sk-valid",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Positive(t, validID)
|
||||
|
||||
valid, err := repo.GetErrorLogByID(ctx, validID)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, valid)
|
||||
|
||||
require.Equal(t, "sk-valid", valid.APIKeyPrefix)
|
||||
require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive")
|
||||
require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner")
|
||||
|
||||
// ── Case 4: account_auth with no inference attempt preserves explicit 0 ──
|
||||
zero := 0
|
||||
credentialFailureID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
|
||||
ErrorPhase: "account_auth",
|
||||
|
||||
@@ -1,36 +0,0 @@
|
||||
//go:build integration
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestOpsRepositoryLookupDeletedKeyAudit(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
_, _ = integrationDB.ExecContext(ctx, "TRUNCATE deleted_api_key_audits RESTART IDENTITY")
|
||||
repo := NewOpsRepository(integrationDB).(*opsRepository)
|
||||
|
||||
// 同一 key 两条审计,取最近一条(deleted_at DESC, id DESC)
|
||||
_, err := integrationDB.ExecContext(ctx, `
|
||||
INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
|
||||
VALUES ('sk-lookup-1', 10, 100, 'old', $1),
|
||||
('sk-lookup-1', 11, 200, 'new', $2)`,
|
||||
time.Now().Add(-time.Hour), time.Now())
|
||||
require.NoError(t, err)
|
||||
|
||||
res, err := repo.LookupDeletedKeyAudit(ctx, "sk-lookup-1")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, res)
|
||||
require.Equal(t, int64(200), res.UserID)
|
||||
require.Equal(t, "new", res.KeyName)
|
||||
|
||||
// 未命中返回 nil
|
||||
miss, err := repo.LookupDeletedKeyAudit(ctx, "sk-never-existed")
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, miss)
|
||||
}
|
||||
@@ -1039,24 +1039,44 @@ func (r *userRepository) syncUserAllowedGroupsWithClient(ctx context.Context, cl
|
||||
return nil
|
||||
}
|
||||
|
||||
// Keep join table as the source of truth for reads.
|
||||
if _, err := client.UserAllowedGroup.Delete().Where(userallowedgroup.UserIDEQ(userID)).Exec(ctx); err != nil {
|
||||
existingRows, err := client.UserAllowedGroup.Query().
|
||||
Where(userallowedgroup.UserIDEQ(userID)).
|
||||
All(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
unique := make(map[int64]struct{}, len(groupIDs))
|
||||
desired := make(map[int64]struct{}, len(groupIDs))
|
||||
for _, id := range groupIDs {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
unique[id] = struct{}{}
|
||||
desired[id] = struct{}{}
|
||||
}
|
||||
|
||||
if len(unique) > 0 {
|
||||
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(unique))
|
||||
for groupID := range unique {
|
||||
existing := make(map[int64]struct{}, len(existingRows))
|
||||
removed := make([]int64, 0)
|
||||
for _, row := range existingRows {
|
||||
existing[row.GroupID] = struct{}{}
|
||||
if _, keep := desired[row.GroupID]; !keep {
|
||||
removed = append(removed, row.GroupID)
|
||||
}
|
||||
}
|
||||
if len(removed) > 0 {
|
||||
if _, err := client.UserAllowedGroup.Delete().
|
||||
Where(userallowedgroup.UserIDEQ(userID), userallowedgroup.GroupIDIn(removed...)).
|
||||
Exec(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
creates := make([]*dbent.UserAllowedGroupCreate, 0, len(desired))
|
||||
for groupID := range desired {
|
||||
if _, present := existing[groupID]; !present {
|
||||
creates = append(creates, client.UserAllowedGroup.Create().SetUserID(userID).SetGroupID(groupID))
|
||||
}
|
||||
}
|
||||
if len(creates) > 0 {
|
||||
if err := client.UserAllowedGroup.
|
||||
CreateBulk(creates...).
|
||||
OnConflictColumns(userallowedgroup.FieldUserID, userallowedgroup.FieldGroupID).
|
||||
|
||||
@@ -14,7 +14,7 @@ import (
|
||||
)
|
||||
|
||||
// TestUserRepository_DeleteUser_AtomicWithAPIKeys 复现 AdminService.DeleteUser 的事务编排场景:
|
||||
// 把"删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
|
||||
// 把"tombstone 并删 API Key"(apiKeyRepo.DeleteWithAudit) 与"删 User"(userRepo.Delete) 放进同一个外部事务时,
|
||||
// userRepo.Delete 必须复用 context 中的事务,而不是用 base client 自起一个独立事务并提前提交。
|
||||
//
|
||||
// 用例用"回滚外层事务"来模拟 commit 失败 / 中止:
|
||||
@@ -91,5 +91,5 @@ func TestUserRepository_DeleteUser_AtomicWithAPIKeys(t *testing.T) {
|
||||
|
||||
require.NoError(t, integrationDB.QueryRowContext(ctx,
|
||||
`SELECT COUNT(*) FROM deleted_api_key_audits WHERE user_id = $1`, user.ID).Scan(&auditCount))
|
||||
require.Equal(t, 2, auditCount, "提交后应为每个被删 Key 写入一行审计")
|
||||
require.Zero(t, auditCount, "提交后也不得保留被删 Key 的凭据材料")
|
||||
}
|
||||
|
||||
@@ -126,6 +126,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewLeaderLockCache,
|
||||
ProvideSchedulerCache,
|
||||
NewSchedulerOutboxRepository,
|
||||
NewAuthCacheInvalidationOutboxRepository,
|
||||
NewProxyLatencyCache,
|
||||
NewTotpCache,
|
||||
NewRefreshTokenCache,
|
||||
|
||||
@@ -16,6 +16,7 @@ var ProviderSet = wire.NewSet(
|
||||
wire.Bind(new(ConfigStore), new(*ConfigManager)),
|
||||
NewPromptService,
|
||||
wire.Bind(new(PromptEngine), new(*PromptService)),
|
||||
wire.Bind(new(PromptAdminService), new(*PromptService)),
|
||||
NewLegacyModerationAdapter,
|
||||
NewCoordinator,
|
||||
NewPromptAdminHandler,
|
||||
|
||||
@@ -103,8 +103,9 @@ func ProvideRouter(
|
||||
func ProvideHTTPServer(cfg *config.Config, router *gin.Engine) *http.Server {
|
||||
httpHandler := http.Handler(router)
|
||||
server := &http.Server{
|
||||
Addr: cfg.Server.Address(),
|
||||
Handler: httpHandler,
|
||||
Addr: cfg.Server.Address(),
|
||||
Handler: httpHandler,
|
||||
MaxHeaderBytes: cfg.Server.MaxHeaderBytes,
|
||||
// ReadHeaderTimeout: 读取请求头的超时时间,防止慢速请求头攻击
|
||||
ReadHeaderTimeout: time.Duration(cfg.Server.ReadHeaderTimeout) * time.Second,
|
||||
// IdleTimeout: 空闲连接超时时间,释放不活跃的连接资源
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
//go:build unit
|
||||
|
||||
package server
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func ingressTestConfig() *config.Config {
|
||||
return &config.Config{
|
||||
Server: config.ServerConfig{
|
||||
Host: "127.0.0.1",
|
||||
ReadHeaderTimeout: 1,
|
||||
IdleTimeout: 5,
|
||||
MaxHeaderBytes: 8 * 1024,
|
||||
MaxRequestBodySize: 1024,
|
||||
},
|
||||
Gateway: config.GatewayConfig{MaxBodySize: 1024},
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvideHTTPServerAppliesIngressLimits(t *testing.T) {
|
||||
srv := ProvideHTTPServer(ingressTestConfig(), gin.New())
|
||||
require.Equal(t, 8*1024, srv.MaxHeaderBytes)
|
||||
require.Equal(t, time.Second, srv.ReadHeaderTimeout)
|
||||
require.Equal(t, 5*time.Second, srv.IdleTimeout)
|
||||
}
|
||||
|
||||
func TestProvideHTTPServerEnablesBoundedH2C(t *testing.T) {
|
||||
cfg := ingressTestConfig()
|
||||
cfg.Server.H2C = config.H2CConfig{
|
||||
Enabled: true,
|
||||
MaxConcurrentStreams: 25,
|
||||
IdleTimeout: 30,
|
||||
MaxReadFrameSize: 64 * 1024,
|
||||
MaxUploadBufferPerConnection: 1024 * 1024,
|
||||
MaxUploadBufferPerStream: 256 * 1024,
|
||||
}
|
||||
srv := ProvideHTTPServer(cfg, gin.New())
|
||||
require.NotNil(t, srv.Protocols)
|
||||
require.True(t, srv.Protocols.UnencryptedHTTP2())
|
||||
require.True(t, srv.Protocols.HTTP1())
|
||||
}
|
||||
|
||||
func TestHTTPServerRejectsOversizedHTTP1Header(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
srv := ProvideHTTPServer(ingressTestConfig(), r)
|
||||
addr, stop := serveIngressTestServer(t, srv)
|
||||
defer stop()
|
||||
|
||||
conn, err := net.DialTimeout("tcp", addr, time.Second)
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = conn.Close() }()
|
||||
_ = conn.SetDeadline(time.Now().Add(3 * time.Second))
|
||||
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Fill: "+strings.Repeat("a", 32*1024)+"\r\n\r\n")
|
||||
require.NoError(t, err)
|
||||
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
require.Equal(t, http.StatusRequestHeaderFieldsTooLarge, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestHTTPServerClosesSlowIncompleteHeader(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
srv := ProvideHTTPServer(ingressTestConfig(), r)
|
||||
addr, stop := serveIngressTestServer(t, srv)
|
||||
defer stop()
|
||||
|
||||
conn, err := net.DialTimeout("tcp", addr, time.Second)
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = conn.Close() }()
|
||||
_, err = io.WriteString(conn, "GET / HTTP/1.1\r\nHost: test\r\nX-Slow:")
|
||||
require.NoError(t, err)
|
||||
time.Sleep(1200 * time.Millisecond)
|
||||
_ = conn.SetReadDeadline(time.Now().Add(time.Second))
|
||||
_, err = bufio.NewReader(conn).ReadByte()
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestHTTPServerGlobalBodyLimit(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.POST("/", func(c *gin.Context) {
|
||||
_, err := io.ReadAll(c.Request.Body)
|
||||
if err != nil {
|
||||
var maxErr *http.MaxBytesError
|
||||
if errors.As(err, &maxErr) {
|
||||
c.Status(http.StatusRequestEntityTooLarge)
|
||||
return
|
||||
}
|
||||
}
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
srv := ProvideHTTPServer(ingressTestConfig(), r)
|
||||
req, err := http.NewRequest(http.MethodPost, "/", strings.NewReader(strings.Repeat("x", 1025)))
|
||||
require.NoError(t, err)
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Handler.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusRequestEntityTooLarge, rec.Code)
|
||||
}
|
||||
|
||||
func serveIngressTestServer(t *testing.T, srv *http.Server) (string, func()) {
|
||||
t.Helper()
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
go func() { _ = srv.Serve(ln) }()
|
||||
return ln.Addr().String(), func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
_ = srv.Shutdown(ctx)
|
||||
}
|
||||
}
|
||||
@@ -15,6 +15,8 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const maxAPIKeyAuthorizationHeaderBytes = service.MaxAPIKeyCredentialBytes + 128
|
||||
|
||||
// NewAPIKeyAuthMiddleware 创建 API Key 认证中间件
|
||||
func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) APIKeyAuthMiddleware {
|
||||
return APIKeyAuthMiddleware(apiKeyAuthWithSubscription(apiKeyService, subscriptionService, cfg))
|
||||
@@ -32,10 +34,23 @@ func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionS
|
||||
func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// ── 1. 提取 API Key ──────────────────────────────────────────
|
||||
if rejectInvalidAuthAbuse(c, apiKeyService) {
|
||||
AbortWithError(c, http.StatusTooManyRequests, "INVALID_AUTH_RATE_LIMITED", "Too many invalid authentication attempts; retry later")
|
||||
return
|
||||
}
|
||||
|
||||
if apiKeyHeadersTooLarge(c) {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
|
||||
return
|
||||
}
|
||||
|
||||
queryKey := strings.TrimSpace(c.Query("key"))
|
||||
queryApiKey := strings.TrimSpace(c.Query("api_key"))
|
||||
if queryKey != "" || queryApiKey != "" {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
|
||||
AbortWithError(c, 400, "api_key_in_query_deprecated", "API key in query parameter is deprecated. Please use Authorization header instead.")
|
||||
return
|
||||
}
|
||||
@@ -56,6 +71,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
if apiKeyString == "" {
|
||||
apiKeyString = c.GetHeader("x-api-key")
|
||||
}
|
||||
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
AbortWithError(c, http.StatusUnauthorized, "INVALID_API_KEY", "Invalid API key")
|
||||
return
|
||||
}
|
||||
|
||||
// 如果x-api-key header中没有,尝试从x-goog-api-key header中提取(Gemini CLI兼容)
|
||||
if apiKeyString == "" {
|
||||
@@ -64,6 +85,12 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
|
||||
// 如果所有header都没有API key
|
||||
if apiKeyString == "" {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
if hasAPIKeyCredentialInput(c) {
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
} else {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
|
||||
}
|
||||
AbortWithError(c, 401, "API_KEY_REQUIRED", "API key is required in Authorization header (Bearer scheme), x-api-key header, or x-goog-api-key header")
|
||||
return
|
||||
}
|
||||
@@ -73,9 +100,16 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrAPIKeyNotFound) {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
AbortWithError(c, 401, "INVALID_API_KEY", "Invalid API key")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
|
||||
AbortWithError(c, http.StatusServiceUnavailable, "API_KEY_AUTH_OVERLOADED", "API key authentication is temporarily unavailable")
|
||||
return
|
||||
}
|
||||
AbortWithError(c, 500, "INTERNAL_ERROR", "Failed to validate API key")
|
||||
return
|
||||
}
|
||||
@@ -90,6 +124,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
if !apiKey.IsActive() &&
|
||||
apiKey.Status != service.StatusAPIKeyExpired &&
|
||||
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
|
||||
AbortWithError(c, 401, "API_KEY_DISABLED", "API key is disabled")
|
||||
return
|
||||
}
|
||||
@@ -104,6 +139,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
clientIP = "unknown"
|
||||
}
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
|
||||
MarkIngressRejected(c, IngressRejectIPRestricted)
|
||||
AbortWithError(c, 403, "ACCESS_DENIED", fmt.Sprintf("Access denied. Your IP is %s", clientIP))
|
||||
return
|
||||
}
|
||||
@@ -117,6 +153,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
|
||||
// 检查用户状态
|
||||
if !apiKey.User.IsActive() {
|
||||
MarkIngressRejected(c, IngressRejectUserInactive)
|
||||
AbortWithError(c, 401, "USER_INACTIVE", "User account is not active")
|
||||
return
|
||||
}
|
||||
@@ -250,6 +287,24 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
}
|
||||
}
|
||||
|
||||
func apiKeyHeadersTooLarge(c *gin.Context) bool {
|
||||
if c == nil {
|
||||
return false
|
||||
}
|
||||
return len(c.GetHeader("Authorization")) > maxAPIKeyAuthorizationHeaderBytes ||
|
||||
len(c.GetHeader("x-api-key")) > service.MaxAPIKeyCredentialBytes ||
|
||||
len(c.GetHeader("x-goog-api-key")) > service.MaxAPIKeyCredentialBytes
|
||||
}
|
||||
|
||||
func hasAPIKeyCredentialInput(c *gin.Context) bool {
|
||||
if c == nil {
|
||||
return false
|
||||
}
|
||||
return c.GetHeader("Authorization") != "" ||
|
||||
c.GetHeader("x-api-key") != "" ||
|
||||
c.GetHeader("x-goog-api-key") != ""
|
||||
}
|
||||
|
||||
func isAsyncImageTaskRead(method, path string) bool {
|
||||
if method != http.MethodGet {
|
||||
return false
|
||||
@@ -321,6 +376,11 @@ func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool
|
||||
return false
|
||||
}
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
|
||||
if code == "GROUP_DELETED" {
|
||||
MarkIngressRejected(c, IngressRejectGroupDeleted)
|
||||
} else {
|
||||
MarkIngressRejected(c, IngressRejectGroupDisabled)
|
||||
}
|
||||
AbortWithError(c, 403, code, message)
|
||||
return true
|
||||
}
|
||||
@@ -330,6 +390,7 @@ func abortIfAPIKeyGroupNotAllowed(c *gin.Context, apiKey *service.APIKey) bool {
|
||||
return false
|
||||
}
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
|
||||
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
|
||||
AbortWithError(c, 403, "GROUP_NOT_ALLOWED", "API Key 所属专属分组不再允许当前用户使用")
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -24,22 +24,53 @@ func APIKeyAuthGoogle(apiKeyService *service.APIKeyService, cfg *config.Config)
|
||||
// It is intended for Gemini native endpoints (/v1beta) to match Gemini SDK expectations.
|
||||
func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if rejectInvalidAuthAbuse(c, apiKeyService) {
|
||||
abortWithGoogleError(c, 429, "Too many invalid authentication attempts; retry later")
|
||||
return
|
||||
}
|
||||
if apiKeyHeadersTooLarge(c) {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
abortWithGoogleError(c, 401, "Invalid API key")
|
||||
return
|
||||
}
|
||||
if v := strings.TrimSpace(c.Query("api_key")); v != "" {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectQueryAPIKeyDeprecated)
|
||||
abortWithGoogleError(c, 400, "Query parameter api_key is deprecated. Use Authorization header or key instead.")
|
||||
return
|
||||
}
|
||||
apiKeyString := extractAPIKeyForGoogle(c)
|
||||
if apiKeyString == "" {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
if hasAPIKeyCredentialInput(c) {
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
} else {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyRequired)
|
||||
}
|
||||
abortWithGoogleError(c, 401, "API key is required")
|
||||
return
|
||||
}
|
||||
if len(apiKeyString) > service.MaxAPIKeyCredentialBytes {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
abortWithGoogleError(c, 401, "Invalid API key")
|
||||
return
|
||||
}
|
||||
|
||||
apiKey, err := apiKeyService.GetByKey(c.Request.Context(), apiKeyString)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrAPIKeyNotFound) {
|
||||
recordInvalidAuthFailure(c, apiKeyService)
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
abortWithGoogleError(c, 401, "Invalid API key")
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrAPIKeyAuthOverloaded) {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyAuthOverloaded)
|
||||
abortWithGoogleError(c, 503, "API key authentication is temporarily unavailable")
|
||||
return
|
||||
}
|
||||
abortWithGoogleError(c, 500, "Failed to validate API key")
|
||||
return
|
||||
}
|
||||
@@ -53,6 +84,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
|
||||
if !apiKey.IsActive() &&
|
||||
apiKey.Status != service.StatusAPIKeyExpired &&
|
||||
apiKey.Status != service.StatusAPIKeyQuotaExhausted {
|
||||
MarkIngressRejected(c, IngressRejectAPIKeyDisabled)
|
||||
abortWithGoogleError(c, 401, "API key is disabled")
|
||||
return
|
||||
}
|
||||
@@ -66,6 +98,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
|
||||
clientIP = "unknown"
|
||||
}
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonIPRestriction)
|
||||
MarkIngressRejected(c, IngressRejectIPRestricted)
|
||||
abortWithGoogleError(c, 403, fmt.Sprintf("Access denied. Your IP is %s", clientIP))
|
||||
return
|
||||
}
|
||||
@@ -76,17 +109,24 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
|
||||
return
|
||||
}
|
||||
if !apiKey.User.IsActive() {
|
||||
MarkIngressRejected(c, IngressRejectUserInactive)
|
||||
abortWithGoogleError(c, 401, "User account is not active")
|
||||
return
|
||||
}
|
||||
if _, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
|
||||
if code, message, ok := validateAPIKeyGroupAvailable(apiKey); !ok {
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
|
||||
if code == "GROUP_DELETED" {
|
||||
MarkIngressRejected(c, IngressRejectGroupDeleted)
|
||||
} else {
|
||||
MarkIngressRejected(c, IngressRejectGroupDisabled)
|
||||
}
|
||||
abortWithGoogleError(c, 403, message)
|
||||
return
|
||||
}
|
||||
// 专属分组授权校验:用户对该专属分组的授权被撤销后应拒绝(与主中间件一致,防止越权)。
|
||||
if !validateAPIKeyGroupAllowed(apiKey) {
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
|
||||
MarkIngressRejected(c, IngressRejectGroupNotAllowed)
|
||||
abortWithGoogleError(c, 403, "API Key 所属专属分组不再允许当前用户使用")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -18,6 +20,59 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGoogleAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
var calls atomic.Int32
|
||||
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
calls.Add(1)
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
}}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
r := gin.New()
|
||||
var reason IngressRejectReason
|
||||
var rejected bool
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
reason, rejected = GetIngressRejectReason(c)
|
||||
})
|
||||
r.Use(APIKeyAuthGoogle(svc, cfg))
|
||||
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
|
||||
req.Header.Set("x-goog-api-key", strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1))
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
require.Zero(t, calls.Load())
|
||||
require.True(t, rejected)
|
||||
require.Equal(t, IngressRejectInvalidAPIKey, reason)
|
||||
}
|
||||
|
||||
func TestGoogleAPIKeyAuthMarksLookupBulkheadRejection(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
return nil, service.ErrAPIKeyAuthOverloaded
|
||||
}}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
r := gin.New()
|
||||
var reason IngressRejectReason
|
||||
var rejected bool
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
reason, rejected = GetIngressRejectReason(c)
|
||||
})
|
||||
r.Use(APIKeyAuthGoogle(svc, cfg))
|
||||
r.GET("/v1beta/models", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1beta/models", nil)
|
||||
req.Header.Set("x-goog-api-key", "valid-shape")
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusServiceUnavailable, w.Code)
|
||||
require.True(t, rejected)
|
||||
require.Equal(t, IngressRejectAPIKeyAuthOverloaded, reason)
|
||||
}
|
||||
|
||||
type fakeAPIKeyRepo struct {
|
||||
getByKey func(ctx context.Context, key string) (*service.APIKey, error)
|
||||
updateLastUsed func(ctx context.Context, id int64, usedAt time.Time) error
|
||||
@@ -376,6 +431,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
},
|
||||
})
|
||||
var rejectReason IngressRejectReason
|
||||
var rejected bool
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
rejectReason, rejected = GetIngressRejectReason(c)
|
||||
})
|
||||
r.Use(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, &config.Config{}))
|
||||
r.GET("/v1beta/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
|
||||
|
||||
@@ -390,6 +451,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_InvalidKey(t *testing.T) {
|
||||
require.Equal(t, http.StatusUnauthorized, resp.Error.Code)
|
||||
require.Equal(t, "Invalid API key", resp.Error.Message)
|
||||
require.Equal(t, "UNAUTHENTICATED", resp.Error.Status)
|
||||
require.True(t, rejected)
|
||||
require.Equal(t, IngressRejectInvalidAPIKey, rejectReason)
|
||||
}
|
||||
|
||||
func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t *testing.T) {
|
||||
@@ -422,9 +485,12 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
|
||||
r := gin.New()
|
||||
var markedBusinessLimited bool
|
||||
var businessLimitedReason string
|
||||
var rejectReason IngressRejectReason
|
||||
var rejected bool
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
|
||||
rejectReason, rejected = GetIngressRejectReason(c)
|
||||
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
|
||||
businessLimitedReason, _ = v.(string)
|
||||
}
|
||||
@@ -452,6 +518,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_MarksUnavailableGroupBusinessLimited(t
|
||||
require.Equal(t, "API Key 所属分组已删除", resp.Error.Message)
|
||||
require.True(t, markedBusinessLimited)
|
||||
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
|
||||
require.True(t, rejected)
|
||||
require.Equal(t, IngressRejectGroupDeleted, rejectReason)
|
||||
}
|
||||
|
||||
func TestApiKeyAuthWithSubscriptionGoogle_RepoError(t *testing.T) {
|
||||
|
||||
@@ -8,6 +8,8 @@ import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -19,6 +21,35 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAPIKeyAuthRejectsOversizedCredentialsBeforeLookup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
var calls atomic.Int32
|
||||
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
calls.Add(1)
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
}}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
|
||||
for _, headers := range []map[string]string{
|
||||
{"x-api-key": strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
|
||||
{"Authorization": "Bearer " + strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1)},
|
||||
{"Authorization": strings.Repeat("x", maxAPIKeyAuthorizationHeaderBytes+1)},
|
||||
} {
|
||||
r := gin.New()
|
||||
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
|
||||
r.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/t", nil)
|
||||
for name, value := range headers {
|
||||
req.Header.Set(name, value)
|
||||
}
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
}
|
||||
require.Zero(t, calls.Load())
|
||||
}
|
||||
|
||||
func TestSimpleModeBypassesQuotaCheck(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -436,6 +467,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
wantStatus int
|
||||
wantCode string
|
||||
wantMarked bool
|
||||
wantReject IngressRejectReason
|
||||
}{
|
||||
{
|
||||
name: "active group passes",
|
||||
@@ -460,6 +492,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
wantStatus: http.StatusForbidden,
|
||||
wantCode: "GROUP_DISABLED",
|
||||
wantMarked: true,
|
||||
wantReject: IngressRejectGroupDisabled,
|
||||
},
|
||||
{
|
||||
name: "deleted status group is forbidden",
|
||||
@@ -473,6 +506,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
wantStatus: http.StatusForbidden,
|
||||
wantCode: "GROUP_DELETED",
|
||||
wantMarked: true,
|
||||
wantReject: IngressRejectGroupDeleted,
|
||||
},
|
||||
{
|
||||
name: "missing group edge is forbidden",
|
||||
@@ -480,6 +514,7 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
wantStatus: http.StatusForbidden,
|
||||
wantCode: "GROUP_DELETED",
|
||||
wantMarked: true,
|
||||
wantReject: IngressRejectGroupDeleted,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -508,9 +543,12 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
router := gin.New()
|
||||
var markedBusinessLimited bool
|
||||
var businessLimitedReason string
|
||||
var rejectReason IngressRejectReason
|
||||
var rejected bool
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
|
||||
rejectReason, rejected = GetIngressRejectReason(c)
|
||||
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
|
||||
businessLimitedReason, _ = v.(string)
|
||||
}
|
||||
@@ -530,6 +568,8 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
require.Contains(t, w.Body.String(), tt.wantCode)
|
||||
}
|
||||
require.Equal(t, tt.wantMarked, markedBusinessLimited)
|
||||
require.Equal(t, tt.wantReject != "", rejected)
|
||||
require.Equal(t, tt.wantReject, rejectReason)
|
||||
if tt.wantMarked {
|
||||
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable, businessLimitedReason)
|
||||
}
|
||||
@@ -537,6 +577,112 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPIKeyAuthMarksOnlyExpectedIngressRejections(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
path string
|
||||
key string
|
||||
authHeader string
|
||||
repoErr error
|
||||
wantStatus int
|
||||
wantCode string
|
||||
wantReason IngressRejectReason
|
||||
}{
|
||||
{
|
||||
name: "query key deprecated",
|
||||
path: "/t?key=legacy",
|
||||
wantStatus: http.StatusBadRequest,
|
||||
wantCode: "api_key_in_query_deprecated",
|
||||
wantReason: IngressRejectQueryAPIKeyDeprecated,
|
||||
},
|
||||
{
|
||||
name: "missing key",
|
||||
path: "/t",
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
wantCode: "API_KEY_REQUIRED",
|
||||
wantReason: IngressRejectAPIKeyRequired,
|
||||
},
|
||||
{
|
||||
name: "malformed authorization",
|
||||
path: "/t",
|
||||
authHeader: "Basic not-a-bearer-key",
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
wantCode: "API_KEY_REQUIRED",
|
||||
wantReason: IngressRejectInvalidAPIKey,
|
||||
},
|
||||
{
|
||||
name: "oversized key",
|
||||
path: "/t",
|
||||
key: strings.Repeat("x", service.MaxAPIKeyCredentialBytes+1),
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
wantCode: "INVALID_API_KEY",
|
||||
wantReason: IngressRejectInvalidAPIKey,
|
||||
},
|
||||
{
|
||||
name: "invalid key",
|
||||
path: "/t",
|
||||
key: "invalid",
|
||||
repoErr: service.ErrAPIKeyNotFound,
|
||||
wantStatus: http.StatusUnauthorized,
|
||||
wantCode: "INVALID_API_KEY",
|
||||
wantReason: IngressRejectInvalidAPIKey,
|
||||
},
|
||||
{
|
||||
name: "repository failure remains operational error",
|
||||
path: "/t",
|
||||
key: "valid-shape",
|
||||
repoErr: errors.New("database unavailable"),
|
||||
wantStatus: http.StatusInternalServerError,
|
||||
wantCode: "INTERNAL_ERROR",
|
||||
},
|
||||
{
|
||||
name: "auth lookup bulkhead rejection is an admission rejection",
|
||||
path: "/t",
|
||||
key: "valid-shape",
|
||||
repoErr: service.ErrAPIKeyAuthOverloaded,
|
||||
wantStatus: http.StatusServiceUnavailable,
|
||||
wantCode: "API_KEY_AUTH_OVERLOADED",
|
||||
wantReason: IngressRejectAPIKeyAuthOverloaded,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
return nil, tt.repoErr
|
||||
}}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
apiKeyService := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
router := gin.New()
|
||||
var reason IngressRejectReason
|
||||
var rejected bool
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
reason, rejected = GetIngressRejectReason(c)
|
||||
})
|
||||
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
|
||||
router.GET("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, tt.path, nil)
|
||||
if tt.key != "" {
|
||||
req.Header.Set("x-api-key", tt.key)
|
||||
}
|
||||
if tt.authHeader != "" {
|
||||
req.Header.Set("Authorization", tt.authHeader)
|
||||
}
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
require.Equal(t, tt.wantStatus, w.Code)
|
||||
require.Contains(t, w.Body.String(), tt.wantCode)
|
||||
require.Equal(t, tt.wantReason != "", rejected)
|
||||
require.Equal(t, tt.wantReason, reason)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPIKeyAuthSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -686,9 +832,12 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
|
||||
router := gin.New()
|
||||
var markedBusinessLimited bool
|
||||
var businessLimitedReason string
|
||||
var rejectReason IngressRejectReason
|
||||
var rejected bool
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
markedBusinessLimited = service.HasOpsClientBusinessLimited(c)
|
||||
rejectReason, rejected = GetIngressRejectReason(c)
|
||||
if v, ok := c.Get(service.OpsClientBusinessLimitedReasonKey); ok {
|
||||
businessLimitedReason, _ = v.(string)
|
||||
}
|
||||
@@ -708,6 +857,8 @@ func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
|
||||
|
||||
require.Equal(t, http.StatusForbidden, w.Code)
|
||||
require.Contains(t, w.Body.String(), "not assigned to any group")
|
||||
require.True(t, rejected)
|
||||
require.Equal(t, IngressRejectGroupUnassigned, rejectReason)
|
||||
require.True(t, markedBusinessLimited)
|
||||
require.Equal(t, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned, businessLimitedReason)
|
||||
}
|
||||
@@ -822,7 +973,7 @@ func TestAPIKeyAuthIPRestrictionIncludesClientIPForBlacklistDenial(t *testing.T)
|
||||
requireAPIKeyAuthError(t, w, "ACCESS_DENIED", "Access denied. Your IP is 9.9.9.9")
|
||||
}
|
||||
|
||||
func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *testing.T) {
|
||||
func TestAPIKeyAuthIPRestrictionUsesConfiguredTrustedProxy(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
user := &service.User{
|
||||
@@ -855,7 +1006,7 @@ func TestAPIKeyAuthIPRestrictionCanTrustForwardedClientIPForReverseProxy(t *test
|
||||
cfg.SetTrustForwardedIPForAPIKeyACL(true)
|
||||
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
|
||||
router := gin.New()
|
||||
require.NoError(t, router.SetTrustedProxies(nil))
|
||||
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
|
||||
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
|
||||
router.GET("/t", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
@@ -906,7 +1057,7 @@ func TestAPIKeyAuthIPRestrictionUsesForwardedClientIPInDenialWhenTrusted(t *test
|
||||
cfg.SetTrustForwardedIPForAPIKeyACL(true)
|
||||
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
|
||||
router := gin.New()
|
||||
require.NoError(t, router.SetTrustedProxies(nil))
|
||||
require.NoError(t, router.SetTrustedProxies([]string{"9.9.9.9"}))
|
||||
router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
|
||||
router.GET("/t", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
|
||||
@@ -24,7 +24,14 @@ func ClientRequestID() gin.HandlerFunc {
|
||||
}
|
||||
|
||||
if v, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(v) != "" {
|
||||
c.Header(clientRequestIDHeader, strings.TrimSpace(v))
|
||||
var valid bool
|
||||
v, valid = normalizeCorrelationID(v)
|
||||
if !valid {
|
||||
v = uuid.New().String()
|
||||
}
|
||||
c.Header(clientRequestIDHeader, v)
|
||||
ctx := context.WithValue(c.Request.Context(), ctxkey.ClientRequestID, v)
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
|
||||
@@ -30,6 +31,23 @@ func TestClientRequestIDGeneratesAndExposesID(t *testing.T) {
|
||||
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
|
||||
}
|
||||
|
||||
func TestClientRequestIDBoundsExistingContextID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.Use(ClientRequestID())
|
||||
router.GET("/", func(c *gin.Context) {
|
||||
value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
|
||||
c.String(http.StatusOK, value)
|
||||
})
|
||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
req = req.WithContext(context.WithValue(req.Context(), ctxkey.ClientRequestID, strings.Repeat("x", 200)))
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
require.Len(t, w.Body.String(), 36)
|
||||
require.NotEqual(t, strings.Repeat("x", maxPersistentRequestIDBytes), w.Body.String())
|
||||
require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
|
||||
}
|
||||
|
||||
func TestClientRequestIDPreservesExistingContextID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"math"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// IngressRejectReason identifies expected gateway admission failures that must
|
||||
// not be treated as operational request errors.
|
||||
type IngressRejectReason string
|
||||
|
||||
const (
|
||||
IngressRejectQueryAPIKeyDeprecated IngressRejectReason = "query_api_key_deprecated"
|
||||
IngressRejectAPIKeyRequired IngressRejectReason = "api_key_required"
|
||||
IngressRejectInvalidAPIKey IngressRejectReason = "invalid_api_key"
|
||||
IngressRejectAPIKeyDisabled IngressRejectReason = "api_key_disabled"
|
||||
IngressRejectIPRestricted IngressRejectReason = "ip_restricted"
|
||||
IngressRejectUserInactive IngressRejectReason = "user_inactive"
|
||||
IngressRejectGroupDeleted IngressRejectReason = "group_deleted"
|
||||
IngressRejectGroupDisabled IngressRejectReason = "group_disabled"
|
||||
IngressRejectGroupNotAllowed IngressRejectReason = "group_not_allowed"
|
||||
IngressRejectGroupUnassigned IngressRejectReason = "group_unassigned"
|
||||
IngressRejectInvalidAuthRateLimited IngressRejectReason = "invalid_auth_rate_limited"
|
||||
IngressRejectAPIKeyAuthOverloaded IngressRejectReason = "api_key_auth_overloaded"
|
||||
)
|
||||
|
||||
const ingressRejectReasonContextKey = "ingress_reject_reason"
|
||||
|
||||
type IngressRejectRecorder interface {
|
||||
RecordIngressReject(reason, routeFamily, protocol, clientIP string, userID, apiKeyID int64)
|
||||
}
|
||||
|
||||
func invalidAuthClientKey(c *gin.Context) string {
|
||||
return normalizeIngressRejectIP(SecurityClientIP(c))
|
||||
}
|
||||
|
||||
func rejectInvalidAuthAbuse(c *gin.Context, apiKeyService interface {
|
||||
CheckInvalidAuthAbuse(string) (time.Duration, bool)
|
||||
}) bool {
|
||||
if c == nil || apiKeyService == nil {
|
||||
return false
|
||||
}
|
||||
retry, blocked := apiKeyService.CheckInvalidAuthAbuse(invalidAuthClientKey(c))
|
||||
if !blocked {
|
||||
return false
|
||||
}
|
||||
retrySeconds := int(math.Ceil(retry.Seconds()))
|
||||
if retrySeconds < 1 {
|
||||
retrySeconds = 1
|
||||
}
|
||||
c.Header("Retry-After", strconv.Itoa(retrySeconds))
|
||||
MarkIngressRejected(c, IngressRejectInvalidAuthRateLimited)
|
||||
return true
|
||||
}
|
||||
|
||||
func recordInvalidAuthFailure(c *gin.Context, apiKeyService interface {
|
||||
RecordInvalidAuthFailure(string)
|
||||
}) {
|
||||
if c == nil || apiKeyService == nil {
|
||||
return
|
||||
}
|
||||
apiKeyService.RecordInvalidAuthFailure(invalidAuthClientKey(c))
|
||||
}
|
||||
|
||||
type ingressRejectRecorderHolder struct{ recorder IngressRejectRecorder }
|
||||
|
||||
var activeIngressRejectRecorder atomic.Pointer[ingressRejectRecorderHolder]
|
||||
|
||||
func SetIngressRejectRecorder(recorder IngressRejectRecorder) {
|
||||
if recorder == nil {
|
||||
activeIngressRejectRecorder.Store(nil)
|
||||
return
|
||||
}
|
||||
activeIngressRejectRecorder.Store(&ingressRejectRecorderHolder{recorder: recorder})
|
||||
}
|
||||
|
||||
// MarkIngressRejected marks a request as rejected before gateway admission.
|
||||
func MarkIngressRejected(c *gin.Context, reason IngressRejectReason) {
|
||||
if c == nil || reason == "" {
|
||||
return
|
||||
}
|
||||
c.Set(ingressRejectReasonContextKey, reason)
|
||||
}
|
||||
|
||||
// GetIngressRejectReason returns the admission rejection reason, if any.
|
||||
func GetIngressRejectReason(c *gin.Context) (IngressRejectReason, bool) {
|
||||
if c == nil {
|
||||
return "", false
|
||||
}
|
||||
value, exists := c.Get(ingressRejectReasonContextKey)
|
||||
if !exists {
|
||||
return "", false
|
||||
}
|
||||
reason, ok := value.(IngressRejectReason)
|
||||
return reason, ok && reason != ""
|
||||
}
|
||||
|
||||
func recordIngressReject(c *gin.Context, reason IngressRejectReason) {
|
||||
holder := activeIngressRejectRecorder.Load()
|
||||
if holder == nil || holder.recorder == nil || c == nil || c.Request == nil {
|
||||
return
|
||||
}
|
||||
routeFamily, protocol := ingressRejectRoute(c.Request.URL.Path)
|
||||
clientIP := normalizeIngressRejectIP(SecurityClientIP(c))
|
||||
var userID, apiKeyID int64
|
||||
if apiKey, ok := GetAPIKeyFromContext(c); ok && apiKey != nil {
|
||||
apiKeyID = apiKey.ID
|
||||
if apiKey.User != nil {
|
||||
userID = apiKey.User.ID
|
||||
}
|
||||
} else if apiKey, ok := GetOpsFallbackAPIKey(c); ok && apiKey != nil {
|
||||
apiKeyID = apiKey.ID
|
||||
if apiKey.User != nil {
|
||||
userID = apiKey.User.ID
|
||||
}
|
||||
}
|
||||
holder.recorder.RecordIngressReject(string(reason), routeFamily, protocol, clientIP, userID, apiKeyID)
|
||||
}
|
||||
|
||||
func normalizeIngressRejectIP(raw string) string {
|
||||
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
return "0.0.0.0"
|
||||
}
|
||||
addr = addr.Unmap()
|
||||
if addr.Is6() {
|
||||
return netip.PrefixFrom(addr, 64).Masked().Addr().String()
|
||||
}
|
||||
return addr.String()
|
||||
}
|
||||
|
||||
func ingressRejectRoute(path string) (string, string) {
|
||||
path = strings.ToLower(strings.TrimSpace(path))
|
||||
switch {
|
||||
case strings.HasPrefix(path, "/antigravity/v1beta"):
|
||||
return "antigravity", "google"
|
||||
case strings.HasPrefix(path, "/v1beta"):
|
||||
return "gemini", "google"
|
||||
case strings.HasPrefix(path, "/backend-api/codex"):
|
||||
return "codex", "openai"
|
||||
case strings.HasPrefix(path, "/antigravity"):
|
||||
return "antigravity", "anthropic"
|
||||
case strings.Contains(path, "/messages"):
|
||||
return "messages", "anthropic"
|
||||
case strings.Contains(path, "/responses"):
|
||||
return "responses", "openai"
|
||||
case strings.Contains(path, "/chat/completions"):
|
||||
return "chat_completions", "openai"
|
||||
case strings.Contains(path, "/images"):
|
||||
return "images", "openai"
|
||||
case strings.Contains(path, "/videos"):
|
||||
return "videos", "openai"
|
||||
case strings.Contains(path, "/embeddings"):
|
||||
return "embeddings", "openai"
|
||||
case strings.Contains(path, "/models"):
|
||||
return "models", "openai"
|
||||
default:
|
||||
return "other", "gateway"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
ingressRejectAccessLogLimit = 20
|
||||
ingressRejectAccessLogWindow = time.Second
|
||||
ingressRejectDroppedSummaryPeriod = 30 * time.Second
|
||||
)
|
||||
|
||||
type ingressRejectAccessSampler struct {
|
||||
mu sync.Mutex
|
||||
limit int
|
||||
window time.Duration
|
||||
summaryPeriod time.Duration
|
||||
windowStart time.Time
|
||||
emitted int
|
||||
dropped uint64
|
||||
lastSummary time.Time
|
||||
}
|
||||
|
||||
func newIngressRejectAccessSampler(limit int, window, summaryPeriod time.Duration) *ingressRejectAccessSampler {
|
||||
return &ingressRejectAccessSampler{limit: limit, window: window, summaryPeriod: summaryPeriod}
|
||||
}
|
||||
|
||||
// allow applies one process-wide fixed-window budget. It stores no attacker
|
||||
// dimensions, so memory remains constant even for rotating keys and addresses.
|
||||
func (s *ingressRejectAccessSampler) allow(now time.Time) (allowed bool, droppedSummary uint64) {
|
||||
if s == nil || s.limit <= 0 || s.window <= 0 {
|
||||
return false, 0
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.windowStart.IsZero() || now.Sub(s.windowStart) >= s.window || now.Before(s.windowStart) {
|
||||
s.windowStart = now
|
||||
s.emitted = 0
|
||||
}
|
||||
if s.emitted < s.limit {
|
||||
s.emitted++
|
||||
return true, 0
|
||||
}
|
||||
s.dropped++
|
||||
if s.summaryPeriod > 0 && (s.lastSummary.IsZero() || now.Sub(s.lastSummary) >= s.summaryPeriod) {
|
||||
droppedSummary = s.dropped
|
||||
s.dropped = 0
|
||||
s.lastSummary = now
|
||||
}
|
||||
return false, droppedSummary
|
||||
}
|
||||
|
||||
var globalIngressRejectAccessSampler = newIngressRejectAccessSampler(
|
||||
ingressRejectAccessLogLimit,
|
||||
ingressRejectAccessLogWindow,
|
||||
ingressRejectDroppedSummaryPeriod,
|
||||
)
|
||||
@@ -0,0 +1,81 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIngressRejectAccessSamplerConcurrentGlobalLimit(t *testing.T) {
|
||||
sampler := newIngressRejectAccessSampler(10, time.Hour, time.Minute)
|
||||
now := time.Now()
|
||||
var allowed atomic.Int64
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 200; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if ok, _ := sampler.allow(now); ok {
|
||||
allowed.Add(1)
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
require.Equal(t, int64(10), allowed.Load())
|
||||
}
|
||||
|
||||
func TestLoggerIngressRejectSamplingIsBoundedAndSummarySkipsOpsSink(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
original := globalIngressRejectAccessSampler
|
||||
globalIngressRejectAccessSampler = newIngressRejectAccessSampler(2, time.Hour, time.Hour)
|
||||
t.Cleanup(func() { globalIngressRejectAccessSampler = original })
|
||||
sink := initMiddlewareTestLogger(t)
|
||||
router := gin.New()
|
||||
router.Use(Logger())
|
||||
router.GET("/v1/messages", func(c *gin.Context) {
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
c.Status(http.StatusUnauthorized)
|
||||
})
|
||||
for i := 0; i < 20; i++ {
|
||||
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/v1/messages", nil))
|
||||
}
|
||||
var accessEvents, summaries int
|
||||
for _, event := range sink.list() {
|
||||
switch event.Message {
|
||||
case "http request completed":
|
||||
accessEvents++
|
||||
case "ingress rejection access logs dropped":
|
||||
summaries++
|
||||
if skipped, _ := event.Fields[logger.OpsSystemLogSkipField].(bool); !skipped {
|
||||
t.Fatalf("dropped summary must skip ops system log sink")
|
||||
}
|
||||
}
|
||||
}
|
||||
require.Equal(t, 2, accessEvents)
|
||||
require.Equal(t, 1, summaries)
|
||||
}
|
||||
|
||||
func TestIngressRejectAccessSamplerDroppedSummaryIsLowFrequency(t *testing.T) {
|
||||
sampler := newIngressRejectAccessSampler(1, time.Hour, time.Second)
|
||||
now := time.Now()
|
||||
allowed, summary := sampler.allow(now)
|
||||
require.True(t, allowed)
|
||||
require.Zero(t, summary)
|
||||
|
||||
allowed, summary = sampler.allow(now.Add(100 * time.Millisecond))
|
||||
require.False(t, allowed)
|
||||
require.Equal(t, uint64(1), summary)
|
||||
allowed, summary = sampler.allow(now.Add(200 * time.Millisecond))
|
||||
require.False(t, allowed)
|
||||
require.Zero(t, summary)
|
||||
allowed, summary = sampler.allow(now.Add(2 * time.Second))
|
||||
require.False(t, allowed)
|
||||
require.Equal(t, uint64(2), summary)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type ingressRejectRecorderStub struct {
|
||||
mu sync.Mutex
|
||||
calls int
|
||||
clientIP string
|
||||
}
|
||||
|
||||
func (r *ingressRejectRecorderStub) RecordIngressReject(_, _, _, clientIP string, _, _ int64) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.calls++
|
||||
r.clientIP = clientIP
|
||||
}
|
||||
|
||||
func TestNormalizeIngressRejectIP(t *testing.T) {
|
||||
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::1"))
|
||||
require.Equal(t, "192.0.2.4", normalizeIngressRejectIP("::ffff:192.0.2.4"))
|
||||
require.Equal(t, "0.0.0.0", normalizeIngressRejectIP("not-an-ip"))
|
||||
}
|
||||
|
||||
func TestLoggerRecordsIngressRejectOnce(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := &ingressRejectRecorderStub{}
|
||||
SetIngressRejectRecorder(recorder)
|
||||
t.Cleanup(func() { SetIngressRejectRecorder(nil) })
|
||||
router := gin.New()
|
||||
router.Use(Logger())
|
||||
router.GET("/v1/messages", func(c *gin.Context) {
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
c.Status(http.StatusUnauthorized)
|
||||
})
|
||||
request := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
|
||||
request.RemoteAddr = "[2001:db8:abcd:1234:ffff::1]:1234"
|
||||
router.ServeHTTP(httptest.NewRecorder(), request)
|
||||
recorder.mu.Lock()
|
||||
require.Equal(t, 1, recorder.calls)
|
||||
require.Equal(t, "2001:db8:abcd:1234::", recorder.clientIP)
|
||||
recorder.mu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
//go:build unit
|
||||
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func invalidAuthAbuseTestConfig(threshold int) *config.Config {
|
||||
return &config.Config{
|
||||
RunMode: config.RunModeSimple,
|
||||
APIKeyAuth: config.APIKeyAuthCacheConfig{InvalidAbuse: config.InvalidAuthAbuseConfig{
|
||||
Enabled: true, Threshold: threshold, WindowSeconds: 60, BlockSeconds: 60, Capacity: 256,
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPIKeyAuthInvalidAbuseReturns429BeforeRepository(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repoCalls := 0
|
||||
repo := &stubApiKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
repoCalls++
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
}}
|
||||
cfg := invalidAuthAbuseTestConfig(3)
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
r := gin.New()
|
||||
var reason IngressRejectReason
|
||||
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
|
||||
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
|
||||
r.POST("/v1/messages", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
|
||||
requests := []*http.Request{
|
||||
httpRequest(t, "/v1/messages", "", ""),
|
||||
httpRequest(t, "/v1/messages", "Basic malformed", ""),
|
||||
httpRequest(t, "/v1/messages", "", "random-invalid-key"),
|
||||
}
|
||||
for _, req := range requests {
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
require.NotEqual(t, http.StatusTooManyRequests, w.Code)
|
||||
}
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, httpRequest(t, "/v1/messages", "", "another-random-key"))
|
||||
require.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||
require.Equal(t, "60", w.Header().Get("Retry-After"))
|
||||
require.Contains(t, w.Body.String(), "INVALID_AUTH_RATE_LIMITED")
|
||||
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
|
||||
require.Equal(t, 1, repoCalls, "rate-limited request must not reach the repository")
|
||||
}
|
||||
|
||||
func TestGoogleAPIKeyAuthInvalidAbuseReturnsProtocol429(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repoCalls := 0
|
||||
repo := fakeAPIKeyRepo{getByKey: func(context.Context, string) (*service.APIKey, error) {
|
||||
repoCalls++
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
}}
|
||||
cfg := invalidAuthAbuseTestConfig(2)
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
r := gin.New()
|
||||
var reason IngressRejectReason
|
||||
r.Use(func(c *gin.Context) { c.Next(); reason, _ = GetIngressRejectReason(c) })
|
||||
r.Use(APIKeyAuthGoogle(svc, cfg))
|
||||
r.POST("/v1beta/models/test:generateContent", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
for _, key := range []string{"random-1", "random-2"} {
|
||||
w := httptest.NewRecorder()
|
||||
req := httpRequest(t, "/v1beta/models/test:generateContent", "", key)
|
||||
req.Header.Del("x-api-key")
|
||||
req.Header.Set("x-goog-api-key", key)
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
req := httpRequest(t, "/v1beta/models/test:generateContent", "", "random-3")
|
||||
req.Header.Del("x-api-key")
|
||||
req.Header.Set("x-goog-api-key", "random-3")
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusTooManyRequests, w.Code)
|
||||
require.Equal(t, "60", w.Header().Get("Retry-After"))
|
||||
require.Contains(t, w.Body.String(), "RESOURCE_EXHAUSTED")
|
||||
require.Equal(t, IngressRejectInvalidAuthRateLimited, reason)
|
||||
require.Equal(t, 2, repoCalls)
|
||||
}
|
||||
|
||||
func TestInvalidAuthAbuseDoesNotCountValidOrOperationalFailures(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
user := &service.User{ID: 1, Status: service.StatusActive, Role: service.RoleUser, Balance: 1}
|
||||
repo := &stubApiKeyRepo{getByKey: func(_ context.Context, key string) (*service.APIKey, error) {
|
||||
switch key {
|
||||
case "valid-key":
|
||||
return &service.APIKey{ID: 1, UserID: 1, Key: key, Status: service.StatusActive, User: user}, nil
|
||||
case "db-error":
|
||||
return nil, errors.New("database unavailable")
|
||||
default:
|
||||
return nil, service.ErrAPIKeyNotFound
|
||||
}
|
||||
}}
|
||||
cfg := invalidAuthAbuseTestConfig(10)
|
||||
svc := service.NewAPIKeyService(repo, nil, nil, nil, nil, nil, cfg)
|
||||
r := gin.New()
|
||||
r.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(svc, nil, cfg)))
|
||||
r.POST("/t", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
|
||||
for _, tc := range []struct {
|
||||
key string
|
||||
want int
|
||||
}{{"invalid", 401}, {"valid-key", 200}, {"db-error", 500}, {"db-error", 500}} {
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, httpRequest(t, "/t", "", tc.key))
|
||||
require.Equal(t, tc.want, w.Code)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
req := httpRequest(t, "/t", "", "")
|
||||
req.Header.Set("x-goog-api-key", "valid-key")
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusOK, w.Code)
|
||||
require.Equal(t, uint64(1), svc.InvalidAuthAbuseHealth().Recorded)
|
||||
}
|
||||
|
||||
func TestNormalizeIngressRejectIPGroupsIPv6By64(t *testing.T) {
|
||||
require.Equal(t, "2001:db8:abcd:1234::", normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"))
|
||||
require.Equal(t, normalizeIngressRejectIP("2001:db8:abcd:1234:1111::1"), normalizeIngressRejectIP("2001:db8:abcd:1234:ffff::2"))
|
||||
}
|
||||
|
||||
func httpRequest(t *testing.T, path, authorization, apiKey string) *http.Request {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, nil)
|
||||
req.RemoteAddr = "203.0.113.10:12345"
|
||||
if authorization != "" {
|
||||
req.Header.Set("Authorization", authorization)
|
||||
}
|
||||
if apiKey != "" {
|
||||
req.Header.Set("x-api-key", apiKey)
|
||||
}
|
||||
return req
|
||||
}
|
||||
@@ -37,6 +37,21 @@ func Logger() gin.HandlerFunc {
|
||||
accountID, hasAccountID := c.Request.Context().Value(ctxkey.AccountID).(int64)
|
||||
platform, _ := c.Request.Context().Value(ctxkey.Platform).(string)
|
||||
model, _ := c.Request.Context().Value(ctxkey.Model).(string)
|
||||
reason, rejected := GetIngressRejectReason(c)
|
||||
if rejected {
|
||||
recordIngressReject(c, reason)
|
||||
allowed, droppedSummary := globalIngressRejectAccessSampler.allow(endTime)
|
||||
if droppedSummary > 0 {
|
||||
logger.FromContext(c.Request.Context()).Info("ingress rejection access logs dropped",
|
||||
zap.String("component", "http.access"),
|
||||
zap.Uint64("dropped_count", droppedSummary),
|
||||
zap.Bool(logger.OpsSystemLogSkipField, true),
|
||||
)
|
||||
}
|
||||
if !allowed {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
fields := []zap.Field{
|
||||
zap.String("component", "http.access"),
|
||||
@@ -47,6 +62,12 @@ func Logger() gin.HandlerFunc {
|
||||
zap.String("method", method),
|
||||
zap.String("path", path),
|
||||
}
|
||||
if rejected {
|
||||
fields = append(fields,
|
||||
zap.String("ingress_reject_reason", string(reason)),
|
||||
zap.Bool(logger.OpsSystemLogSkipField, true),
|
||||
)
|
||||
}
|
||||
if hasAccountID && accountID > 0 {
|
||||
fields = append(fields, zap.Int64("account_id", accountID))
|
||||
}
|
||||
|
||||
@@ -121,6 +121,7 @@ func RequireGroupAssignment(settingService *service.SettingService, writeError G
|
||||
return
|
||||
}
|
||||
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnassigned)
|
||||
MarkIngressRejected(c, IngressRejectGroupUnassigned)
|
||||
writeError(c, http.StatusForbidden, "API Key is not assigned to any group and cannot be used. Please contact the administrator to assign it to a group.")
|
||||
c.Abort()
|
||||
}
|
||||
|
||||
@@ -112,6 +112,26 @@ func TestRequestLogger_KeepIncomingRequestID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestLoggerBoundsIncomingRequestID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(RequestLogger())
|
||||
r.GET("/t", func(c *gin.Context) {
|
||||
reqID, _ := c.Request.Context().Value(ctxkey.RequestID).(string)
|
||||
if len(reqID) != 36 {
|
||||
t.Fatalf("request_id length=%d", len(reqID))
|
||||
}
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/t", nil)
|
||||
req.Header.Set(requestIDHeader, strings.Repeat("r", 1024))
|
||||
r.ServeHTTP(w, req)
|
||||
if got := len(w.Header().Get(requestIDHeader)); got != 36 {
|
||||
t.Fatalf("response request_id length=%d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
sink := initMiddlewareTestLogger(t)
|
||||
@@ -180,11 +200,43 @@ func TestLogger_AccessLogIncludesCoreFields(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
|
||||
func TestLogger_IngressRejectRemainsInStandardAccessLog(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
sink := initMiddlewareTestLogger(t)
|
||||
r := gin.New()
|
||||
r.Use(Logger())
|
||||
r.GET("/v1/messages", func(c *gin.Context) {
|
||||
MarkIngressRejected(c, IngressRejectInvalidAPIKey)
|
||||
c.Status(http.StatusUnauthorized)
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/messages", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status=%d", w.Code)
|
||||
}
|
||||
events := sink.list()
|
||||
if len(events) != 1 {
|
||||
t.Fatalf("events=%d, want 1", len(events))
|
||||
}
|
||||
if got := events[0].Fields["ingress_reject_reason"]; got != string(IngressRejectInvalidAPIKey) {
|
||||
t.Fatalf("ingress_reject_reason=%v", got)
|
||||
}
|
||||
if got, _ := events[0].Fields[logger.OpsSystemLogSkipField].(bool); !got {
|
||||
t.Fatalf("%s must be true", logger.OpsSystemLogSkipField)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogger_AccessLogUsesForwardedClientIPFromTrustedProxy(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
sink := initMiddlewareTestLogger(t)
|
||||
|
||||
r := gin.New()
|
||||
if err := r.SetTrustedProxies([]string{"104.23.251.120"}); err != nil {
|
||||
t.Fatalf("set trusted proxies: %v", err)
|
||||
}
|
||||
r.Use(Logger())
|
||||
r.GET("/api/test", func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
@@ -193,7 +245,7 @@ func TestLogger_AccessLogUsesForwardedClientIP(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/test", nil)
|
||||
req.RemoteAddr = "104.23.251.120:443"
|
||||
req.Header.Set("CF-Connecting-IP", "203.0.113.42")
|
||||
req.Header.Set("X-Forwarded-For", "203.0.113.42")
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d", w.Code)
|
||||
|
||||
@@ -21,14 +21,15 @@ func RequestLogger() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
requestID := strings.TrimSpace(c.GetHeader(requestIDHeader))
|
||||
if requestID == "" {
|
||||
requestID, validRequestID := normalizeCorrelationID(c.GetHeader(requestIDHeader))
|
||||
if !validRequestID {
|
||||
requestID = uuid.NewString()
|
||||
}
|
||||
c.Header(requestIDHeader, requestID)
|
||||
|
||||
ctx := context.WithValue(c.Request.Context(), ctxkey.RequestID, requestID)
|
||||
clientRequestID, _ := ctx.Value(ctxkey.ClientRequestID).(string)
|
||||
clientRequestID, _ = normalizeCorrelationID(clientRequestID)
|
||||
|
||||
requestLogger := logger.With(
|
||||
zap.String("component", "http"),
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
maxPersistentRequestIDBytes = 64
|
||||
maxPersistentUserAgentBytes = 512
|
||||
)
|
||||
|
||||
// normalizePersistentText bounds attacker-controlled metadata before it reaches
|
||||
// logs or database columns while preserving valid UTF-8 content.
|
||||
func normalizePersistentText(value string, maxBytes int) string {
|
||||
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
|
||||
if maxBytes <= 0 || len(value) <= maxBytes {
|
||||
return value
|
||||
}
|
||||
value = value[:maxBytes]
|
||||
for !utf8.ValidString(value) {
|
||||
value = value[:len(value)-1]
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeCorrelationID(value string) (string, bool) {
|
||||
value = strings.TrimSpace(strings.ToValidUTF8(value, ""))
|
||||
return value, value != "" && len(value) <= maxPersistentRequestIDBytes
|
||||
}
|
||||
@@ -13,14 +13,15 @@ import (
|
||||
// SessionBindingContext 全局中间件:将请求的客户端 IP 与 User-Agent 注入
|
||||
// request context,供 token 签发路径(登录 / 刷新 / OAuth 回调)读取并写入会话绑定,
|
||||
// 同时作为审计日志、会话绑定校验的统一客户端 IP 来源。
|
||||
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关:
|
||||
// 开启时信任反代转发头(CF-Connecting-IP / X-Real-IP / X-Forwarded-For),
|
||||
// 关闭时走 trusted_proxies 解析链,避免不可信头伪造绕过绑定。
|
||||
// IP 取值与 API Key IP 限制共用 Gin trusted_proxies 解析链;旧设置开关
|
||||
// 仅为配置兼容保留,不能单独使直连请求的转发头变为可信。
|
||||
func SessionBindingContext(cfg *config.Config) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userAgent := normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes)
|
||||
c.Request.Header.Set("User-Agent", userAgent)
|
||||
binding := &service.SessionBinding{
|
||||
IP: ip.GetSecurityClientIP(c, cfg.TrustForwardedIPForAPIKeyACL()),
|
||||
UserAgent: c.Request.UserAgent(),
|
||||
UserAgent: userAgent,
|
||||
}
|
||||
c.Request = c.Request.WithContext(service.WithSessionBinding(c.Request.Context(), binding))
|
||||
c.Next()
|
||||
@@ -36,7 +37,7 @@ func requestSessionBinding(c *gin.Context) *service.SessionBinding {
|
||||
}
|
||||
return &service.SessionBinding{
|
||||
IP: ip.GetTrustedClientIP(c),
|
||||
UserAgent: c.Request.UserAgent(),
|
||||
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -93,7 +94,7 @@ func enforceSessionBinding(
|
||||
Method: c.Request.Method,
|
||||
Path: path,
|
||||
ClientIP: binding.IP,
|
||||
UserAgent: c.Request.UserAgent(),
|
||||
UserAgent: normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes),
|
||||
StatusCode: 401,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ package middleware
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
@@ -13,9 +14,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// 反代场景:RemoteAddr 为 127.0.0.1,真实客户端 IP 在 X-Real-IP 中。
|
||||
// 会话绑定注入与审计 IP 必须与 API Key IP 限制共用「信任反代传递的客户端 IP」开关语义。
|
||||
func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
|
||||
func TestSessionBindingContextDoesNotTrustHeadersWithoutTrustedProxy(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
for _, tc := range []struct {
|
||||
@@ -24,7 +23,7 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
|
||||
wantIP string
|
||||
}{
|
||||
{name: "trust disabled records proxy address", trustForwarded: false, wantIP: "127.0.0.1"},
|
||||
{name: "trust enabled records forwarded client IP", trustForwarded: true, wantIP: "1.2.3.4"},
|
||||
{name: "legacy trust toggle cannot bypass trusted proxies", trustForwarded: true, wantIP: "127.0.0.1"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cfg := &config.Config{}
|
||||
@@ -54,6 +53,23 @@ func TestSessionBindingContextHonorsTrustForwardedToggle(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionBindingContextBoundsPersistedUserAgent(t *testing.T) {
|
||||
cfg := &config.Config{}
|
||||
r := gin.New()
|
||||
r.Use(SessionBindingContext(cfg))
|
||||
r.GET("/t", func(c *gin.Context) {
|
||||
binding := service.SessionBindingFromContext(c.Request.Context())
|
||||
require.Len(t, binding.UserAgent, maxPersistentUserAgentBytes)
|
||||
require.Equal(t, binding.UserAgent, c.Request.UserAgent())
|
||||
c.Status(200)
|
||||
})
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest("GET", "/t", nil)
|
||||
req.Header.Set("User-Agent", strings.Repeat("u", 2048))
|
||||
r.ServeHTTP(w, req)
|
||||
require.Equal(t, 200, w.Code)
|
||||
}
|
||||
|
||||
// 未经过 SessionBindingContext 注入时(异常挂载顺序/单测直调),回退 trusted_proxies 链,
|
||||
// 等价于开关关闭时的历史行为。
|
||||
func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
|
||||
@@ -75,8 +91,6 @@ func TestSecurityClientIPFallsBackWithoutInjectedBinding(t *testing.T) {
|
||||
require.Equal(t, "9.9.9.9", w.Body.String())
|
||||
}
|
||||
|
||||
// requestSessionBinding 优先取注入值:开关开启时校验哈希必须基于注入的转发 IP 计算,
|
||||
// 与 token 签发路径取值一致,否则同一客户端会被误判为指纹变化。
|
||||
func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -84,7 +98,7 @@ func TestRequestSessionBindingPrefersInjectedBinding(t *testing.T) {
|
||||
cfg.SetTrustForwardedIPForAPIKeyACL(true)
|
||||
|
||||
r := gin.New()
|
||||
require.NoError(t, r.SetTrustedProxies(nil))
|
||||
require.NoError(t, r.SetTrustedProxies([]string{"127.0.0.1"}))
|
||||
r.Use(SessionBindingContext(cfg))
|
||||
r.GET("/t", func(c *gin.Context) {
|
||||
issued := &service.SessionBinding{IP: "1.2.3.4", UserAgent: "test-agent"}
|
||||
|
||||
@@ -35,6 +35,7 @@ func SetupRouter(
|
||||
cfg *config.Config,
|
||||
redisClient *redis.Client,
|
||||
) *gin.Engine {
|
||||
middleware2.SetIngressRejectRecorder(opsService)
|
||||
// 缓存 iframe 页面的 origin 列表,用于动态注入 CSP frame-src
|
||||
var cachedFrameOrigins atomic.Pointer[[]string]
|
||||
emptyOrigins := []string{}
|
||||
@@ -55,7 +56,7 @@ func SetupRouter(
|
||||
// 应用中间件
|
||||
r.Use(middleware2.RequestLogger())
|
||||
// 将客户端 IP + UA 注入 request context,供 token 签发/会话绑定/审计日志统一读取。
|
||||
// IP 取值与 API Key IP 限制共用「信任反代传递的客户端 IP」系统开关。
|
||||
// IP 取值与 API Key IP 限制共用 server.trusted_proxies 信任链。
|
||||
r.Use(middleware2.SessionBindingContext(cfg))
|
||||
r.Use(middleware2.Logger())
|
||||
r.Use(middleware2.CORS(cfg.CORS))
|
||||
|
||||
@@ -235,6 +235,11 @@ func registerOpsRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
|
||||
ops.GET("/request-errors/:id/upstream-errors", h.Admin.Ops.ListRequestErrorUpstreamErrors)
|
||||
ops.PUT("/request-errors/:id/resolve", h.Admin.Ops.ResolveRequestError)
|
||||
|
||||
// Bounded ingress-admission rejection aggregates.
|
||||
ops.GET("/ingress-rejections", h.Admin.Ops.ListIngressRejects)
|
||||
ops.GET("/ingress-rejections/health", h.Admin.Ops.GetIngressRejectHealth)
|
||||
ops.GET("/auth-cache-invalidation/health", h.Admin.Ops.GetAuthCacheInvalidationHealth)
|
||||
|
||||
// Upstream errors (independent upstream failures)
|
||||
ops.GET("/upstream-errors", h.Admin.Ops.ListUpstreamErrors)
|
||||
ops.GET("/upstream-errors/:id", h.Admin.Ops.GetUpstreamError)
|
||||
|
||||
@@ -23,6 +23,7 @@ func RegisterGatewayRoutes(
|
||||
cfg *config.Config,
|
||||
) {
|
||||
bodyLimit := middleware.RequestBodyLimit(cfg.Gateway.MaxBodySize)
|
||||
textBodyLimit := middleware.RequestBodyLimit(cfg.Gateway.TextMaxBodySize)
|
||||
clientRequestID := middleware.ClientRequestID()
|
||||
opsErrorLogger := handler.OpsErrorLoggerMiddleware(opsService)
|
||||
endpointNorm := handler.InboundEndpointMiddleware()
|
||||
@@ -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
|
||||
|
||||
@@ -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.cancel()
|
||||
a.recordMu.Unlock()
|
||||
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, ©Item)
|
||||
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, ©Item)
|
||||
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 {
|
||||
|
||||
@@ -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(默认)保持精确 =,管理端语义不变。
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
# =========================================================================
|
||||
|
||||
@@ -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.
|
||||
@@ -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 |
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -709,6 +721,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
|
||||
|
||||
@@ -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: '自动刷新',
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user