diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index 496473bc88..c5f7e315a9 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -103,6 +103,7 @@ func provideCleanup( paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, quotaFlusher *service.UserPlatformQuotaUsageFlusher, + upstreamBillingProbe *service.UpstreamBillingProbeService, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -279,6 +280,12 @@ func provideCleanup( } return nil }}, + {"UpstreamBillingProbeService", func() error { + if upstreamBillingProbe != nil { + upstreamBillingProbe.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 0c466a6e8a..c7ab9eeb88 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -250,7 +250,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { paymentHandler := admin.NewPaymentHandler(paymentService, paymentConfigService) affiliateHandler := admin.NewAffiliateHandler(affiliateService, adminService) complianceHandler := admin.NewComplianceHandler(settingService) - adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler) + upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db) + adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler, upstreamBillingProbeService) usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig) userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) @@ -290,7 +291,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) + 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) application := &Application{ Server: httpServer, Cleanup: v, @@ -350,6 +351,7 @@ func provideCleanup( paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, quotaFlusher *service.UserPlatformQuotaUsageFlusher, + upstreamBillingProbe *service.UpstreamBillingProbeService, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -525,6 +527,12 @@ func provideCleanup( } return nil }}, + {"UpstreamBillingProbeService", func() error { + if upstreamBillingProbe != nil { + upstreamBillingProbe.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 27707bc8c6..e1a16dc098 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -83,6 +83,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // paymentOrderExpiry nil, // channelMonitorRunner nil, // quotaFlusher + nil, // upstreamBillingProbe ) require.NotPanics(t, func() { diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 661647f681..d9b742003f 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -62,6 +62,12 @@ type AccountHandler struct { rpmCache service.RPMCache tokenCacheInvalidator service.TokenCacheInvalidator grokImportProber grokUsageProber + upstreamBillingProbe *service.UpstreamBillingProbeService +} + +// SetUpstreamBillingProbeService attaches the optional remote billing probe service. +func (h *AccountHandler) SetUpstreamBillingProbeService(probe *service.UpstreamBillingProbeService) { + h.upstreamBillingProbe = probe } // NewAccountHandler creates a new admin account handler diff --git a/backend/internal/handler/admin/account_upstream_billing_probe.go b/backend/internal/handler/admin/account_upstream_billing_probe.go new file mode 100644 index 0000000000..efaca8f73c --- /dev/null +++ b/backend/internal/handler/admin/account_upstream_billing_probe.go @@ -0,0 +1,122 @@ +package admin + +import ( + "strconv" + + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +type upstreamBillingProbeEnabledRequest struct { + Enabled *bool `json:"enabled" binding:"required"` +} + +type upstreamBillingProbeBatchRequest struct { + AccountIDs []int64 `json:"account_ids" binding:"required"` +} + +func (h *AccountHandler) GetUpstreamBillingProbeSettings(c *gin.Context) { + if h.upstreamBillingProbe == nil { + response.ErrorFrom(c, service.ErrUpstreamBillingProbeUnavailable) + return + } + settings, err := h.upstreamBillingProbe.GetSettings(c.Request.Context()) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, settings) +} + +func (h *AccountHandler) UpdateUpstreamBillingProbeSettings(c *gin.Context) { + if h.upstreamBillingProbe == nil { + response.ErrorFrom(c, service.ErrUpstreamBillingProbeUnavailable) + return + } + var req service.UpstreamBillingProbeSettings + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + if err := h.upstreamBillingProbe.UpdateSettings(c.Request.Context(), &req); err != nil { + response.ErrorFrom(c, err) + return + } + settings, err := h.upstreamBillingProbe.GetSettings(c.Request.Context()) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, settings) +} + +func (h *AccountHandler) SetUpstreamBillingProbeEnabled(c *gin.Context) { + if h.upstreamBillingProbe == nil { + response.ErrorFrom(c, service.ErrUpstreamBillingProbeUnavailable) + return + } + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil || accountID <= 0 { + response.BadRequest(c, "Invalid account ID") + return + } + var req upstreamBillingProbeEnabledRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + if err := h.upstreamBillingProbe.SetAccountEnabled(c.Request.Context(), accountID, *req.Enabled); err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, gin.H{"account_id": accountID, "enabled": *req.Enabled}) +} + +func (h *AccountHandler) ProbeUpstreamBilling(c *gin.Context) { + if h.upstreamBillingProbe == nil { + response.ErrorFrom(c, service.ErrUpstreamBillingProbeUnavailable) + return + } + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil || accountID <= 0 { + response.BadRequest(c, "Invalid account ID") + return + } + snapshot, err := h.upstreamBillingProbe.ProbeAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, service.UpstreamBillingProbeResult{AccountID: accountID, Snapshot: snapshot}) +} + +func (h *AccountHandler) ProbeUpstreamBillingBatch(c *gin.Context) { + if h.upstreamBillingProbe == nil { + response.ErrorFrom(c, service.ErrUpstreamBillingProbeUnavailable) + return + } + var req upstreamBillingProbeBatchRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + if len(req.AccountIDs) == 0 || len(req.AccountIDs) > service.UpstreamBillingProbeMaxBatchSize { + response.BadRequest(c, "account_ids must contain between 1 and 20 items") + return + } + seen := make(map[int64]struct{}, len(req.AccountIDs)) + accountIDs := make([]int64, 0, len(req.AccountIDs)) + for _, accountID := range req.AccountIDs { + if accountID <= 0 { + response.BadRequest(c, "account_ids must contain positive IDs") + return + } + if _, exists := seen[accountID]; exists { + continue + } + seen[accountID] = struct{}{} + accountIDs = append(accountIDs, accountID) + } + response.Success(c, gin.H{"results": h.upstreamBillingProbe.ProbeAccounts(c.Request.Context(), accountIDs)}) +} diff --git a/backend/internal/handler/admin/account_upstream_billing_probe_test.go b/backend/internal/handler/admin/account_upstream_billing_probe_test.go new file mode 100644 index 0000000000..cb181ebc52 --- /dev/null +++ b/backend/internal/handler/admin/account_upstream_billing_probe_test.go @@ -0,0 +1,71 @@ +package admin + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func setupUpstreamBillingProbeRouter() *gin.Engine { + gin.SetMode(gin.TestMode) + handler := NewAccountHandler(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + handler.SetUpstreamBillingProbeService(service.NewUpstreamBillingProbeService(nil, nil, nil)) + + router := gin.New() + router.GET("/admin/accounts/upstream-billing-probe/settings", handler.GetUpstreamBillingProbeSettings) + router.POST("/admin/accounts/upstream-billing-probe/batch", handler.ProbeUpstreamBillingBatch) + router.PUT("/admin/accounts/:id/upstream-billing-probe", handler.SetUpstreamBillingProbeEnabled) + return router +} + +func TestAccountHandlerGetUpstreamBillingProbeSettingsReturnsDefaults(t *testing.T) { + router := setupUpstreamBillingProbeRouter() + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/admin/accounts/upstream-billing-probe/settings", nil)) + + require.Equal(t, http.StatusOK, recorder.Code) + var response struct { + Data service.UpstreamBillingProbeSettings `json:"data"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + require.True(t, response.Data.Enabled) + require.Equal(t, 30, response.Data.IntervalMinutes) +} + +func TestAccountHandlerProbeUpstreamBillingBatchValidatesIDs(t *testing.T) { + router := setupUpstreamBillingProbeRouter() + + for _, body := range []string{`{"account_ids":[]}`, `{"account_ids":[0]}`} { + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/admin/accounts/upstream-billing-probe/batch", bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, request) + require.Equal(t, http.StatusBadRequest, recorder.Code) + } +} + +func TestAccountHandlerSetUpstreamBillingProbeEnabledRejectsInvalidID(t *testing.T) { + router := setupUpstreamBillingProbeRouter() + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPut, "/admin/accounts/not-an-id/upstream-billing-probe", bytes.NewBufferString(`{"enabled":true}`)) + request.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) +} + +func TestAccountHandlerSetUpstreamBillingProbeEnabledRequiresValue(t *testing.T) { + router := setupUpstreamBillingProbeRouter() + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPut, "/admin/accounts/1/upstream-billing-probe", bytes.NewBufferString(`{}`)) + request.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) +} diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 67380fed33..82004c5119 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -41,7 +41,9 @@ func ProvideAdminHandlers( paymentHandler *admin.PaymentHandler, affiliateHandler *admin.AffiliateHandler, complianceHandler *admin.ComplianceHandler, + upstreamBillingProbe *service.UpstreamBillingProbeService, ) *AdminHandlers { + accountHandler.SetUpstreamBillingProbeService(upstreamBillingProbe) return &AdminHandlers{ Dashboard: dashboardHandler, User: userHandler, diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 04beff8c62..f014155bb2 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -57,6 +57,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{ "codex_5h_", "codex_7d_", "passive_usage_", + "upstream_billing_probe", } var schedulerNeutralExtraKeys = map[string]struct{}{ @@ -395,21 +396,80 @@ func (r *accountRepository) ListCRSAccountIDs(ctx context.Context) (map[string]i } func (r *accountRepository) Update(ctx context.Context, account *service.Account) error { + return r.updateAccount(ctx, account, nil) +} + +// UpdateWithUpstreamBillingProbeEnabled applies an explicit probe switch in the +// same row-lock transaction as the rest of an admin account edit. +func (r *accountRepository) UpdateWithUpstreamBillingProbeEnabled(ctx context.Context, account *service.Account, enabled bool) error { + return r.updateAccount(ctx, account, &enabled) +} + +func (r *accountRepository) updateAccount(ctx context.Context, account *service.Account, explicitProbeEnabled *bool) error { if account == nil { return nil } + + baseCtx := ctx + contextTx := dbent.TxFromContext(ctx) + client := r.client + var tx *dbent.Tx + if contextTx != nil { + client = contextTx.Client() + } else { + var err error + tx, err = r.client.Tx(ctx) + if err != nil && !errors.Is(err, dbent.ErrTxStarted) { + return err + } + if tx != nil { + defer func() { _ = tx.Rollback() }() + ctx = dbent.NewTxContext(ctx, tx) + client = tx.Client() + } + } + + updated, err := r.updateLockedAccount(ctx, client, account, explicitProbeEnabled) + if err != nil { + return translatePersistenceError(err, service.ErrAccountNotFound, nil) + } + if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, buildSchedulerGroupPayload(account.GroupIDs)); err != nil { + return err + } + if tx != nil { + if err := tx.Commit(); err != nil { + return err + } + } + + account.UpdatedAt = updated.UpdatedAt + // 普通账号编辑(如 model_mapping / credentials)也需要立即刷新单账号快照, + // 否则网关在 outbox worker 延迟或异常时仍可能读到旧配置。 + if contextTx == nil { + r.syncSchedulerAccountSnapshot(baseCtx, account.ID) + } + return nil +} + +func (r *accountRepository) updateLockedAccount(ctx context.Context, client *dbent.Client, account *service.Account, explicitProbeEnabled *bool) (*dbent.Account, error) { + extra, err := lockAndMergeAccountProbeExtra(ctx, client, account, explicitProbeEnabled) + if err != nil { + return nil, err + } + account.Extra = extra + schedulable := account.Schedulable if account.Status == service.StatusError { schedulable = false } - builder := r.client.Account.UpdateOneID(account.ID). + builder := client.Account.UpdateOneID(account.ID). SetName(account.Name). SetNillableNotes(account.Notes). SetPlatform(account.Platform). SetType(account.Type). SetCredentials(normalizeJSONMap(account.Credentials)). - SetExtra(normalizeJSONMap(account.Extra)). + SetExtra(extra). SetConcurrency(account.Concurrency). SetPriority(account.Priority). SetStatus(account.Status). @@ -478,31 +538,140 @@ func (r *accountRepository) Update(ctx context.Context, account *service.Account builder.SetQuotaDimension(dbaccount.QuotaDimension(account.QuotaDimensionOrDefault())) builder.SetNillableParentAccountID(account.ParentAccountID) - updated, err := builder.Save(ctx) + return builder.Save(ctx) +} + +func lockAndMergeAccountProbeExtra(ctx context.Context, client *dbent.Client, account *service.Account, explicitProbeEnabled *bool) (map[string]any, error) { + credentials, err := json.Marshal(normalizeJSONMap(account.Credentials)) if err != nil { - return translatePersistenceError(err, service.ErrAccountNotFound, nil) + return nil, err } - account.UpdatedAt = updated.UpdatedAt - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, buildSchedulerGroupPayload(account.GroupIDs)); err != nil { - logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue account update failed: account=%d err=%v", account.ID, err) + var proxyID any + if account.ProxyID != nil { + proxyID = *account.ProxyID } - // 普通账号编辑(如 model_mapping / credentials)也需要立即刷新单账号快照, - // 否则网关在 outbox worker 延迟或异常时仍可能读到旧配置。 - r.syncSchedulerAccountSnapshot(ctx, account.ID) - return nil + rows, err := client.QueryContext(ctx, ` + SELECT + platform = $2 + AND type = $3 + AND credentials = $4::jsonb + AND proxy_id IS NOT DISTINCT FROM $5, + extra -> 'upstream_billing_probe_enabled', + extra -> 'upstream_billing_probe' + FROM accounts + WHERE id = $1 AND deleted_at IS NULL + FOR NO KEY UPDATE + `, account.ID, account.Platform, account.Type, string(credentials), proxyID) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + if !rows.Next() { + if err := rows.Err(); err != nil { + return nil, err + } + return nil, service.ErrAccountNotFound + } + + var ( + identityUnchanged bool + currentEnabled []byte + currentSnapshot []byte + ) + if err := rows.Scan(&identityUnchanged, ¤tEnabled, ¤tSnapshot); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + + extra := copyJSONMap(normalizeJSONMap(account.Extra)) + delete(extra, service.UpstreamBillingProbeEnabledExtraKey) + delete(extra, service.UpstreamBillingProbeExtraKey) + probeExplicitlyDisabled := false + probeAccount := account.Platform == service.PlatformOpenAI && account.Type == service.AccountTypeAPIKey + if probeAccount && explicitProbeEnabled != nil { + extra[service.UpstreamBillingProbeEnabledExtraKey] = *explicitProbeEnabled + probeExplicitlyDisabled = !*explicitProbeEnabled + } else if probeAccount && len(currentEnabled) > 0 && string(currentEnabled) != "null" { + var enabled any + if err := json.Unmarshal(currentEnabled, &enabled); err != nil { + return nil, err + } + extra[service.UpstreamBillingProbeEnabledExtraKey] = enabled + if value, ok := enabled.(bool); ok && !value { + probeExplicitlyDisabled = true + } + } + if !identityUnchanged || probeExplicitlyDisabled || len(currentSnapshot) == 0 || string(currentSnapshot) == "null" { + return extra, nil + } + var snapshot any + if err := json.Unmarshal(currentSnapshot, &snapshot); err != nil { + return nil, err + } + extra[service.UpstreamBillingProbeExtraKey] = snapshot + return extra, nil } func (r *accountRepository) UpdateCredentials(ctx context.Context, id int64, credentials map[string]any) error { - _, err := r.client.Account.UpdateOneID(id). - SetCredentials(normalizeJSONMap(credentials)). - Save(ctx) + payload, err := json.Marshal(normalizeJSONMap(credentials)) if err != nil { - return translatePersistenceError(err, service.ErrAccountNotFound, nil) + return err } - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { - logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue credentials update failed: account=%d err=%v", id, err) + baseCtx := ctx + contextTx := dbent.TxFromContext(ctx) + client := r.client + var tx *dbent.Tx + if contextTx != nil { + client = contextTx.Client() + } else if r.client != nil { + var txErr error + tx, txErr = r.client.Tx(ctx) + if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) { + return txErr + } + if tx != nil { + defer func() { _ = tx.Rollback() }() + ctx = dbent.NewTxContext(ctx, tx) + client = tx.Client() + } + } + result, err := client.ExecContext(ctx, ` + UPDATE accounts + SET + credentials = $1::jsonb, + extra = CASE + WHEN platform = 'openai' + AND type = 'apikey' + AND credentials IS DISTINCT FROM $1::jsonb + THEN COALESCE(extra, '{}'::jsonb) - 'upstream_billing_probe' + ELSE extra + END, + updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL + `, string(payload), id) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrAccountNotFound + } + if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { + return err + } + if tx != nil { + if err := tx.Commit(); err != nil { + return err + } + } + if contextTx == nil { + r.syncSchedulerAccountSnapshot(baseCtx, id) } - r.syncSchedulerAccountSnapshot(ctx, id) return nil } @@ -2117,10 +2286,31 @@ func (r *accountRepository) UpdateExtra(ctx context.Context, id int64, updates m return err } + clearProbeSnapshot := upstreamBillingProbeExplicitlyDisabled(updates) || upstreamBillingProbeSnapshotClearRequested(updates) + durableSchedulerChange := shouldEnqueueSchedulerOutboxForExtraUpdates(updates) || clearProbeSnapshot + baseCtx := ctx + contextTx := dbent.TxFromContext(ctx) client := clientFromContext(ctx, r.client) + var tx *dbent.Tx + if durableSchedulerChange && contextTx == nil { + var txErr error + tx, txErr = r.client.Tx(ctx) + if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) { + return txErr + } + if tx != nil { + defer func() { _ = tx.Rollback() }() + ctx = dbent.NewTxContext(ctx, tx) + client = tx.Client() + } + } + extraExpression := "COALESCE(extra, '{}'::jsonb) || $1::jsonb" + if clearProbeSnapshot { + extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'" + } result, err := client.ExecContext( ctx, - "UPDATE accounts SET extra = COALESCE(extra, '{}'::jsonb) || $1::jsonb, updated_at = NOW() WHERE id = $2 AND deleted_at IS NULL", + "UPDATE accounts SET extra = "+extraExpression+", updated_at = NOW() WHERE id = $2 AND deleted_at IS NULL", string(payload), id, ) @@ -2135,19 +2325,159 @@ func (r *accountRepository) UpdateExtra(ctx context.Context, id int64, updates m if affected == 0 { return service.ErrAccountNotFound } - if shouldEnqueueSchedulerOutboxForExtraUpdates(updates) { - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { - logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue extra update failed: account=%d err=%v", id, err) + if durableSchedulerChange { + if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { + return err + } + if tx != nil { + if err := tx.Commit(); err != nil { + return err + } + } + if contextTx == nil { + r.syncSchedulerAccountSnapshot(baseCtx, id) } } else { // 观测型 extra 字段不需要触发 bucket 重建,但仍同步单账号快照, // 让 sticky session / GetAccount 命中缓存时也能读到最新数据, // 同时避免缓存局部 patch 覆盖掉并发写入的其它账号字段。 - r.syncSchedulerAccountSnapshot(ctx, id) + if dbent.TxFromContext(ctx) == nil { + r.syncSchedulerAccountSnapshot(ctx, id) + } } return nil } +// UpdateUpstreamBillingProbeSnapshot stores a probe result only while the +// network identity used by that probe is still current. +func (r *accountRepository) UpdateUpstreamBillingProbeSnapshot( + ctx context.Context, + account *service.Account, + snapshot *service.UpstreamBillingProbeSnapshot, +) error { + if account == nil || snapshot == nil { + return service.ErrAccountNilInput + } + if dbent.TxFromContext(ctx) == nil { + tx, err := r.client.Tx(ctx) + if errors.Is(err, dbent.ErrTxStarted) { + return r.updateUpstreamBillingProbeSnapshotInTx(ctx, account, snapshot) + } + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + + if err := r.updateUpstreamBillingProbeSnapshotInTx(dbent.NewTxContext(ctx, tx), account, snapshot); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return err + } + // The durable outbox event is committed with the snapshot. This direct + // cache write only reduces visibility latency on the current instance. + r.syncSchedulerAccountSnapshot(ctx, account.ID) + return nil + } + return r.updateUpstreamBillingProbeSnapshotInTx(ctx, account, snapshot) +} + +func (r *accountRepository) updateUpstreamBillingProbeSnapshotInTx( + ctx context.Context, + account *service.Account, + snapshot *service.UpstreamBillingProbeSnapshot, +) error { + payload, err := json.Marshal(map[string]any{service.UpstreamBillingProbeExtraKey: snapshot}) + if err != nil { + return err + } + credentials, err := json.Marshal(account.Credentials) + if err != nil { + return err + } + var expectedSnapshot any + if account.Extra != nil { + expectedSnapshot = account.Extra[service.UpstreamBillingProbeExtraKey] + } + expectedSnapshotJSON, err := json.Marshal(expectedSnapshot) + if err != nil { + return err + } + var expectedEnabled any + if account.Extra != nil { + expectedEnabled = account.Extra[service.UpstreamBillingProbeEnabledExtraKey] + } + expectedEnabledJSON, err := json.Marshal(expectedEnabled) + if err != nil { + return err + } + client := clientFromContext(ctx, r.client) + proxyMatches, err := lockAndMatchProbeProxyIdentity(ctx, client, account) + if err != nil { + return err + } + if !proxyMatches { + return service.ErrUpstreamBillingProbeIdentityChanged + } + var proxyID any + if account.ProxyID != nil { + proxyID = *account.ProxyID + } + result, err := client.ExecContext(ctx, ` + UPDATE accounts + SET extra = COALESCE(extra, '{}'::jsonb) || $1::jsonb, updated_at = NOW() + WHERE id = $2 + AND platform = $3 + AND type = $4 + AND credentials = $5::jsonb + AND proxy_id IS NOT DISTINCT FROM $6 + AND COALESCE(extra -> 'upstream_billing_probe', 'null'::jsonb) = $7::jsonb + AND COALESCE(extra -> 'upstream_billing_probe_enabled', 'null'::jsonb) = $8::jsonb + AND deleted_at IS NULL + `, string(payload), account.ID, account.Platform, account.Type, string(credentials), proxyID, string(expectedSnapshotJSON), string(expectedEnabledJSON)) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + return service.ErrUpstreamBillingProbeIdentityChanged + } + return enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, nil) +} + +func lockAndMatchProbeProxyIdentity(ctx context.Context, client *dbent.Client, account *service.Account) (bool, error) { + if account.ProxyID == nil { + return true, nil + } + rows, err := client.QueryContext(ctx, ` + SELECT protocol, host, port, COALESCE(username, ''), COALESCE(password, ''), status + FROM proxies + WHERE id = $1 AND deleted_at IS NULL + FOR SHARE + `, *account.ProxyID) + if err != nil { + return false, err + } + defer func() { _ = rows.Close() }() + if !rows.Next() { + if err := rows.Err(); err != nil { + return false, err + } + return account.Proxy == nil, nil + } + if account.Proxy == nil || account.Proxy.ID != *account.ProxyID { + return false, nil + } + var current proxyProbeIdentity + if err := rows.Scan(¤t.protocol, ¤t.host, ¤t.port, ¤t.username, ¤t.password, ¤t.status); err != nil { + return false, err + } + return current == proxyProbeIdentityFromService(account.Proxy), rows.Err() +} + func shouldEnqueueSchedulerOutboxForExtraUpdates(updates map[string]any) bool { if len(updates) == 0 { return false @@ -2177,6 +2507,16 @@ func isSchedulerNeutralExtraKey(key string) bool { return false } +func upstreamBillingProbeExplicitlyDisabled(extra map[string]any) bool { + enabled, ok := extra[service.UpstreamBillingProbeEnabledExtraKey].(bool) + return ok && !enabled +} + +func upstreamBillingProbeSnapshotClearRequested(extra map[string]any) bool { + value, ok := extra[service.UpstreamBillingProbeExtraKey] + return ok && value == nil +} + func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates service.AccountBulkUpdate) (int64, error) { if len(ids) == 0 { return 0, nil @@ -2250,7 +2590,11 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates if err != nil { return 0, err } - setClauses = append(setClauses, "extra = COALESCE(extra, '{}'::jsonb) || $"+itoa(idx)+"::jsonb") + extraExpression := "COALESCE(extra, '{}'::jsonb) || $" + itoa(idx) + "::jsonb" + if upstreamBillingProbeExplicitlyDisabled(updates.Extra) || upstreamBillingProbeSnapshotClearRequested(updates.Extra) { + extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'" + } + setClauses = append(setClauses, "extra = "+extraExpression) args = append(args, payload) idx++ } @@ -2264,7 +2608,26 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates query := "UPDATE accounts SET " + joinClauses(setClauses, ", ") + " WHERE id = ANY($" + itoa(idx) + ") AND deleted_at IS NULL" args = append(args, pq.Array(ids)) - result, err := r.sql.ExecContext(ctx, query, args...) + baseCtx := ctx + contextTx := dbent.TxFromContext(ctx) + exec := r.sql + var tx *dbent.Tx + if contextTx != nil { + exec = contextTx.Client() + } else if r.client != nil { + var txErr error + tx, txErr = r.client.Tx(ctx) + if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) { + return 0, txErr + } + if tx != nil { + defer func() { _ = tx.Rollback() }() + ctx = dbent.NewTxContext(ctx, tx) + exec = tx.Client() + } + } + + result, err := exec.ExecContext(ctx, query, args...) if err != nil { return 0, err } @@ -2274,9 +2637,16 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates } if rows > 0 { payload := map[string]any{"account_ids": ids} - if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil { - logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue bulk update failed: err=%v", err) + if err := enqueueSchedulerOutbox(ctx, exec, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil { + return 0, err } + } + if tx != nil { + if err := tx.Commit(); err != nil { + return 0, err + } + } + if rows > 0 && contextTx == nil { shouldSync := false if updates.Status != nil && (*updates.Status == service.StatusError || *updates.Status == service.StatusDisabled) { shouldSync = true @@ -2285,7 +2655,7 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates shouldSync = true } if shouldSync { - r.syncSchedulerAccountSnapshots(ctx, ids) + r.syncSchedulerAccountSnapshots(baseCtx, ids) } } return rows, nil @@ -2733,6 +3103,107 @@ func (r *accountRepository) FindByExtraField(ctx context.Context, key string, va return r.accountsToService(ctx, accounts) } +// ListDueUpstreamBillingProbeAccounts bounds result hydration and network work +// to limit. PostgreSQL must still filter and order all enabled candidates; +// MATERIALIZED avoids repeating the defensive timestamp parse expression. +func (r *accountRepository) ListDueUpstreamBillingProbeAccounts(ctx context.Context, now time.Time, limit int) ([]service.Account, error) { + if limit <= 0 { + return []service.Account{}, nil + } + if r.sql == nil { + return nil, errors.New("account repository SQL executor not configured") + } + + rows, err := r.sql.QueryContext(ctx, ` + WITH candidates AS ( + SELECT + id, + extra #>> '{upstream_billing_probe,status}' AS probe_status, + extra #>> '{upstream_billing_probe,next_probe_at}' AS next_probe_at + FROM accounts + WHERE deleted_at IS NULL + AND status = 'active' + AND platform = 'openai' + AND type = 'apikey' + AND extra @> '{"upstream_billing_probe_enabled": true}'::jsonb + ), parsed AS MATERIALIZED ( + SELECT + id, + probe_status, + next_probe_at, + next_probe_at ~ '^[0-9]{4}-[0-9]{2}-[0-9]{2}T[0-9]{2}:[0-9]{2}:[0-9]{2}(\.[0-9]+)?(Z|[+-][0-9]{2}:[0-9]{2})$' AS rfc3339_shape, + jsonb_path_query_first_tz( + jsonb_build_object( + 'value', + replace(regexp_replace(next_probe_at, 'Z$', '+00:00'), 'T', ' ') + ), + '$.value.datetime()', + '{}'::jsonb, + true + ) #>> '{}' AS parsed_next_probe_at + FROM candidates + ), normalized AS ( + SELECT + id, + probe_status, + next_probe_at, + parsed_next_probe_at, + rfc3339_shape AND parsed_next_probe_at IS NOT NULL AS valid_next_probe_at + FROM parsed + ) + SELECT id + FROM normalized + WHERE probe_status NOT IN ('ok', 'unsupported', 'failed') + OR probe_status IS NULL + OR next_probe_at IS NULL + OR NOT valid_next_probe_at + OR CASE WHEN valid_next_probe_at THEN parsed_next_probe_at::timestamptz <= $1 ELSE FALSE END + ORDER BY + CASE + WHEN probe_status NOT IN ('ok', 'unsupported', 'failed') + OR probe_status IS NULL + OR next_probe_at IS NULL + OR NOT valid_next_probe_at + THEN 0 + ELSE 1 + END ASC, + CASE WHEN valid_next_probe_at THEN parsed_next_probe_at::timestamptz END ASC NULLS FIRST, + id ASC + LIMIT $2 + `, now.UTC(), limit) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + ids := make([]int64, 0, limit) + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + ids = append(ids, id) + } + if err := rows.Err(); err != nil { + return nil, err + } + if len(ids) == 0 { + return []service.Account{}, nil + } + + accounts, err := r.GetByIDs(ctx, ids) + if err != nil { + return nil, err + } + out := make([]service.Account, 0, len(accounts)) + for _, account := range accounts { + if account != nil { + out = append(out, *account) + } + } + return out, nil +} + // nowUTC is a SQL expression to generate a UTC RFC3339 timestamp string. const nowUTC = `to_char(NOW() AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.US"Z"')` diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_cas_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_cas_test.go new file mode 100644 index 0000000000..4de2187915 --- /dev/null +++ b/backend/internal/repository/account_repo_upstream_billing_probe_cas_test.go @@ -0,0 +1,186 @@ +package repository + +import ( + "context" + "errors" + "regexp" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +func TestUpdateUpstreamBillingProbeSnapshotRequiresSameIdentityAndSnapshot(t *testing.T) { + tests := []struct { + name string + affected int64 + wantErr error + }{ + {name: "same identity and snapshot", affected: 1}, + {name: "identity or snapshot changed", affected: 0, wantErr: service.ErrUpstreamBillingProbeIdentityChanged}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + tx, err := client.Tx(context.Background()) + require.NoError(t, err) + mock.ExpectQuery(`(?s)` + regexp.QuoteMeta("SELECT protocol, host, port") + `.*` + regexp.QuoteMeta("FOR SHARE")). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"protocol", "host", "port", "username", "password", "status"}). + AddRow("http", "127.0.0.1", 3128, "user", "pass", service.StatusActive)) + mock.ExpectExec(`(?s)`+regexp.QuoteMeta("UPDATE accounts")+`.*`+regexp.QuoteMeta("WHERE id = $2")+`.*`+regexp.QuoteMeta("AND platform = $3")+`.*`+regexp.QuoteMeta("AND type = $4")+`.*`+regexp.QuoteMeta("AND credentials = $5::jsonb")+`.*`+regexp.QuoteMeta("AND proxy_id IS NOT DISTINCT FROM $6")+`.*`+regexp.QuoteMeta("COALESCE(extra -> 'upstream_billing_probe', 'null'::jsonb) = $7::jsonb")+`.*`+regexp.QuoteMeta("COALESCE(extra -> 'upstream_billing_probe_enabled', 'null'::jsonb) = $8::jsonb")). + WithArgs(sqlmock.AnyArg(), int64(17), service.PlatformOpenAI, service.AccountTypeAPIKey, `{"api_key":"sk-test","base_url":"http://127.0.0.1:8080"}`, int64(9), `{"status":"stale"}`, "null"). + WillReturnResult(sqlmock.NewResult(0, tt.affected)) + if tt.affected > 0 { + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")). + WithArgs(service.SchedulerOutboxEventAccountChanged, int64(17), nil, nil, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + } + repo := newAccountRepositoryWithSQL(client, &recordingSQLExecutor{err: errors.New("must use transaction client")}, nil) + proxyID := int64(9) + account := &service.Account{ + ID: 17, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "http://127.0.0.1:8080", + }, + ProxyID: &proxyID, + Proxy: &service.Proxy{ + ID: proxyID, + Protocol: "http", + Host: "127.0.0.1", + Port: 3128, + Username: "user", + Password: "pass", + Status: service.StatusActive, + }, + Extra: map[string]any{ + service.UpstreamBillingProbeExtraKey: map[string]any{"status": "stale"}, + }, + } + + txCtx := dbent.NewTxContext(context.Background(), tx) + err = repo.UpdateUpstreamBillingProbeSnapshot(txCtx, account, &service.UpstreamBillingProbeSnapshot{Status: service.UpstreamBillingProbeStatusOK}) + + if tt.wantErr != nil { + require.ErrorIs(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + mock.ExpectRollback() + require.NoError(t, tx.Rollback()) + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestUpdateUpstreamBillingProbeSnapshotCommitsSnapshotAndOutboxAtomically(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)`+regexp.QuoteMeta("UPDATE accounts")+`.*`+regexp.QuoteMeta("AND credentials = $5::jsonb")+`.*`+regexp.QuoteMeta("AND proxy_id IS NOT DISTINCT FROM $6")+`.*`+regexp.QuoteMeta("COALESCE(extra -> 'upstream_billing_probe', 'null'::jsonb) = $7::jsonb")). + WithArgs(sqlmock.AnyArg(), int64(17), service.PlatformOpenAI, service.AccountTypeAPIKey, `{"api_key":"sk-test"}`, nil, "null", "null"). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")). + WithArgs(service.SchedulerOutboxEventAccountChanged, int64(17), nil, nil, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + repo := newAccountRepositoryWithSQL(client, db, nil) + account := &service.Account{ + ID: 17, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + } + + err = repo.UpdateUpstreamBillingProbeSnapshot(context.Background(), account, &service.UpstreamBillingProbeSnapshot{Status: service.UpstreamBillingProbeStatusOK}) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateUpstreamBillingProbeSnapshotRejectsChangedProxyIdentity(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + tx, err := client.Tx(context.Background()) + require.NoError(t, err) + mock.ExpectQuery(`(?s)` + regexp.QuoteMeta("SELECT protocol, host, port") + `.*` + regexp.QuoteMeta("FOR SHARE")). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"protocol", "host", "port", "username", "password", "status"}). + AddRow("http", "new.example", 3128, "user", "pass", service.StatusActive)) + + proxyID := int64(9) + account := &service.Account{ + ID: 17, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + ProxyID: &proxyID, + Proxy: &service.Proxy{ + ID: proxyID, Protocol: "http", Host: "old.example", Port: 3128, + Username: "user", Password: "pass", Status: service.StatusActive, + }, + } + repo := newAccountRepositoryWithSQL(client, db, nil) + err = repo.UpdateUpstreamBillingProbeSnapshot(dbent.NewTxContext(context.Background(), tx), account, &service.UpstreamBillingProbeSnapshot{Status: service.UpstreamBillingProbeStatusOK}) + + require.ErrorIs(t, err, service.ErrUpstreamBillingProbeIdentityChanged) + mock.ExpectRollback() + require.NoError(t, tx.Rollback()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateUpstreamBillingProbeSnapshotRollsBackWhenOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)`+regexp.QuoteMeta("UPDATE accounts")+`.*`+regexp.QuoteMeta("AND proxy_id IS NOT DISTINCT FROM $6")+`.*`+regexp.QuoteMeta("COALESCE(extra -> 'upstream_billing_probe', 'null'::jsonb) = $7::jsonb")). + WithArgs(sqlmock.AnyArg(), int64(18), service.PlatformOpenAI, service.AccountTypeAPIKey, `{"api_key":"sk-test"}`, nil, "null", "null"). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")).WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newAccountRepositoryWithSQL(client, db, nil) + account := &service.Account{ + ID: 18, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + } + + err = repo.UpdateUpstreamBillingProbeSnapshot(context.Background(), account, &service.UpstreamBillingProbeSnapshot{Status: service.UpstreamBillingProbeStatusOK}) + + require.EqualError(t, err, "outbox failed") + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_due_integration_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_due_integration_test.go new file mode 100644 index 0000000000..64d29553fe --- /dev/null +++ b/backend/internal/repository/account_repo_upstream_billing_probe_due_integration_test.go @@ -0,0 +1,51 @@ +//go:build integration + +package repository + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestListDueUpstreamBillingProbeAccountsHandlesInvalidCalendarDate(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + now := time.Date(2026, time.July, 14, 12, 0, 0, 0, time.UTC) + _, err := tx.ExecContext(ctx, ` + UPDATE accounts + SET extra = extra - 'upstream_billing_probe_enabled' - 'upstream_billing_probe' + `) + require.NoError(t, err) + + insert := func(name, nextProbeAt string) int64 { + t.Helper() + var id int64 + extra := fmt.Sprintf(`{ + "upstream_billing_probe_enabled": true, + "upstream_billing_probe": {"status": "ok", "next_probe_at": %q} + }`, nextProbeAt) + err := scanSingleRow(ctx, tx, ` + INSERT INTO accounts (name, platform, type, status, extra) + VALUES ($1, 'openai', $2, 'active', $3::jsonb) + RETURNING id + `, []any{name, service.AccountTypeAPIKey, extra}, &id) + require.NoError(t, err) + return id + } + + invalidID := insert("probe-invalid-calendar-date", "2026-99-99T12:00:00Z") + dueID := insert("probe-due", "2026-07-14T11:59:59Z") + _ = insert("probe-not-due", "2026-07-14T12:00:01Z") + + accounts, err := repo.ListDueUpstreamBillingProbeAccounts(ctx, now, 20) + require.NoError(t, err) + require.Len(t, accounts, 2) + require.Equal(t, invalidID, accounts[0].ID) + require.Equal(t, dueID, accounts[1].ID) +} diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_due_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_due_test.go new file mode 100644 index 0000000000..350b95d4a2 --- /dev/null +++ b/backend/internal/repository/account_repo_upstream_billing_probe_due_test.go @@ -0,0 +1,48 @@ +package repository + +import ( + "context" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/require" +) + +func TestAccountRepositoryListDueUpstreamBillingProbeAccountsBoundsQuery(t *testing.T) { + db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp)) + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + + now := time.Date(2026, time.July, 14, 12, 0, 0, 0, time.UTC) + var capturedSQL string + mock.ExpectQuery("WITH candidates AS"). + WithArgs(now, 20). + WillReturnRows(sqlmock.NewRows([]string{"id"})) + repo := newAccountRepositoryWithSQL(nil, captureQuerySQL{db: db, captured: &capturedSQL}, nil) + + accounts, err := repo.ListDueUpstreamBillingProbeAccounts(context.Background(), now, 20) + + require.NoError(t, err) + require.Empty(t, accounts) + normalized := normalizeSQLWhitespace(capturedSQL) + require.Contains(t, normalized, "deleted_at IS NULL") + require.Contains(t, normalized, "status = 'active'") + require.Contains(t, normalized, "platform = 'openai'") + require.Contains(t, normalized, "type = 'apikey'") + require.Contains(t, normalized, `extra @> '{"upstream_billing_probe_enabled": true}'::jsonb`) + require.Contains(t, normalized, "jsonb_path_query_first_tz") + require.Contains(t, normalized, "parsed AS MATERIALIZED") + require.Contains(t, normalized, "parsed_next_probe_at::timestamptz <= $1") + require.Contains(t, normalized, "LIMIT $2") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestAccountRepositoryListDueUpstreamBillingProbeAccountsRejectsNonPositiveLimit(t *testing.T) { + repo := newAccountRepositoryWithSQL(nil, nil, nil) + + accounts, err := repo.ListDueUpstreamBillingProbeAccounts(context.Background(), time.Now(), 0) + + require.NoError(t, err) + require.Empty(t, accounts) +} diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_test.go new file mode 100644 index 0000000000..bcfa59ddd3 --- /dev/null +++ b/backend/internal/repository/account_repo_upstream_billing_probe_test.go @@ -0,0 +1,16 @@ +package repository + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestUpstreamBillingProbeExtraIsSchedulerNeutral(t *testing.T) { + require.True(t, isSchedulerNeutralExtraKey("upstream_billing_probe")) + require.True(t, isSchedulerNeutralExtraKey("upstream_billing_probe_enabled")) + require.False(t, shouldEnqueueSchedulerOutboxForExtraUpdates(map[string]any{ + "upstream_billing_probe": map[string]any{"status": "ok"}, + "upstream_billing_probe_enabled": true, + })) +} diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go new file mode 100644 index 0000000000..26fbb74dcf --- /dev/null +++ b/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go @@ -0,0 +1,306 @@ +package repository + +import ( + "context" + "errors" + "regexp" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + dbaccount "github.com/Wei-Shaw/sub2api/ent/account" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +func TestLockAndMergeAccountProbeExtraUsesCurrentDatabaseSnapshot(t *testing.T) { + tests := []struct { + name string + identityUnchanged bool + databaseEnabled any + databaseSnapshot any + inputExtra map[string]any + wantSnapshot any + wantEnabled any + }{ + { + name: "ordinary edit preserves current enable flag and snapshot created after account load", + identityUnchanged: true, + databaseEnabled: []byte(`true`), + databaseSnapshot: []byte(`{"status":"ok"}`), + inputExtra: map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false}, + wantSnapshot: map[string]any{"status": "ok"}, + wantEnabled: true, + }, + { + name: "identity change clears stale snapshot", + identityUnchanged: false, + databaseEnabled: []byte(`true`), + databaseSnapshot: []byte(`{"status":"ok"}`), + inputExtra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": "stale"}, + }, + wantEnabled: true, + }, + { + name: "current explicit disable clears snapshot", + identityUnchanged: true, + databaseEnabled: []byte(`false`), + databaseSnapshot: []byte(`{"status":"ok"}`), + inputExtra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": "stale"}, + }, + wantEnabled: false, + }, + { + name: "missing database snapshot is not resurrected from stale input", + identityUnchanged: true, + databaseEnabled: []byte(`true`), + databaseSnapshot: nil, + inputExtra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": "stale"}, + }, + wantEnabled: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectQuery(`(?s)`+regexp.QuoteMeta("SELECT")+`.*`+regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(int64(27), service.PlatformOpenAI, service.AccountTypeAPIKey, `{"api_key":"sk-test"}`, nil). + WillReturnRows(sqlmock.NewRows([]string{"identity_unchanged", "enabled", "snapshot"}). + AddRow(tt.identityUnchanged, tt.databaseEnabled, tt.databaseSnapshot)) + + account := &service.Account{ + ID: 27, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: tt.inputExtra, + } + got, err := lockAndMergeAccountProbeExtra(context.Background(), client, account, nil) + require.NoError(t, err) + if tt.wantSnapshot == nil { + require.NotContains(t, got, service.UpstreamBillingProbeExtraKey) + } else { + require.Equal(t, tt.wantSnapshot, got[service.UpstreamBillingProbeExtraKey]) + } + require.Equal(t, tt.wantEnabled, got[service.UpstreamBillingProbeEnabledExtraKey]) + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestUpdateExtraExplicitProbeDisableRemovesSnapshot(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts SET extra = .* - 'upstream_billing_probe'`). + WithArgs(`{"upstream_billing_probe_enabled":false}`, int64(27)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")). + WithArgs(service.SchedulerOutboxEventAccountChanged, int64(27), nil, nil, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + repo := newAccountRepositoryWithSQL(client, db, nil) + + err = repo.UpdateExtra(context.Background(), 27, map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false}) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateExtraNilProbeRemovesKeyInsteadOfWritingJSONNull(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts SET extra = .* - 'upstream_billing_probe'`). + WithArgs(`{"upstream_billing_probe":null}`, int64(27)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")). + WithArgs(service.SchedulerOutboxEventAccountChanged, int64(27), nil, nil, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + repo := newAccountRepositoryWithSQL(client, db, nil) + + err = repo.UpdateExtra(context.Background(), 27, map[string]any{service.UpstreamBillingProbeExtraKey: nil}) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestBulkUpdateNilProbeRemovesKeyInsteadOfWritingJSONNull(t *testing.T) { + exec := &recordingSQLExecutor{result: rowsAffectedResult(1)} + repo := newAccountRepositoryWithSQL(nil, exec, nil) + + _, err := repo.BulkUpdate(context.Background(), []int64{27}, service.AccountBulkUpdate{ + Extra: map[string]any{service.UpstreamBillingProbeExtraKey: nil}, + }) + + require.NoError(t, err) + require.NotEmpty(t, exec.execQueries) + require.Contains(t, normalizeSQLWhitespace(exec.execQueries[0]), "- 'upstream_billing_probe'") +} + +func TestUpdateCredentialsAtomicallyClearsProbeForOpenAIAPIKeyIdentityChange(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts.*credentials IS DISTINCT FROM \$1::jsonb.*- 'upstream_billing_probe'`). + WithArgs(`{"api_key":"sk-new"}`, int64(27)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")). + WithArgs(service.SchedulerOutboxEventAccountChanged, int64(27), nil, nil, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + repo := newAccountRepositoryWithSQL(client, db, nil) + + err = repo.UpdateCredentials(context.Background(), 27, map[string]any{"api_key": "sk-new"}) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateWithUpstreamBillingProbeEnabledRollsBackWhenOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectQuery(`(?s)`+regexp.QuoteMeta("SELECT")+`.*`+regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(int64(27), service.PlatformOpenAI, service.AccountTypeAPIKey, `{"api_key":"sk-test"}`, nil). + WillReturnRows(sqlmock.NewRows([]string{"identity_unchanged", "enabled", "snapshot"}). + AddRow(true, []byte(`true`), []byte(`{"status":"ok"}`))) + mock.ExpectExec(`(?s)UPDATE .*accounts.*SET.*WHERE .*id.*`). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectQuery(`(?s)SELECT .* FROM "accounts" WHERE "id" = \$1`). + WithArgs(int64(27)). + WillReturnRows(updatedAccountRows(27, `{"upstream_billing_probe_enabled":false}`)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")).WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newAccountRepositoryWithSQL(client, db, nil) + account := &service.Account{ + ID: 27, + Name: "test", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{ + service.UpstreamBillingProbeExtraKey: map[string]any{"status": "stale"}, + }, + Concurrency: 1, + Priority: 1, + Status: service.StatusActive, + Schedulable: true, + } + + err = repo.UpdateWithUpstreamBillingProbeEnabled(context.Background(), account, false) + + require.EqualError(t, err, "outbox failed") + require.Equal(t, false, account.Extra[service.UpstreamBillingProbeEnabledExtraKey]) + require.NotContains(t, account.Extra, service.UpstreamBillingProbeExtraKey) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateExtraRollsBackWhenOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts SET extra = .* - 'upstream_billing_probe'`). + WithArgs(`{"upstream_billing_probe_enabled":false}`, int64(27)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")).WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newAccountRepositoryWithSQL(client, db, nil) + err = repo.UpdateExtra(context.Background(), 27, map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false}) + + require.EqualError(t, err, "outbox failed") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestUpdateCredentialsRollsBackWhenOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts.*credentials IS DISTINCT FROM \$1::jsonb.*- 'upstream_billing_probe'`). + WithArgs(`{"api_key":"sk-new"}`, int64(27)). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")).WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newAccountRepositoryWithSQL(client, db, nil) + err = repo.UpdateCredentials(context.Background(), 27, map[string]any{"api_key": "sk-new"}) + + require.EqualError(t, err, "outbox failed") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestBulkUpdateRollsBackWhenOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + name := "renamed" + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts SET name = \$1.*WHERE id = ANY\(\$2\)`). + WithArgs(name, `{27,28}`). + WillReturnResult(sqlmock.NewResult(0, 2)) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox")).WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newAccountRepositoryWithSQL(client, db, nil) + rows, err := repo.BulkUpdate(context.Background(), []int64{27, 28}, service.AccountBulkUpdate{Name: &name}) + + require.EqualError(t, err, "outbox failed") + require.Zero(t, rows) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func updatedAccountRows(id int64, extra string) *sqlmock.Rows { + now := time.Now() + return sqlmock.NewRows(dbaccount.Columns).AddRow( + id, now, now, nil, "test", nil, service.PlatformOpenAI, service.AccountTypeAPIKey, + []byte(`{"api_key":"sk-test"}`), []byte(extra), nil, nil, 1, nil, 1, 1.0, + service.StatusActive, nil, nil, nil, false, true, nil, nil, nil, nil, nil, nil, + nil, nil, nil, service.QuotaDimensionGlobal, + ) +} diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index 58b1d345d4..06113efc6c 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -195,7 +195,7 @@ func (s *httpUpstreamService) Do(req *http.Request, proxyURL string, accountID i } // 执行请求 - resp, err := servertiming.Do(entry.client, req) + resp, err := servertiming.Do(httpClientForUpstreamRequest(entry.client, req), req) if err != nil { s.recordOpenAIHTTP2Failure(profile, entry.protocolMode, entry.proxyKey, err) // 请求失败,立即减少计数 @@ -226,6 +226,11 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco if profile == nil { return s.Do(req, proxyURL, accountID, accountConcurrency) } + // Plain HTTP has no TLS handshake to fingerprint. Reuse the normal transport + // so a configured HTTP or SOCKS proxy is not bypassed. + if req != nil && req.URL != nil && strings.EqualFold(req.URL.Scheme, "http") { + return s.Do(req, proxyURL, accountID, accountConcurrency) + } applyGrokCLIProxyHeaders(req) upstreamProfile := service.HTTPUpstreamProfileDefault if req != nil { @@ -252,7 +257,7 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco return nil, err } - resp, err := servertiming.Do(entry.client, req) + resp, err := servertiming.Do(httpClientForUpstreamRequest(entry.client, req), req) if err != nil { atomic.AddInt64(&entry.inFlight, -1) atomic.StoreInt64(&entry.lastUsed, time.Now().UnixNano()) @@ -270,6 +275,17 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco return resp, nil } +func httpClientForUpstreamRequest(client *http.Client, req *http.Request) *http.Client { + if client == nil || req == nil || !service.HTTPUpstreamRedirectsDisabled(req.Context()) { + return client + } + clone := *client + clone.CheckRedirect = func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + } + return &clone +} + // applyGrokCLIProxyHeaders applies the official Grok Build client identity at // the final shared transport boundary. Keying this behavior to the exact CLI // proxy host keeps direct api.x.ai traffic unchanged and automatically covers @@ -1186,7 +1202,11 @@ func buildUpstreamTransportWithTLSFingerprint(settings poolSettings, proxyURL *u slog.Debug("tls_fingerprint_transport_socks5", "proxy", proxyURL.Host) socks5Dialer := tlsfingerprint.NewSOCKS5ProxyDialer(profile, proxyURL) transport.DialTLSContext = socks5Dialer.DialTLSContext - case "http", "https": + case "https": + // The fingerprint dialer emits a plaintext CONNECT preface and cannot + // establish TLS to an HTTPS proxy. Keep proxy routing via net/http. + return buildUpstreamTransport(settings, proxyURL, upstreamProtocolModeDefault) + case "http": // HTTP/HTTPS 代理:使用 HTTPProxyDialer(CONNECT 隧道) slog.Debug("tls_fingerprint_transport_http_connect", "proxy", proxyURL.Host) httpDialer := tlsfingerprint.NewHTTPProxyDialer(profile, proxyURL) diff --git a/backend/internal/repository/http_upstream_test.go b/backend/internal/repository/http_upstream_test.go index cb8fe70d5b..151e227146 100644 --- a/backend/internal/repository/http_upstream_test.go +++ b/backend/internal/repository/http_upstream_test.go @@ -1,9 +1,14 @@ package repository import ( + "encoding/binary" "errors" + "fmt" "io" + "net" "net/http" + "net/http/httptest" + "net/url" "sync/atomic" "testing" "time" @@ -15,6 +20,172 @@ import ( "github.com/stretchr/testify/suite" ) +func TestHTTPUpstreamDoCanDisableRedirectsPerRequest(t *testing.T) { + var redirectedCalls atomic.Int64 + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + redirectedCalls.Add(1) + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(target.Close) + redirector := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, target.URL, http.StatusFound) + })) + t.Cleanup(redirector.Close) + + upstream := NewHTTPUpstream(nil) + req, err := http.NewRequestWithContext( + service.WithHTTPUpstreamRedirectsDisabled(t.Context()), + http.MethodGet, + redirector.URL, + nil, + ) + require.NoError(t, err) + + resp, err := upstream.Do(req, "", 1, 1) + require.NoError(t, err) + require.Equal(t, http.StatusFound, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + require.Zero(t, redirectedCalls.Load()) +} + +func TestHTTPUpstreamDoWithTLSPlainHTTPUsesConfiguredHTTPProxy(t *testing.T) { + var upstreamCalls atomic.Int64 + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + upstreamCalls.Add(1) + w.WriteHeader(http.StatusTeapot) + })) + t.Cleanup(upstream.Close) + var proxyCalls atomic.Int64 + proxy := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + proxyCalls.Add(1) + w.WriteHeader(http.StatusNoContent) + })) + t.Cleanup(proxy.Close) + + req, err := http.NewRequest(http.MethodGet, upstream.URL, nil) + require.NoError(t, err) + client := NewHTTPUpstream(nil) + resp, err := client.DoWithTLS(req, proxy.URL, 41, 1, &tlsfingerprint.Profile{Name: "unused-for-http"}) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + require.Equal(t, int64(1), proxyCalls.Load()) + require.Zero(t, upstreamCalls.Load(), "plain HTTP must not bypass the configured proxy") +} + +func TestHTTPUpstreamDoWithTLSPlainHTTPUsesConfiguredSOCKSProxy(t *testing.T) { + var upstreamCalls atomic.Int64 + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + upstreamCalls.Add(1) + w.WriteHeader(http.StatusNoContent) + })) + t.Cleanup(upstream.Close) + proxyURL, proxyCalls := startTestSOCKS5Proxy(t) + + req, err := http.NewRequest(http.MethodGet, upstream.URL, nil) + require.NoError(t, err) + client := NewHTTPUpstream(nil) + resp, err := client.DoWithTLS(req, proxyURL, 42, 1, &tlsfingerprint.Profile{Name: "unused-for-http"}) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, resp.StatusCode) + require.NoError(t, resp.Body.Close()) + require.Equal(t, int64(1), proxyCalls.Load()) + require.Equal(t, int64(1), upstreamCalls.Load()) +} + +func TestTLSFingerprintHTTPSProxyFallsBackWithoutBypassingProxy(t *testing.T) { + proxyURL, err := url.Parse("https://user:pass@proxy.example:8443") + require.NoError(t, err) + transport, err := buildUpstreamTransportWithTLSFingerprint(poolSettings{}, proxyURL, &tlsfingerprint.Profile{Name: "test"}) + require.NoError(t, err) + require.NotNil(t, transport.Proxy) + require.Nil(t, transport.DialTLSContext) + req := &http.Request{URL: &url.URL{Scheme: "https", Host: "upstream.example"}} + resolved, err := transport.Proxy(req) + require.NoError(t, err) + require.Equal(t, "https://user:pass@proxy.example:8443", resolved.String()) +} + +func startTestSOCKS5Proxy(t *testing.T) (string, *atomic.Int64) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + calls := &atomic.Int64{} + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + calls.Add(1) + go serveTestSOCKS5Conn(conn) + } + }() + return "socks5h://" + listener.Addr().String(), calls +} + +func serveTestSOCKS5Conn(client net.Conn) { + defer func() { _ = client.Close() }() + header := make([]byte, 2) + if _, err := io.ReadFull(client, header); err != nil || header[0] != 5 { + return + } + methods := make([]byte, int(header[1])) + if _, err := io.ReadFull(client, methods); err != nil { + return + } + if _, err := client.Write([]byte{5, 0}); err != nil { + return + } + request := make([]byte, 4) + if _, err := io.ReadFull(client, request); err != nil || request[0] != 5 || request[1] != 1 { + return + } + var host string + switch request[3] { + case 1: + address := make([]byte, net.IPv4len) + if _, err := io.ReadFull(client, address); err != nil { + return + } + host = net.IP(address).String() + case 3: + length := make([]byte, 1) + if _, err := io.ReadFull(client, length); err != nil { + return + } + address := make([]byte, int(length[0])) + if _, err := io.ReadFull(client, address); err != nil { + return + } + host = string(address) + case 4: + address := make([]byte, net.IPv6len) + if _, err := io.ReadFull(client, address); err != nil { + return + } + host = net.IP(address).String() + default: + return + } + portBytes := make([]byte, 2) + if _, err := io.ReadFull(client, portBytes); err != nil { + return + } + target, err := net.Dial("tcp", net.JoinHostPort(host, fmt.Sprintf("%d", binary.BigEndian.Uint16(portBytes)))) + if err != nil { + _, _ = client.Write([]byte{5, 1, 0, 1, 0, 0, 0, 0, 0, 0}) + return + } + defer func() { _ = target.Close() }() + if _, err := client.Write([]byte{5, 0, 0, 1, 0, 0, 0, 0, 0, 0}); err != nil { + return + } + go func() { _, _ = io.Copy(target, client); _ = target.Close() }() + _, _ = io.Copy(client, target) +} + func TestHTTPUpstreamDoAppliesGrokCLIIdentityBeforeOAuthRoundTrip(t *testing.T) { t.Setenv("XAI_GROK_CLI_VERSION", "") diff --git a/backend/internal/repository/proxy_repo.go b/backend/internal/repository/proxy_repo.go index fcb2e53b87..6801f074e1 100644 --- a/backend/internal/repository/proxy_repo.go +++ b/backend/internal/repository/proxy_repo.go @@ -23,6 +23,8 @@ type proxyRepository struct { sql sqlExecutor } +const proxyProbeOutboxAccountChunkSize = 500 + func NewProxyRepository(client *dbent.Client, sqlDB *sql.DB) service.ProxyRepository { return newProxyRepositoryWithSQL(client, sqlDB) } @@ -91,7 +93,62 @@ func (r *proxyRepository) ListByIDs(ctx context.Context, ids []int64) ([]service } func (r *proxyRepository) Update(ctx context.Context, proxyIn *service.Proxy) error { - builder := r.client.Proxy.UpdateOneID(proxyIn.ID). + client := r.client + var tx *dbent.Tx + if contextTx := dbent.TxFromContext(ctx); contextTx != nil { + client = contextTx.Client() + } else { + var err error + tx, err = r.client.Tx(ctx) + if err != nil && err != dbent.ErrTxStarted { + return err + } + if tx != nil { + defer func() { _ = tx.Rollback() }() + ctx = dbent.NewTxContext(ctx, tx) + client = tx.Client() + } + } + + updated, err := updateProxyAndInvalidateProbeSnapshots(ctx, client, proxyIn) + if err != nil { + return err + } + if tx != nil { + if err := tx.Commit(); err != nil { + return err + } + } + applyProxyEntityToService(proxyIn, updated) + return nil +} + +type proxyProbeIdentity struct { + protocol string + host string + port int + username string + password string + status string +} + +func proxyProbeIdentityFromService(proxyIn *service.Proxy) proxyProbeIdentity { + return proxyProbeIdentity{ + protocol: proxyIn.Protocol, + host: proxyIn.Host, + port: proxyIn.Port, + username: proxyIn.Username, + password: proxyIn.Password, + status: proxyIn.Status, + } +} + +func updateProxyAndInvalidateProbeSnapshots(ctx context.Context, client *dbent.Client, proxyIn *service.Proxy) (*dbent.Proxy, error) { + currentIdentity, err := lockProxyProbeIdentity(ctx, client, proxyIn.ID) + if err != nil { + return nil, err + } + builder := client.Proxy.UpdateOneID(proxyIn.ID). SetName(proxyIn.Name). SetProtocol(proxyIn.Protocol). SetHost(proxyIn.Host). @@ -121,14 +178,92 @@ func (r *proxyRepository) Update(ctx context.Context, proxyIn *service.Proxy) er } updated, err := builder.Save(ctx) - if err == nil { - applyProxyEntityToService(proxyIn, updated) - return nil - } if dbent.IsNotFound(err) { - return service.ErrProxyNotFound + return nil, service.ErrProxyNotFound } - return err + if err != nil { + return nil, err + } + if currentIdentity == proxyProbeIdentityFromService(proxyIn) { + return updated, nil + } + accountIDs, err := invalidateProxyProbeSnapshots(ctx, client, proxyIn.ID) + if err != nil { + return nil, err + } + if err := enqueueProxyProbeAccountChanges(ctx, client, accountIDs); err != nil { + return nil, err + } + return updated, nil +} + +func lockProxyProbeIdentity(ctx context.Context, client *dbent.Client, proxyID int64) (proxyProbeIdentity, error) { + rows, err := client.QueryContext(ctx, ` + SELECT protocol, host, port, COALESCE(username, ''), COALESCE(password, ''), status + FROM proxies + WHERE id = $1 AND deleted_at IS NULL + FOR NO KEY UPDATE + `, proxyID) + if err != nil { + return proxyProbeIdentity{}, err + } + defer func() { _ = rows.Close() }() + if !rows.Next() { + if err := rows.Err(); err != nil { + return proxyProbeIdentity{}, err + } + return proxyProbeIdentity{}, service.ErrProxyNotFound + } + var identity proxyProbeIdentity + if err := rows.Scan(&identity.protocol, &identity.host, &identity.port, &identity.username, &identity.password, &identity.status); err != nil { + return proxyProbeIdentity{}, err + } + return identity, rows.Err() +} + +func invalidateProxyProbeSnapshots(ctx context.Context, exec sqlExecutor, proxyID int64) ([]int64, error) { + rows, err := exec.QueryContext(ctx, ` + UPDATE accounts + SET extra = COALESCE(extra, '{}'::jsonb) - 'upstream_billing_probe', updated_at = NOW() + WHERE proxy_id = $1 + AND platform = 'openai' + AND type = 'apikey' + AND extra ? 'upstream_billing_probe' + AND extra -> 'upstream_billing_probe' <> 'null'::jsonb + AND deleted_at IS NULL + RETURNING id + `, proxyID) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + accountIDs := make([]int64, 0) + for rows.Next() { + var accountID int64 + if err := rows.Scan(&accountID); err != nil { + return nil, err + } + accountIDs = append(accountIDs, accountID) + } + if err := rows.Err(); err != nil { + return nil, err + } + return accountIDs, nil +} + +func enqueueProxyProbeAccountChanges(ctx context.Context, exec sqlExecutor, accountIDs []int64) error { + accountIDs = sortedUniqueAccountIDs(accountIDs) + for start := 0; start < len(accountIDs); start += proxyProbeOutboxAccountChunkSize { + end := start + proxyProbeOutboxAccountChunkSize + if end > len(accountIDs) { + end = len(accountIDs) + } + payload := map[string]any{"account_ids": accountIDs[start:end]} + if err := enqueueSchedulerOutbox(ctx, exec, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil { + return err + } + } + return nil } func (r *proxyRepository) Delete(ctx context.Context, id int64) error { @@ -587,6 +722,13 @@ func (r *proxyRepository) sweepOneExpiredProxyOnExec(ctx context.Context, exec s return nil, err } if !change { + accountIDs, err := invalidateProxyProbeSnapshots(ctx, exec, proxyID) + if err != nil { + return nil, err + } + if err := enqueueProxyProbeAccountChanges(ctx, exec, accountIDs); err != nil { + return nil, err + } return nil, nil } var ( @@ -595,12 +737,24 @@ func (r *proxyRepository) sweepOneExpiredProxyOnExec(ctx context.Context, exec s ) if target == nil { rows, err = exec.QueryContext(ctx, ` - UPDATE accounts SET proxy_id=NULL, proxy_fallback_origin_id=$1, updated_at=NOW() + UPDATE accounts SET proxy_id=NULL, proxy_fallback_origin_id=$1, + extra=CASE + WHEN platform='openai' AND type='apikey' AND extra ? 'upstream_billing_probe' + THEN extra - 'upstream_billing_probe' + ELSE extra + END, + updated_at=NOW() WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL RETURNING id`, proxyID) } else { rows, err = exec.QueryContext(ctx, ` - UPDATE accounts SET proxy_id=$2, proxy_fallback_origin_id=$1, updated_at=NOW() + UPDATE accounts SET proxy_id=$2, proxy_fallback_origin_id=$1, + extra=CASE + WHEN platform='openai' AND type='apikey' AND extra ? 'upstream_billing_probe' + THEN extra - 'upstream_billing_probe' + ELSE extra + END, + updated_at=NOW() WHERE proxy_id=$1 AND proxy_fallback_origin_id IS NULL AND deleted_at IS NULL RETURNING id`, proxyID, *target) } diff --git a/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go b/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go new file mode 100644 index 0000000000..27a3ac1242 --- /dev/null +++ b/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go @@ -0,0 +1,158 @@ +package repository + +import ( + "context" + "errors" + "regexp" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +func TestProxyUpdateInvalidatesBoundProbeSnapshotsAndEnqueuesOutboxAtomically(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectQuery(`(?s)` + regexp.QuoteMeta("SELECT protocol, host, port") + `.*` + regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"protocol", "host", "port", "username", "password", "status"}). + AddRow("http", "old.example", 8080, "user", "pass", service.StatusActive)) + mock.ExpectExec(`(?s)UPDATE "proxies" SET`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`UPDATE "proxies" SET "backup_proxy_id" = NULL WHERE "backup_proxy_id" = \$1`). + WithArgs(int64(9)). + WillReturnResult(sqlmock.NewResult(0, 0)) + expectProxyUpdateReload(mock, 9, "new.example", "user", "pass") + mock.ExpectQuery(`(?s)UPDATE accounts.*platform = 'openai'.*type = 'apikey'.*extra \? 'upstream_billing_probe'.*extra -> 'upstream_billing_probe' <> 'null'::jsonb.*RETURNING id`). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(int64(17)).AddRow(int64(18))) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload)")). + WithArgs(service.SchedulerOutboxEventAccountBulkChanged, nil, nil, accountIDsPayloadMatcher{want: []int64{17, 18}}). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectCommit() + + repo := newProxyRepositoryWithSQL(client, db) + proxy := &service.Proxy{ + ID: 9, + Name: "proxy", + Protocol: "http", + Host: "new.example", + Port: 8080, + Username: "user", + Password: "pass", + Status: service.StatusActive, + } + + err = repo.Update(context.Background(), proxy) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestProxyUpdateRollsBackWhenProbeInvalidationOutboxFails(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectQuery(`(?s)` + regexp.QuoteMeta("SELECT protocol, host, port") + `.*` + regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"protocol", "host", "port", "username", "password", "status"}). + AddRow("http", "old.example", 8080, "", "", service.StatusActive)) + mock.ExpectExec(`(?s)UPDATE "proxies" SET`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`UPDATE "proxies" SET "backup_proxy_id" = NULL WHERE "backup_proxy_id" = \$1`). + WithArgs(int64(9)). + WillReturnResult(sqlmock.NewResult(0, 0)) + expectProxyUpdateReload(mock, 9, "new.example", "", "") + mock.ExpectQuery(`(?s)UPDATE accounts.*platform = 'openai'.*type = 'apikey'.*extra \? 'upstream_billing_probe'.*extra -> 'upstream_billing_probe' <> 'null'::jsonb.*RETURNING id`). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(int64(17))) + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload)")). + WillReturnError(errors.New("outbox failed")) + mock.ExpectRollback() + + repo := newProxyRepositoryWithSQL(client, db) + proxy := &service.Proxy{ID: 9, Name: "proxy", Protocol: "http", Host: "new.example", Port: 8080, Status: service.StatusActive} + + err = repo.Update(context.Background(), proxy) + + require.EqualError(t, err, "outbox failed") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestProxyUpdateSkipsProbeInvalidationForNonIdentityChange(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + + mock.ExpectBegin() + mock.ExpectQuery(`(?s)` + regexp.QuoteMeta("SELECT protocol, host, port") + `.*` + regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(int64(9)). + WillReturnRows(sqlmock.NewRows([]string{"protocol", "host", "port", "username", "password", "status"}). + AddRow("http", "same.example", 8080, "", "", service.StatusActive)) + mock.ExpectExec(`(?s)UPDATE "proxies" SET`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec(`UPDATE "proxies" SET "backup_proxy_id" = NULL WHERE "backup_proxy_id" = \$1`). + WithArgs(int64(9)). + WillReturnResult(sqlmock.NewResult(0, 0)) + expectProxyUpdateReload(mock, 9, "same.example", "", "") + mock.ExpectCommit() + + repo := newProxyRepositoryWithSQL(client, db) + proxy := &service.Proxy{ID: 9, Name: "renamed", Protocol: "http", Host: "same.example", Port: 8080, Status: service.StatusActive} + + err = repo.Update(context.Background(), proxy) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func expectProxyUpdateReload(mock sqlmock.Sqlmock, id int64, host, username, password string) { + now := time.Now() + mock.ExpectQuery(`(?s)SELECT .* FROM "proxies" WHERE "id" = \$1`). + WithArgs(id). + WillReturnRows(sqlmock.NewRows([]string{ + "id", "created_at", "updated_at", "deleted_at", "name", "protocol", "host", "port", + "username", "password", "status", "expires_at", "fallback_mode", "backup_proxy_id", "expiry_warn_days", + }).AddRow( + id, now, now, nil, "proxy", "http", host, 8080, + username, password, service.StatusActive, nil, service.FallbackModeNone, nil, 0, + )) +} + +func TestEnqueueProxyAccountChangesChunksLargePayloads(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + accountIDs := make([]int64, 1001) + for i := range accountIDs { + accountIDs[i] = int64(i + 1) + } + for start := 0; start < len(accountIDs); start += proxyProbeOutboxAccountChunkSize { + end := start + proxyProbeOutboxAccountChunkSize + if end > len(accountIDs) { + end = len(accountIDs) + } + mock.ExpectExec(regexp.QuoteMeta("INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload)")). + WithArgs(service.SchedulerOutboxEventAccountBulkChanged, nil, nil, accountIDsPayloadMatcher{want: accountIDs[start:end]}). + WillReturnResult(sqlmock.NewResult(1, 1)) + } + + err = enqueueProxyProbeAccountChanges(context.Background(), db, accountIDs) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/upstream_billing_probe_persistence_integration_test.go b/backend/internal/repository/upstream_billing_probe_persistence_integration_test.go new file mode 100644 index 0000000000..252fa47cf9 --- /dev/null +++ b/backend/internal/repository/upstream_billing_probe_persistence_integration_test.go @@ -0,0 +1,395 @@ +//go:build integration + +package repository + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestAccountUpdatePreservesConcurrentProbeSnapshot(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "probe-update-preserve", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-old"}, + Extra: map[string]any{service.UpstreamBillingProbeEnabledExtraKey: true}, + }) + + stale, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.NotContains(t, stale.Extra, service.UpstreamBillingProbeExtraKey) + require.NoError(t, repo.UpdateUpstreamBillingProbeSnapshot(ctx, stale, &service.UpstreamBillingProbeSnapshot{ + Status: service.UpstreamBillingProbeStatusOK, + LastAttemptAt: time.Now().UTC(), + })) + + stale.Name = "ordinary-edit" + require.NoError(t, repo.Update(ctx, stale)) + got, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + snapshot, ok := got.Extra[service.UpstreamBillingProbeExtraKey].(map[string]any) + require.True(t, ok) + require.Equal(t, service.UpstreamBillingProbeStatusOK, snapshot["status"]) + + require.NoError(t, repo.UpdateExtra(ctx, got.ID, map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false})) + disabled, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.NotContains(t, disabled.Extra, service.UpstreamBillingProbeExtraKey) +} + +func TestAccountUpdatePreservesConcurrentProbeEnableFlag(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "probe-update-enable", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": service.UpstreamBillingProbeStatusOK}, + }, + }) + + stale, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.NoError(t, repo.UpdateExtra(ctx, account.ID, map[string]any{service.UpstreamBillingProbeEnabledExtraKey: false})) + stale.Name = "ordinary-edit" + require.NoError(t, repo.Update(ctx, stale)) + + got, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.Equal(t, false, got.Extra[service.UpstreamBillingProbeEnabledExtraKey]) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) +} + +func TestAccountUpdateClearsProbeSnapshotWhenIdentityChanges(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "probe-update-identity", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-old"}, + Extra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": service.UpstreamBillingProbeStatusOK}, + }, + }) + + loaded, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + loaded.Credentials["api_key"] = "sk-new" + require.NoError(t, repo.Update(ctx, loaded)) + + got, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) +} + +func TestBulkUpdateAndCredentialUpdateDeleteProbeKey(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + newAccount := func(name string) *service.Account { + return mustCreateAccount(t, tx.Client(), &service.Account{ + Name: name, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-old"}, + Extra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": service.UpstreamBillingProbeStatusOK}, + }, + }) + } + + bulkAccount := newAccount("probe-bulk-clear") + _, err := repo.BulkUpdate(ctx, []int64{bulkAccount.ID}, service.AccountBulkUpdate{ + Extra: map[string]any{service.UpstreamBillingProbeExtraKey: nil}, + }) + require.NoError(t, err) + got, err := repo.GetByID(ctx, bulkAccount.ID) + require.NoError(t, err) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + + credentialAccount := newAccount("probe-credentials-clear") + require.NoError(t, repo.UpdateCredentials(ctx, credentialAccount.ID, map[string]any{"api_key": "sk-new"})) + got, err = repo.GetByID(ctx, credentialAccount.ID) + require.NoError(t, err) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) +} + +func TestProbeSnapshotCASIncludesLoadedEnabledState(t *testing.T) { + tests := []struct { + name string + loadedEnabled bool + concurrentFlip *bool + wantConflict bool + }{ + {name: "manual_false_stays_false", loadedEnabled: false}, + {name: "periodic_true_disabled_in_flight", loadedEnabled: true, concurrentFlip: boolPtr(false), wantConflict: true}, + {name: "manual_false_enabled_in_flight", loadedEnabled: false, concurrentFlip: boolPtr(true), wantConflict: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "probe-enabled-cas-" + tt.name, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{service.UpstreamBillingProbeEnabledExtraKey: tt.loadedEnabled}, + }) + inFlight, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + if tt.concurrentFlip != nil { + require.NoError(t, repo.UpdateExtra(ctx, account.ID, map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: *tt.concurrentFlip, + })) + } + + err = repo.UpdateUpstreamBillingProbeSnapshot(ctx, inFlight, &service.UpstreamBillingProbeSnapshot{ + Status: service.UpstreamBillingProbeStatusOK, + LastAttemptAt: time.Now().UTC(), + }) + if tt.wantConflict { + require.ErrorIs(t, err, service.ErrUpstreamBillingProbeIdentityChanged) + } else { + require.NoError(t, err) + } + got, err := repo.GetByID(ctx, account.ID) + require.NoError(t, err) + if tt.wantConflict { + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + } else { + require.Contains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + } + }) + } +} + +func boolPtr(value bool) *bool { + return &value +} + +func TestProxyIdentityUpdateInvalidatesProbeAndRejectsInFlightSnapshot(t *testing.T) { + tests := []struct { + name string + includeProbeKey bool + probeValue any + wantInvalidation bool + }{ + {name: "missing_snapshot"}, + {name: "json_null_snapshot", includeProbeKey: true}, + {name: "existing_snapshot", includeProbeKey: true, probeValue: map[string]any{"status": service.UpstreamBillingProbeStatusOK}, wantInvalidation: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + accountRepo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + proxyRepo := newProxyRepositoryWithSQL(tx.Client(), tx) + proxy := mustCreateProxy(t, tx.Client(), &service.Proxy{ + Name: "probe-proxy", + Protocol: "http", + Host: "old.example", + Port: 8080, + Username: "old-user", + Password: "old-pass", + Status: service.StatusActive, + }) + extra := map[string]any{service.UpstreamBillingProbeEnabledExtraKey: true} + if tt.includeProbeKey { + extra[service.UpstreamBillingProbeExtraKey] = tt.probeValue + } + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "proxy-probe-account", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: extra, + ProxyID: &proxy.ID, + }) + inFlight, err := accountRepo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.NotNil(t, inFlight.Proxy) + require.Equal(t, "old.example", inFlight.Proxy.Host) + + proxyToUpdate, err := proxyRepo.GetByID(ctx, proxy.ID) + require.NoError(t, err) + proxyToUpdate.Host = "new.example" + require.NoError(t, proxyRepo.Update(ctx, proxyToUpdate)) + + got, err := accountRepo.GetByID(ctx, account.ID) + require.NoError(t, err) + if tt.wantInvalidation || !tt.includeProbeKey { + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + } else { + require.Contains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + require.Nil(t, got.Extra[service.UpstreamBillingProbeExtraKey]) + } + if !tt.wantInvalidation { + require.Equal(t, inFlight.UpdatedAt, got.UpdatedAt, "missing/null snapshots must not cause an account row write") + } + err = accountRepo.UpdateUpstreamBillingProbeSnapshot(ctx, inFlight, &service.UpstreamBillingProbeSnapshot{ + Status: service.UpstreamBillingProbeStatusOK, + LastAttemptAt: time.Now().UTC(), + }) + require.ErrorIs(t, err, service.ErrUpstreamBillingProbeIdentityChanged) + + rows, err := tx.QueryContext(ctx, ` + SELECT COUNT(*), COALESCE(MAX(payload::text), '') + FROM scheduler_outbox + WHERE event_type = $1 + `, service.SchedulerOutboxEventAccountBulkChanged) + require.NoError(t, err) + require.True(t, rows.Next()) + var ( + outboxCount int + payloadJSON string + ) + require.NoError(t, rows.Scan(&outboxCount, &payloadJSON)) + require.NoError(t, rows.Close()) + if tt.wantInvalidation { + require.Equal(t, 1, outboxCount) + var payload struct { + AccountIDs []int64 `json:"account_ids"` + } + require.NoError(t, json.Unmarshal([]byte(payloadJSON), &payload)) + require.Equal(t, []int64{account.ID}, payload.AccountIDs) + } else { + require.Zero(t, outboxCount, "no snapshot change means no PR2 cache invalidation event") + } + }) + } +} + +func TestSweepExpiredProxyWithoutFallbackInvalidatesOnlyExistingProbeSnapshot(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + proxyRepo := newProxyRepositoryWithSQL(tx.Client(), tx) + accountRepo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + past := time.Now().Add(-time.Hour) + proxy := &service.Proxy{ + Name: "expired-probe-proxy-none", + Protocol: "http", + Host: "127.0.0.1", + Port: 8080, + Status: service.StatusActive, + ExpiresAt: &past, + FallbackMode: service.FallbackModeNone, + ExpiryWarnDays: 7, + } + require.NoError(t, proxyRepo.Create(ctx, proxy)) + newAccount := func(name string, probe any, includeProbe bool) *service.Account { + extra := map[string]any{service.UpstreamBillingProbeEnabledExtraKey: true} + if includeProbe { + extra[service.UpstreamBillingProbeExtraKey] = probe + } + return mustCreateAccount(t, tx.Client(), &service.Account{ + Name: name, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: extra, + ProxyID: &proxy.ID, + }) + } + withSnapshot := newAccount("expired-proxy-with-snapshot", map[string]any{"status": service.UpstreamBillingProbeStatusOK}, true) + withoutSnapshot := newAccount("expired-proxy-without-snapshot", nil, false) + withJSONNull := newAccount("expired-proxy-null-snapshot", nil, true) + untouchedUpdatedAt := make(map[int64]time.Time, 2) + for _, untouched := range []*service.Account{withoutSnapshot, withJSONNull} { + loaded, err := accountRepo.GetByID(ctx, untouched.ID) + require.NoError(t, err) + untouchedUpdatedAt[untouched.ID] = loaded.UpdatedAt + } + + changed, err := proxyRepo.SweepExpiredProxies(ctx, time.Now()) + require.NoError(t, err) + require.Zero(t, changed, "probe invalidation must not inflate the rerouted account count") + + got, err := accountRepo.GetByID(ctx, withSnapshot.ID) + require.NoError(t, err) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + for _, untouched := range []*service.Account{withoutSnapshot, withJSONNull} { + got, err = accountRepo.GetByID(ctx, untouched.ID) + require.NoError(t, err) + require.Equal(t, untouchedUpdatedAt[untouched.ID], got.UpdatedAt) + } + + payload := latestBulkAccountOutboxPayload(t, ctx, tx) + require.Equal(t, []int64{withSnapshot.ID}, payload) +} + +func TestSweepExpiredProxyFallbackRerouteDeletesProbeSnapshot(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + proxyRepo := newProxyRepositoryWithSQL(tx.Client(), tx) + accountRepo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + past := time.Now().Add(-time.Hour) + proxy := &service.Proxy{ + Name: "expired-probe-proxy-direct", + Protocol: "http", + Host: "127.0.0.1", + Port: 8080, + Status: service.StatusActive, + ExpiresAt: &past, + FallbackMode: service.FallbackModeDirect, + ExpiryWarnDays: 7, + } + require.NoError(t, proxyRepo.Create(ctx, proxy)) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "expired-proxy-rerouted-snapshot", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{ + service.UpstreamBillingProbeEnabledExtraKey: true, + service.UpstreamBillingProbeExtraKey: map[string]any{"status": service.UpstreamBillingProbeStatusOK}, + }, + ProxyID: &proxy.ID, + }) + + changed, err := proxyRepo.SweepExpiredProxies(ctx, time.Now()) + require.NoError(t, err) + require.EqualValues(t, 1, changed) + + got, err := accountRepo.GetByID(ctx, account.ID) + require.NoError(t, err) + require.Nil(t, got.ProxyID) + require.NotContains(t, got.Extra, service.UpstreamBillingProbeExtraKey) + require.Equal(t, []int64{account.ID}, latestBulkAccountOutboxPayload(t, ctx, tx)) +} + +func latestBulkAccountOutboxPayload(t *testing.T, ctx context.Context, tx sqlQueryer) []int64 { + t.Helper() + var payloadJSON []byte + require.NoError(t, scanSingleRow(ctx, tx, ` + SELECT payload + FROM scheduler_outbox + WHERE event_type = $1 + ORDER BY id DESC + LIMIT 1 + `, []any{service.SchedulerOutboxEventAccountBulkChanged}, &payloadJSON)) + var payload struct { + AccountIDs []int64 `json:"account_ids"` + } + require.NoError(t, json.Unmarshal(payloadJSON, &payload)) + return payload.AccountIDs +} diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index a60adb1826..b6aebb7029 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -295,6 +295,9 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) { accounts := admin.Group("/accounts") { accounts.GET("", h.Admin.Account.List) + accounts.GET("/upstream-billing-probe/settings", h.Admin.Account.GetUpstreamBillingProbeSettings) + accounts.PUT("/upstream-billing-probe/settings", h.Admin.Account.UpdateUpstreamBillingProbeSettings) + accounts.POST("/upstream-billing-probe/batch", h.Admin.Account.ProbeUpstreamBillingBatch) accounts.GET("/:id", h.Admin.Account.GetByID) accounts.POST("", h.Admin.Account.Create) accounts.POST("/:id/duplicate", h.Admin.Account.Duplicate) @@ -303,6 +306,8 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) { accounts.POST("/sync/crs", h.Admin.Account.SyncFromCRS) accounts.POST("/sync/crs/preview", h.Admin.Account.PreviewFromCRS) accounts.PUT("/:id", h.Admin.Account.Update) + accounts.PUT("/:id/upstream-billing-probe", h.Admin.Account.SetUpstreamBillingProbeEnabled) + accounts.POST("/:id/upstream-billing-probe", h.Admin.Account.ProbeUpstreamBilling) accounts.DELETE("/:id", h.Admin.Account.Delete) accounts.POST("/:id/test", h.Admin.Account.Test) accounts.POST("/:id/recover-state", h.Admin.Account.RecoverState) diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index f4c3375650..904153f5aa 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -9,6 +9,7 @@ import ( "log/slog" "maps" "net/http" + "reflect" "strconv" "strings" "time" @@ -394,6 +395,9 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat } func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { + // Probe state is system-managed. New accounts always start with auto probe disabled. + delete(accountExtra, UpstreamBillingProbeEnabledExtraKey) + delete(accountExtra, UpstreamBillingProbeExtraKey) account := &Account{ Name: input.Name, Notes: normalizeAccountNotes(input.Notes), @@ -516,6 +520,10 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou return account, nil } +type accountProbeEnabledAtomicUpdater interface { + UpdateWithUpstreamBillingProbeEnabled(context.Context, *Account, bool) error +} + func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *UpdateAccountInput) (*Account, error) { account, err := s.accountRepo.GetByID(ctx, id) if err != nil { @@ -528,6 +536,7 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U return nil, err } } + previousProbeIdentity := upstreamBillingProbeIdentity(account) // 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用, // 必须在此守住,否则仅在创建时的保证可被这些路径绕过。 if account.IsCredentialShadow() { @@ -579,13 +588,39 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U } // Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。 // 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。 + var requestedProbeEnabledUpdate *bool if input.Extra != nil { + requestedProbeEnabled, hasRequestedProbeEnabled := normalizedExtra[UpstreamBillingProbeEnabledExtraKey] + if hasRequestedProbeEnabled { + enabled, ok := requestedProbeEnabled.(bool) + if !ok { + return nil, infraerrors.BadRequest("INVALID_UPSTREAM_BILLING_PROBE_ENABLED", "upstream_billing_probe_enabled must be a boolean") + } + requestedProbeEnabledUpdate = &enabled + } + delete(normalizedExtra, UpstreamBillingProbeEnabledExtraKey) + delete(normalizedExtra, UpstreamBillingProbeExtraKey) // 保留配额用量字段,防止编辑账号时意外重置 - for _, key := range []string{"quota_used", "quota_daily_used", "quota_daily_start", "quota_weekly_used", "quota_weekly_start"} { + for _, key := range []string{ + "quota_used", + "quota_daily_used", + "quota_daily_start", + "quota_weekly_used", + "quota_weekly_start", + UpstreamBillingProbeEnabledExtraKey, + UpstreamBillingProbeExtraKey, + } { if v, ok := account.Extra[key]; ok { normalizedExtra[key] = v } } + if hasRequestedProbeEnabled { + if isUpstreamBillingProbeAccount(account) { + normalizedExtra[UpstreamBillingProbeEnabledExtraKey] = requestedProbeEnabled + } else { + delete(normalizedExtra, UpstreamBillingProbeEnabledExtraKey) + } + } account.Extra = normalizedExtra if account.Platform == PlatformAntigravity && wasOveragesEnabled && !account.IsOveragesEnabled() { delete(account.Extra, "antigravity_credits_overages") // 清理旧版 overages 运行态 @@ -616,6 +651,12 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U } account.Proxy = nil // 清除关联对象,防止 GORM Save 时根据 Proxy.ID 覆盖 ProxyID } + if !reflect.DeepEqual(previousProbeIdentity, upstreamBillingProbeIdentity(account)) && account.Extra != nil { + delete(account.Extra, UpstreamBillingProbeExtraKey) + if !isUpstreamBillingProbeAccount(account) { + delete(account.Extra, UpstreamBillingProbeEnabledExtraKey) + } + } // 只在指针非 nil 时更新 Concurrency(支持设置为 0) if input.Concurrency != nil { account.Concurrency = normalizeAccountConcurrency(account.Platform, account.Type, *input.Concurrency) @@ -668,8 +709,26 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U } } - if err := s.accountRepo.Update(ctx, account); err != nil { - return nil, err + probeEnabledAppliedAtomically := false + if requestedProbeEnabledUpdate != nil && isUpstreamBillingProbeAccount(account) { + if updater, ok := s.accountRepo.(accountProbeEnabledAtomicUpdater); ok { + if err := updater.UpdateWithUpstreamBillingProbeEnabled(ctx, account, *requestedProbeEnabledUpdate); err != nil { + return nil, err + } + probeEnabledAppliedAtomically = true + } + } + if !probeEnabledAppliedAtomically { + if err := s.accountRepo.Update(ctx, account); err != nil { + return nil, err + } + if requestedProbeEnabledUpdate != nil && isUpstreamBillingProbeAccount(account) { + if err := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + UpstreamBillingProbeEnabledExtraKey: *requestedProbeEnabledUpdate, + }); err != nil { + return nil, err + } + } } // 将 proxy 变更传播到 spark 影子账号(同步;Update 内部已触发调度快照)。 @@ -716,6 +775,10 @@ func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, upd // BulkUpdateAccounts updates multiple accounts in one request. // It merges credentials/extra keys instead of overwriting the whole object. func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUpdateAccountsInput) (*BulkUpdateAccountsResult, error) { + // Probe state is updated only through its dedicated endpoints. + delete(input.Extra, UpstreamBillingProbeEnabledExtraKey) + delete(input.Extra, UpstreamBillingProbeExtraKey) + if len(input.AccountIDs) == 0 && input.Filters != nil { accountIDs, err := s.resolveBulkUpdateTargetIDs(ctx, input.Filters) if err != nil { @@ -825,6 +888,14 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp Credentials: input.Credentials, Extra: input.Extra, } + if updatesUpstreamBillingProbeIdentity(input.Credentials) || input.ProxyID != nil { + if repoUpdates.Extra == nil { + repoUpdates.Extra = make(map[string]any) + } + // JSON null makes every reader treat the old snapshot as absent and lets the + // next enabled runner cycle probe the new upstream identity immediately. + repoUpdates.Extra[UpstreamBillingProbeExtraKey] = nil + } if input.Name != "" { repoUpdates.Name = &input.Name } @@ -898,6 +969,31 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp return result, nil } +func updatesUpstreamBillingProbeIdentity(credentials map[string]any) bool { + for _, key := range []string{"api_key", "base_url", credKeyHeaderOverrideEnabled, credKeyHeaderOverrides} { + if _, ok := credentials[key]; ok { + return true + } + } + return false +} + +func upstreamBillingProbeIdentity(account *Account) map[string]any { + if account == nil { + return nil + } + identity := map[string]any{"platform": account.Platform, "type": account.Type, "proxy_id": nil} + if account.ProxyID != nil { + identity["proxy_id"] = *account.ProxyID + } + for _, key := range []string{"api_key", "base_url", credKeyHeaderOverrideEnabled, credKeyHeaderOverrides} { + if value, ok := account.Credentials[key]; ok { + identity[key] = value + } + } + return identity +} + func (s *adminServiceImpl) resolveBulkUpdateTargetIDs(ctx context.Context, filters *BulkUpdateAccountFilters) ([]int64, error) { if filters == nil { return nil, nil diff --git a/backend/internal/service/admin_account_upstream_billing_probe_test.go b/backend/internal/service/admin_account_upstream_billing_probe_test.go new file mode 100644 index 0000000000..317423aaec --- /dev/null +++ b/backend/internal/service/admin_account_upstream_billing_probe_test.go @@ -0,0 +1,378 @@ +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type upstreamBillingProbeAdminRepo struct { + *upstreamBillingProbeAccountRepo +} + +func (r *upstreamBillingProbeAdminRepo) ListShadowsByParent(context.Context, int64) ([]*Account, error) { + return nil, nil +} + +func TestCreateAccountDropsManagedUpstreamBillingProbeState(t *testing.T) { + repo := &upstreamBillingProbeAccountRepo{} + svc := &adminServiceImpl{accountRepo: repo} + + created, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Name: "upstream", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-test"}, + SkipDefaultGroupBind: true, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }) + + require.NoError(t, err) + require.NotContains(t, created.Extra, UpstreamBillingProbeEnabledExtraKey) + require.NotContains(t, created.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpdateAccountPreservesManagedUpstreamBillingProbeStateForUnrelatedEdit(t *testing.T) { + accountID := int64(110) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + svc := &adminServiceImpl{accountRepo: repo} + updated, err := svc.UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Extra: map[string]any{"custom": "value"}, + }) + + require.NoError(t, err) + require.Equal(t, true, updated.Extra[UpstreamBillingProbeEnabledExtraKey]) + require.Contains(t, updated.Extra, UpstreamBillingProbeExtraKey) + require.Equal(t, "value", updated.Extra["custom"]) +} + +func TestUpdateAccountPreservesProbeSnapshotWhenIdentityValuesAreUnchanged(t *testing.T) { + accountID := int64(119) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Credentials: map[string]any{ + "api_key": "sk-existing", + "base_url": "https://upstream.example", + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"x-route": "stable"}, + }, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + updated, err := (&adminServiceImpl{accountRepo: repo}).UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Credentials: map[string]any{ + "base_url": "https://upstream.example", + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"x-route": "stable"}, + }, + }) + + require.NoError(t, err) + require.Contains(t, updated.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpdateAccountInvalidatesProbeSnapshotWhenUpstreamIdentityChanges(t *testing.T) { + tests := []struct { + name string + input *UpdateAccountInput + wantEnabled bool + }{ + { + name: "api key", + input: &UpdateAccountInput{Credentials: map[string]any{"api_key": "sk-new"}}, + wantEnabled: true, + }, + { + name: "base url", + input: &UpdateAccountInput{Credentials: map[string]any{"base_url": "https://new.example"}}, + wantEnabled: true, + }, + { + name: "header override", + input: &UpdateAccountInput{Credentials: map[string]any{ + credKeyHeaderOverrideEnabled: true, + credKeyHeaderOverrides: map[string]any{"x-route": "new"}, + }}, + wantEnabled: true, + }, + { + name: "account type", + input: &UpdateAccountInput{Type: AccountTypeOAuth}, + wantEnabled: false, + }, + } + + for i, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + accountID := int64(120 + i) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Credentials: map[string]any{ + "api_key": "sk-old", + "base_url": "https://old.example", + }, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + updated, err := (&adminServiceImpl{accountRepo: repo}).UpdateAccount(context.Background(), accountID, tt.input) + + require.NoError(t, err) + require.NotContains(t, updated.Extra, UpstreamBillingProbeExtraKey) + if tt.wantEnabled { + require.Equal(t, true, updated.Extra[UpstreamBillingProbeEnabledExtraKey]) + } else { + require.NotContains(t, updated.Extra, UpstreamBillingProbeEnabledExtraKey) + } + }) + } +} + +func TestUpdateAccountInvalidatesProbeSnapshotWhenProxyChanges(t *testing.T) { + accountID := int64(140) + oldProxyID := int64(7) + newProxyID := int64(8) + baseRepo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Credentials: map[string]any{"api_key": "sk-test"}, + ProxyID: &oldProxyID, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + updated, err := (&adminServiceImpl{accountRepo: &upstreamBillingProbeAdminRepo{baseRepo}}).UpdateAccount( + context.Background(), + accountID, + &UpdateAccountInput{ProxyID: &newProxyID}, + ) + + require.NoError(t, err) + require.Equal(t, newProxyID, *updated.ProxyID) + require.NotContains(t, updated.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpdateAccountPreservesProbeSnapshotWhenProxyIsUnchanged(t *testing.T) { + accountID := int64(141) + existingProxyID := int64(7) + unchangedProxyID := int64(7) + baseRepo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Credentials: map[string]any{"api_key": "sk-test"}, + ProxyID: &existingProxyID, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + updated, err := (&adminServiceImpl{accountRepo: &upstreamBillingProbeAdminRepo{baseRepo}}).UpdateAccount( + context.Background(), + accountID, + &UpdateAccountInput{ProxyID: &unchangedProxyID}, + ) + + require.NoError(t, err) + require.Contains(t, updated.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpdateAccountAcceptsProbeEnabledAndRejectsInjectedSnapshot(t *testing.T) { + accountID := int64(111) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{}, + }, + }} + + svc := &adminServiceImpl{accountRepo: repo} + updated, err := svc.UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }) + + require.NoError(t, err) + require.Equal(t, true, updated.Extra[UpstreamBillingProbeEnabledExtraKey]) + require.NotContains(t, updated.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpdateAccountExplicitProbeDisableUsesDedicatedExtraUpdate(t *testing.T) { + accountID := int64(113) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + }, + }} + + _, err := (&adminServiceImpl{accountRepo: repo}).UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: false}, + }) + + require.NoError(t, err) + require.Len(t, repo.updates[accountID], 1) + require.Equal(t, false, repo.updates[accountID][0][UpstreamBillingProbeEnabledExtraKey]) +} + +func TestUpdateAccountExplicitUnchangedProbeEnabledStillUsesDedicatedExtraUpdate(t *testing.T) { + accountID := int64(114) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + }, + }} + + _, err := (&adminServiceImpl{accountRepo: repo}).UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + }) + + require.NoError(t, err) + require.Len(t, repo.updates[accountID], 1) + require.Equal(t, true, repo.updates[accountID][0][UpstreamBillingProbeEnabledExtraKey]) +} + +func TestUpdateAccountRejectsInvalidProbeEnabled(t *testing.T) { + accountID := int64(112) + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + accountID: { + ID: accountID, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Extra: map[string]any{}, + }, + }} + + svc := &adminServiceImpl{accountRepo: repo} + _, err := svc.UpdateAccount(context.Background(), accountID, &UpdateAccountInput{ + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: "true"}, + }) + + require.Error(t, err) +} + +func TestBulkUpdateAccountsDropsManagedUpstreamBillingProbeState(t *testing.T) { + repo := &upstreamBillingProbeAccountRepo{} + svc := &adminServiceImpl{accountRepo: repo} + input := &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{ + "custom": "value", + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "ok"}, + }, + } + + result, err := svc.BulkUpdateAccounts(context.Background(), input) + + require.NoError(t, err) + require.Equal(t, 1, result.Success) + require.Len(t, repo.bulkUpdates, 1) + require.Equal(t, "value", repo.bulkUpdates[0].Extra["custom"]) + require.NotContains(t, repo.bulkUpdates[0].Extra, UpstreamBillingProbeEnabledExtraKey) + require.NotContains(t, repo.bulkUpdates[0].Extra, UpstreamBillingProbeExtraKey) +} + +func TestBulkUpdateAccountsInvalidatesProbeSnapshotForIdentityCredentials(t *testing.T) { + repo := &upstreamBillingProbeAccountRepo{} + input := &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: map[string]any{"api_key": "sk-new"}, + } + + result, err := (&adminServiceImpl{accountRepo: repo}).BulkUpdateAccounts(context.Background(), input) + + require.NoError(t, err) + require.Equal(t, 1, result.Success) + require.Len(t, repo.bulkUpdates, 1) + require.Contains(t, repo.bulkUpdates[0].Extra, UpstreamBillingProbeExtraKey) + require.Nil(t, repo.bulkUpdates[0].Extra[UpstreamBillingProbeExtraKey]) +} + +func TestBulkUpdateAccountsInvalidatesProbeSnapshotForProxyUpdate(t *testing.T) { + proxyID := int64(9) + baseRepo := &upstreamBillingProbeAccountRepo{} + input := &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + ProxyID: &proxyID, + } + + result, err := (&adminServiceImpl{accountRepo: &upstreamBillingProbeAdminRepo{baseRepo}}).BulkUpdateAccounts(context.Background(), input) + + require.NoError(t, err) + require.Equal(t, 1, result.Success) + require.Len(t, baseRepo.bulkUpdates, 1) + require.Contains(t, baseRepo.bulkUpdates[0].Extra, UpstreamBillingProbeExtraKey) + require.Nil(t, baseRepo.bulkUpdates[0].Extra[UpstreamBillingProbeExtraKey]) +} + +func TestBulkUpdateAccountsKeepsProbeSnapshotForUnrelatedCredentials(t *testing.T) { + repo := &upstreamBillingProbeAccountRepo{} + input := &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Credentials: map[string]any{"model_mapping": map[string]any{"gpt-old": "gpt-new"}}, + } + + _, err := (&adminServiceImpl{accountRepo: repo}).BulkUpdateAccounts(context.Background(), input) + + require.NoError(t, err) + require.Len(t, repo.bulkUpdates, 1) + require.NotContains(t, repo.bulkUpdates[0].Extra, UpstreamBillingProbeExtraKey) +} diff --git a/backend/internal/service/admin_service_duplicate_account_test.go b/backend/internal/service/admin_service_duplicate_account_test.go index f89d887035..7b19dbf6f5 100644 --- a/backend/internal/service/admin_service_duplicate_account_test.go +++ b/backend/internal/service/admin_service_duplicate_account_test.go @@ -140,6 +140,8 @@ func TestDuplicateAccountCopiesConfigurationAndResetsRuntimeState(t *testing.T) SessionWindowEnd: &sessionWindowEnd, SessionWindowStatus: "active", } + source.Extra[UpstreamBillingProbeEnabledExtraKey] = true + source.Extra[UpstreamBillingProbeExtraKey] = map[string]any{"status": "ok"} require.NoError(t, repo.Create(ctx, source)) duplicate, err := svc.DuplicateAccount(ctx, source.ID, "admin:1", "") diff --git a/backend/internal/service/crs_sync_helpers_test.go b/backend/internal/service/crs_sync_helpers_test.go index 0dc053353d..bd6434c1f5 100644 --- a/backend/internal/service/crs_sync_helpers_test.go +++ b/backend/internal/service/crs_sync_helpers_test.go @@ -2,6 +2,8 @@ package service import ( "testing" + + "github.com/stretchr/testify/require" ) func TestBuildSelectedSet(t *testing.T) { @@ -110,3 +112,61 @@ func TestShouldCreateAccount(t *testing.T) { }) } } + +func TestReconcileCRSUpstreamBillingProbeExtra(t *testing.T) { + remote := map[string]any{ + "crs_account_id": "remote-1", + UpstreamBillingProbeEnabledExtraKey: true, + UpstreamBillingProbeExtraKey: map[string]any{"status": "remote"}, + } + + t.Run("create drops remote managed fields", func(t *testing.T) { + extra := mergeMap(nil, remote) + reconcileCRSUpstreamBillingProbeExtra(nil, PlatformOpenAI, AccountTypeAPIKey, map[string]any{"api_key": "new"}, extra) + require.NotContains(t, extra, UpstreamBillingProbeEnabledExtraKey) + require.NotContains(t, extra, UpstreamBillingProbeExtraKey) + }) + + existing := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "local", "base_url": "http://127.0.0.1:8080"}, + Extra: map[string]any{ + UpstreamBillingProbeEnabledExtraKey: false, + UpstreamBillingProbeExtraKey: map[string]any{"status": "local"}, + }, + } + + t.Run("same identity keeps local state", func(t *testing.T) { + extra := mergeMap(existing.Extra, remote) + reconcileCRSUpstreamBillingProbeExtra(existing, existing.Platform, existing.Type, mergeMap(existing.Credentials, nil), extra) + require.Equal(t, false, extra[UpstreamBillingProbeEnabledExtraKey]) + require.Equal(t, map[string]any{"status": "local"}, extra[UpstreamBillingProbeExtraKey]) + }) + + t.Run("identity change keeps enabled and clears snapshot", func(t *testing.T) { + extra := mergeMap(existing.Extra, remote) + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformOpenAI, AccountTypeAPIKey, map[string]any{"api_key": "changed"}, extra) + require.Equal(t, false, extra[UpstreamBillingProbeEnabledExtraKey]) + require.NotContains(t, extra, UpstreamBillingProbeExtraKey) + }) + + for _, target := range []struct { + name string + platform string + typeName string + }{ + {name: "anthropic oauth", platform: PlatformAnthropic, typeName: AccountTypeOAuth}, + {name: "anthropic api key", platform: PlatformAnthropic, typeName: AccountTypeAPIKey}, + {name: "openai oauth", platform: PlatformOpenAI, typeName: AccountTypeOAuth}, + {name: "gemini oauth", platform: PlatformGemini, typeName: AccountTypeOAuth}, + {name: "gemini api key", platform: PlatformGemini, typeName: AccountTypeAPIKey}, + } { + t.Run(target.name+" removes inapplicable state", func(t *testing.T) { + extra := mergeMap(existing.Extra, remote) + reconcileCRSUpstreamBillingProbeExtra(existing, target.platform, target.typeName, existing.Credentials, extra) + require.NotContains(t, extra, UpstreamBillingProbeEnabledExtraKey) + require.NotContains(t, extra, UpstreamBillingProbeExtraKey) + }) + } +} diff --git a/backend/internal/service/crs_sync_service.go b/backend/internal/service/crs_sync_service.go index d0abc74038..8a86c5127a 100644 --- a/backend/internal/service/crs_sync_service.go +++ b/backend/internal/service/crs_sync_service.go @@ -9,6 +9,7 @@ import ( "io" "log/slog" "net/http" + "reflect" "strconv" "strings" "time" @@ -363,6 +364,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + extra = mergeMap(existing.Extra, extra) + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformAnthropic, targetType, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -413,11 +419,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput } // Update existing - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformAnthropic existing.Type = targetType - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -494,6 +500,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + extra = mergeMap(existing.Extra, extra) + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformAnthropic, AccountTypeAPIKey, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -537,11 +548,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformAnthropic existing.Type = AccountTypeAPIKey - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -645,6 +656,10 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformOpenAI, AccountTypeOAuth, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -687,7 +702,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeOAuth - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -792,6 +807,10 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformOpenAI, AccountTypeAPIKey, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -840,7 +859,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeAPIKey - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -917,6 +936,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + extra = mergeMap(existing.Extra, extra) + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformGemini, AccountTypeOAuth, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -963,11 +987,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformGemini existing.Type = AccountTypeOAuth - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -1042,6 +1066,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + if existing != nil { + extra = mergeMap(existing.Extra, extra) + credentials = mergeMap(existing.Credentials, credentials) + } + reconcileCRSUpstreamBillingProbeExtra(existing, PlatformGemini, AccountTypeAPIKey, credentials, extra) if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -1085,11 +1114,11 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformGemini existing.Type = AccountTypeAPIKey - existing.Credentials = mergeMap(existing.Credentials, credentials) + existing.Credentials = credentials if proxyID != nil { existing.ProxyID = proxyID } @@ -1125,6 +1154,31 @@ func mergeMap(existing map[string]any, updates map[string]any) map[string]any { return out } +func reconcileCRSUpstreamBillingProbeExtra( + existing *Account, + targetPlatform, targetType string, + targetCredentials map[string]any, + extra map[string]any, +) { + delete(extra, UpstreamBillingProbeEnabledExtraKey) + delete(extra, UpstreamBillingProbeExtraKey) + if existing == nil { + return + } + if targetPlatform != PlatformOpenAI || targetType != AccountTypeAPIKey { + return + } + if enabled, ok := existing.Extra[UpstreamBillingProbeEnabledExtraKey]; ok { + extra[UpstreamBillingProbeEnabledExtraKey] = enabled + } + target := &Account{Platform: targetPlatform, Type: targetType, Credentials: targetCredentials} + if reflect.DeepEqual(upstreamBillingProbeIdentity(existing), upstreamBillingProbeIdentity(target)) { + if snapshot, ok := existing.Extra[UpstreamBillingProbeExtraKey]; ok { + extra[UpstreamBillingProbeExtraKey] = snapshot + } + } +} + func mergeCRSOpenAILongContextBillingExtra(existing, updates map[string]any) (map[string]any, error) { return normalizeOpenAILongContextBillingExtra(PlatformOpenAI, mergeMap(existing, updates)) } diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 44631e3f3c..7fb78b9b6d 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -369,6 +369,10 @@ const ( // sidebar entry is hidden. Defaults to false (opt-in feature). SettingKeyAvailableChannelsEnabled = "available_channels_enabled" + // SettingKeyUpstreamBillingProbeSettings stores the global enable switch and interval + // for probing remote Sub2API API-key billing metadata. + SettingKeyUpstreamBillingProbeSettings = "upstream_billing_probe_settings" + // ========================= // Overload Cooldown (529) // ========================= diff --git a/backend/internal/service/http_upstream_profile.go b/backend/internal/service/http_upstream_profile.go index 2d63bbd5e3..83bd6e70fc 100644 --- a/backend/internal/service/http_upstream_profile.go +++ b/backend/internal/service/http_upstream_profile.go @@ -12,6 +12,7 @@ const ( ) type httpUpstreamProfileContextKey struct{} +type httpUpstreamDisableRedirectsContextKey struct{} // WithHTTPUpstreamProfile injects an upstream transport profile into ctx. func WithHTTPUpstreamProfile(ctx context.Context, profile HTTPUpstreamProfile) context.Context { @@ -40,3 +41,16 @@ func HTTPUpstreamProfileFromContext(ctx context.Context) HTTPUpstreamProfile { return HTTPUpstreamProfileDefault } } + +// WithHTTPUpstreamRedirectsDisabled prevents credential-bearing probes from +// following redirects through the shared upstream client. +func WithHTTPUpstreamRedirectsDisabled(ctx context.Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, httpUpstreamDisableRedirectsContextKey{}, true) +} + +func HTTPUpstreamRedirectsDisabled(ctx context.Context) bool { + return ctx != nil && ctx.Value(httpUpstreamDisableRedirectsContextKey{}) == true +} diff --git a/backend/internal/service/http_upstream_profile_test.go b/backend/internal/service/http_upstream_profile_test.go index 96f0cd3109..9cd4bf4ff9 100644 --- a/backend/internal/service/http_upstream_profile_test.go +++ b/backend/internal/service/http_upstream_profile_test.go @@ -19,3 +19,14 @@ func TestWithHTTPUpstreamProfile_OpenAI(t *testing.T) { t.Fatalf("expected profile %q, got %q", HTTPUpstreamProfileOpenAI, profile) } } + +func TestWithHTTPUpstreamRedirectsDisabled(t *testing.T) { + //nolint:staticcheck // Exercises the defensive nil-context fallback. + ctx := WithHTTPUpstreamRedirectsDisabled(nil) + if !HTTPUpstreamRedirectsDisabled(ctx) { + t.Fatal("expected redirects to be disabled") + } + if HTTPUpstreamRedirectsDisabled(context.Background()) { + t.Fatal("redirects should remain enabled by default") + } +} diff --git a/backend/internal/service/openai_endpoint_url.go b/backend/internal/service/openai_endpoint_url.go index a08af44cdf..5b0bf68b43 100644 --- a/backend/internal/service/openai_endpoint_url.go +++ b/backend/internal/service/openai_endpoint_url.go @@ -6,16 +6,25 @@ import ( ) func buildOpenAIEndpointURL(base string, endpoint string) string { - normalized := strings.TrimRight(strings.TrimSpace(base), "/") + normalized := strings.TrimSpace(base) endpoint = "/" + strings.TrimLeft(strings.TrimSpace(endpoint), "/") relative := strings.TrimPrefix(endpoint, "/v1") - if strings.HasSuffix(normalized, endpoint) || strings.HasSuffix(normalized, relative) { - return normalized + parsed, err := url.Parse(normalized) + if err != nil { + return strings.TrimRight(normalized, "/") + endpoint } - if openAIBaseURLHasVersionSuffix(normalized) { - return normalized + relative + path := strings.TrimRight(parsed.Path, "/") + if !strings.HasSuffix(path, endpoint) && !strings.HasSuffix(path, relative) { + if openAIBaseURLHasVersionSuffix(path) { + path += relative + } else { + path += endpoint + } } - return normalized + endpoint + parsed.Path = path + parsed.RawPath = "" + parsed.Fragment = "" + return parsed.String() } func buildOpenAIResponsesInputTokensURL(base string) string { diff --git a/backend/internal/service/openai_endpoint_url_test.go b/backend/internal/service/openai_endpoint_url_test.go new file mode 100644 index 0000000000..814db68103 --- /dev/null +++ b/backend/internal/service/openai_endpoint_url_test.go @@ -0,0 +1,31 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestBuildOpenAIEndpointURLPreservesURLComponents(t *testing.T) { + tests := []struct { + name string + base string + endpoint string + want string + }{ + {name: "root", base: "https://upstream.example", endpoint: "/v1/models", want: "https://upstream.example/v1/models"}, + {name: "v1", base: "https://upstream.example/v1", endpoint: "/v1/responses", want: "https://upstream.example/v1/responses"}, + {name: "prefix", base: "https://upstream.example/openai", endpoint: "/v1/chat/completions", want: "https://upstream.example/openai/v1/chat/completions"}, + {name: "version", base: "https://upstream.example/openai/v2", endpoint: "/v1/embeddings", want: "https://upstream.example/openai/v2/embeddings"}, + {name: "query", base: "https://upstream.example/v1?redirect=/", endpoint: "/v1/sub2api/billing", want: "https://upstream.example/v1/sub2api/billing?redirect=/"}, + {name: "fragment is removed", base: "https://upstream.example/v1#stale", endpoint: "/v1/alpha/search", want: "https://upstream.example/v1/alpha/search"}, + {name: "ipv6", base: "http://[2001:db8::1]:8080/v1?tenant=a#stale", endpoint: "/v1/responses/input_tokens", want: "http://[2001:db8::1]:8080/v1/responses/input_tokens?tenant=a"}, + {name: "already complete", base: "https://upstream.example/v1/images/generations?tenant=a", endpoint: "/v1/images/generations", want: "https://upstream.example/v1/images/generations?tenant=a"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, buildOpenAIEndpointURL(tt.base, tt.endpoint)) + }) + } +} diff --git a/backend/internal/service/ops_advisory_lock.go b/backend/internal/service/ops_advisory_lock.go index f7ef4ceec1..f020d1dc09 100644 --- a/backend/internal/service/ops_advisory_lock.go +++ b/backend/internal/service/ops_advisory_lock.go @@ -3,6 +3,7 @@ package service import ( "context" "database/sql" + "fmt" "hash/fnv" "time" ) @@ -14,8 +15,13 @@ func hashAdvisoryLockID(key string) int64 { } func tryAcquireDBAdvisoryLock(ctx context.Context, db *sql.DB, lockID int64) (func(), bool) { + release, acquired, _ := tryAcquireDBAdvisoryLockWithError(ctx, db, lockID) + return release, acquired +} + +func tryAcquireDBAdvisoryLockWithError(ctx context.Context, db *sql.DB, lockID int64) (func(), bool, error) { if db == nil { - return nil, false + return nil, false, nil } if ctx == nil { ctx = context.Background() @@ -23,17 +29,17 @@ func tryAcquireDBAdvisoryLock(ctx context.Context, db *sql.DB, lockID int64) (fu conn, err := db.Conn(ctx) if err != nil { - return nil, false + return nil, false, fmt.Errorf("open advisory-lock connection: %w", err) } acquired := false if err := conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", lockID).Scan(&acquired); err != nil { _ = conn.Close() - return nil, false + return nil, false, fmt.Errorf("query advisory lock: %w", err) } if !acquired { _ = conn.Close() - return nil, false + return nil, false, nil } release := func() { @@ -42,5 +48,5 @@ func tryAcquireDBAdvisoryLock(ctx context.Context, db *sql.DB, lockID int64) (fu _, _ = conn.ExecContext(unlockCtx, "SELECT pg_advisory_unlock($1)", lockID) _ = conn.Close() } - return release, true + return release, true, nil } diff --git a/backend/internal/service/proxy_update_probe_invalidation_test.go b/backend/internal/service/proxy_update_probe_invalidation_test.go new file mode 100644 index 0000000000..ac5215e0e9 --- /dev/null +++ b/backend/internal/service/proxy_update_probe_invalidation_test.go @@ -0,0 +1,71 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type updatingProxyRepoStub struct { + *proxyRepoStub + proxy *Proxy + updateCalls int +} + +func (s *updatingProxyRepoStub) GetByID(context.Context, int64) (*Proxy, error) { + copy := *s.proxy + return ©, nil +} + +func (s *updatingProxyRepoStub) Update(_ context.Context, proxy *Proxy) error { + s.updateCalls++ + copy := *proxy + s.proxy = © + return nil +} + +func TestBothProxyUpdateServicesUseRepositoryUpdateBoundary(t *testing.T) { + t.Run("ProxyService", func(t *testing.T) { + repo := &updatingProxyRepoStub{ + proxyRepoStub: &proxyRepoStub{}, + proxy: &Proxy{ID: 9, Protocol: "http", Host: "old.example", Port: 8080, Status: StatusActive}, + } + svc := NewProxyService(repo) + host := "new.example" + + _, err := svc.Update(context.Background(), 9, UpdateProxyRequest{Host: &host}) + + require.NoError(t, err) + require.Equal(t, 1, repo.updateCalls) + require.Equal(t, host, repo.proxy.Host) + }) + + t.Run("adminService", func(t *testing.T) { + repo := &updatingProxyRepoStub{ + proxyRepoStub: &proxyRepoStub{}, + proxy: &Proxy{ + ID: 9, + Protocol: "http", + Host: "old.example", + Port: 8080, + Status: StatusActive, + FallbackMode: FallbackModeNone, + ExpiryWarnDays: 7, + }, + } + svc := &adminServiceImpl{proxyRepo: repo} + + _, err := svc.UpdateProxy(context.Background(), 9, &UpdateProxyInput{ + Host: "new.example", + FallbackMode: FallbackModeNone, + ExpiryWarnDays: 7, + }) + + require.NoError(t, err) + require.Equal(t, 1, repo.updateCalls) + require.Equal(t, "new.example", repo.proxy.Host) + }) +} diff --git a/backend/internal/service/upstream_billing_probe.go b/backend/internal/service/upstream_billing_probe.go new file mode 100644 index 0000000000..70dbf8f8f2 --- /dev/null +++ b/backend/internal/service/upstream_billing_probe.go @@ -0,0 +1,926 @@ +package service + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "math/rand/v2" + "net/http" + "sort" + "strconv" + "strings" + "sync" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/google/uuid" + "golang.org/x/sync/errgroup" + "golang.org/x/sync/singleflight" +) + +const ( + // These values live in accounts.extra so PR2 does not require a schema migration. + UpstreamBillingProbeExtraKey = "upstream_billing_probe" + UpstreamBillingProbeEnabledExtraKey = "upstream_billing_probe_enabled" + + upstreamBillingProbeDefaultIntervalMinutes = 30 + upstreamBillingProbeMinIntervalMinutes = 5 + upstreamBillingProbeMaxIntervalMinutes = 24 * 60 + upstreamBillingProbeCycleInterval = time.Minute + upstreamBillingProbeRequestTimeout = 10 * time.Second + upstreamBillingProbeMaxBodyBytes = 64 * 1024 + upstreamBillingProbeMaxPerCycle = 20 + upstreamBillingProbeConcurrency = 4 + upstreamBillingProbeMaxBackoff = 24 * time.Hour + upstreamBillingProbeLeaderLockKey = "upstream:billing:probe:leader" + upstreamBillingProbeLeaderLockTTL = 2 * time.Minute +) + +// UpstreamBillingProbeMaxBatchSize limits one manual batch and one runner cycle. +const UpstreamBillingProbeMaxBatchSize = upstreamBillingProbeMaxPerCycle + +var ( + ErrUpstreamBillingProbeUnavailable = infraerrors.ServiceUnavailable( + "UPSTREAM_BILLING_PROBE_UNAVAILABLE", "upstream billing probe is unavailable", + ) + ErrUpstreamBillingProbeAccountInvalid = infraerrors.BadRequest( + "UPSTREAM_BILLING_PROBE_ACCOUNT_INVALID", "account is not an OpenAI API key account", + ) + ErrUpstreamBillingProbeIdentityChanged = infraerrors.Conflict( + "UPSTREAM_BILLING_PROBE_IDENTITY_CHANGED", "account identity changed during upstream billing probe; retry the probe", + ) +) + +const ( + UpstreamBillingProbeStatusOK = "ok" + UpstreamBillingProbeStatusUnsupported = "unsupported" + UpstreamBillingProbeStatusFailed = "failed" +) + +// UpstreamBillingProbeSettings controls the periodic probe runner. +type UpstreamBillingProbeSettings struct { + Enabled bool `json:"enabled"` + IntervalMinutes int `json:"interval_minutes"` +} + +// UpstreamBillingProbeSnapshot is persisted in accounts.extra. Data is kept as +// a sanitized map so future response fields do not require a database change. +type UpstreamBillingProbeSnapshot struct { + Status string `json:"status"` + Data map[string]any `json:"data,omitempty"` + ReceivedAt *time.Time `json:"received_at,omitempty"` + FreshUntil *time.Time `json:"fresh_until,omitempty"` + LastAttemptAt time.Time `json:"last_attempt_at"` + NextProbeAt time.Time `json:"next_probe_at"` + FailureCount int `json:"failure_count,omitempty"` + HTTPStatus int `json:"http_status,omitempty"` + LastError string `json:"last_error,omitempty"` +} + +// UpstreamBillingProbeResult is returned by manual probe endpoints. +type UpstreamBillingProbeResult struct { + AccountID int64 `json:"account_id"` + Snapshot *UpstreamBillingProbeSnapshot `json:"snapshot,omitempty"` + Error string `json:"error,omitempty"` +} + +type upstreamBillingProbeResponse struct { + Object string `json:"object"` + SchemaVersion int `json:"schema_version"` + BillingScope string `json:"billing_scope"` + GroupRateMultiplier *float64 `json:"group_rate_multiplier"` + UserRateMultiplier *float64 `json:"user_rate_multiplier"` + ResolvedRateMultiplier *float64 `json:"resolved_rate_multiplier"` + PeakRateEnabled *bool `json:"peak_rate_enabled"` + PeakStart *string `json:"peak_start"` + PeakEnd *string `json:"peak_end"` + PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + AppliedPeakMultiplier *float64 `json:"applied_peak_multiplier"` + EffectiveRateMultiplier *float64 `json:"effective_rate_multiplier"` + Timezone *string `json:"timezone"` + ObservedAt string `json:"observed_at"` +} + +// GetUpstreamBillingProbeSettings returns defaults when the setting is absent. +func (s *SettingService) GetUpstreamBillingProbeSettings(ctx context.Context) (*UpstreamBillingProbeSettings, error) { + defaults := defaultUpstreamBillingProbeSettings() + if s == nil || s.settingRepo == nil { + return defaults, nil + } + value, err := s.settingRepo.GetValue(ctx, SettingKeyUpstreamBillingProbeSettings) + if err != nil { + if errors.Is(err, ErrSettingNotFound) { + return defaults, nil + } + return nil, fmt.Errorf("get upstream billing probe settings: %w", err) + } + if strings.TrimSpace(value) == "" { + return defaults, nil + } + settings := *defaults + if err := json.Unmarshal([]byte(value), &settings); err != nil { + return nil, fmt.Errorf("parse upstream billing probe settings: %w", err) + } + if settings.IntervalMinutes == 0 { + settings.IntervalMinutes = defaults.IntervalMinutes + } + normalizeUpstreamBillingProbeSettings(&settings) + return &settings, nil +} + +// SetUpstreamBillingProbeSettings validates and persists the runner settings. +func (s *SettingService) SetUpstreamBillingProbeSettings(ctx context.Context, settings *UpstreamBillingProbeSettings) error { + if s == nil || s.settingRepo == nil { + return fmt.Errorf("setting repository is unavailable") + } + if settings == nil { + return infraerrors.BadRequest("INVALID_UPSTREAM_BILLING_PROBE_SETTINGS", "settings cannot be nil") + } + if settings.IntervalMinutes < upstreamBillingProbeMinIntervalMinutes || settings.IntervalMinutes > upstreamBillingProbeMaxIntervalMinutes { + return infraerrors.BadRequest( + "INVALID_UPSTREAM_BILLING_PROBE_INTERVAL", + fmt.Sprintf("interval_minutes must be between %d and %d", upstreamBillingProbeMinIntervalMinutes, upstreamBillingProbeMaxIntervalMinutes), + ) + } + normalizeUpstreamBillingProbeSettings(settings) + data, err := json.Marshal(settings) + if err != nil { + return fmt.Errorf("marshal upstream billing probe settings: %w", err) + } + return s.settingRepo.Set(ctx, SettingKeyUpstreamBillingProbeSettings, string(data)) +} + +func defaultUpstreamBillingProbeSettings() *UpstreamBillingProbeSettings { + return &UpstreamBillingProbeSettings{Enabled: true, IntervalMinutes: upstreamBillingProbeDefaultIntervalMinutes} +} + +func normalizeUpstreamBillingProbeSettings(settings *UpstreamBillingProbeSettings) { + if settings.IntervalMinutes < upstreamBillingProbeMinIntervalMinutes { + settings.IntervalMinutes = upstreamBillingProbeMinIntervalMinutes + } + if settings.IntervalMinutes > upstreamBillingProbeMaxIntervalMinutes { + settings.IntervalMinutes = upstreamBillingProbeMaxIntervalMinutes + } +} + +// UpstreamBillingProbeService discovers a remote Sub2API billing snapshot. +type UpstreamBillingProbeService struct { + accountRepo AccountRepository + accountTestService *AccountTestService + settingService *SettingService + + parentCtx context.Context + parentCancel context.CancelFunc + wg sync.WaitGroup + mu sync.Mutex + started bool + stopped bool + cycleMu sync.Mutex + probeGroup singleflight.Group + probeSlots chan struct{} + now func() time.Time + lockCache LeaderLockCache + db *sql.DB + instanceID string +} + +type upstreamBillingProbeSnapshotWriter interface { + UpdateUpstreamBillingProbeSnapshot(context.Context, *Account, *UpstreamBillingProbeSnapshot) error +} + +type upstreamBillingProbeDueAccountLister interface { + ListDueUpstreamBillingProbeAccounts(context.Context, time.Time, int) ([]Account, error) +} + +func NewUpstreamBillingProbeService( + accountRepo AccountRepository, + accountTestService *AccountTestService, + settingService *SettingService, +) *UpstreamBillingProbeService { + ctx, cancel := context.WithCancel(context.Background()) + return &UpstreamBillingProbeService{ + accountRepo: accountRepo, + accountTestService: accountTestService, + settingService: settingService, + parentCtx: ctx, + parentCancel: cancel, + probeSlots: make(chan struct{}, upstreamBillingProbeConcurrency), + now: time.Now, + instanceID: uuid.NewString(), + } +} + +func (s *UpstreamBillingProbeService) SetLeaderLock(lockCache LeaderLockCache, db *sql.DB) { + if s == nil { + return + } + s.lockCache = lockCache + s.db = db +} + +// ProvideUpstreamBillingProbeService starts the process-wide periodic runner. +func ProvideUpstreamBillingProbeService( + accountRepo AccountRepository, + accountTestService *AccountTestService, + settingService *SettingService, + lockCache LeaderLockCache, + db *sql.DB, +) *UpstreamBillingProbeService { + svc := NewUpstreamBillingProbeService(accountRepo, accountTestService, settingService) + svc.SetLeaderLock(lockCache, db) + svc.Start() + return svc +} + +func (s *UpstreamBillingProbeService) Start() { + if s == nil { + return + } + s.mu.Lock() + if s.started || s.stopped { + s.mu.Unlock() + return + } + s.started = true + s.wg.Add(1) + s.mu.Unlock() + go s.runLoop() +} + +func (s *UpstreamBillingProbeService) Stop() { + if s == nil { + return + } + s.mu.Lock() + if s.stopped { + s.mu.Unlock() + return + } + s.stopped = true + s.parentCancel() + s.mu.Unlock() + s.wg.Wait() +} + +func (s *UpstreamBillingProbeService) runLoop() { + defer s.wg.Done() + _ = s.RunDue(s.parentCtx) + ticker := time.NewTicker(upstreamBillingProbeCycleInterval) + defer ticker.Stop() + for { + select { + case <-s.parentCtx.Done(): + return + case <-ticker.C: + if err := s.RunDue(s.parentCtx); err != nil { + logger.LegacyPrintf("service.upstream_billing_probe", "run_due_failed: err=%v", err) + } + } + } +} + +// RunDue executes at most one bounded batch of due accounts. +func (s *UpstreamBillingProbeService) RunDue(ctx context.Context) error { + if s == nil || s.accountRepo == nil { + return nil + } + s.cycleMu.Lock() + defer s.cycleMu.Unlock() + + settings, err := s.getSettings(ctx) + if err != nil { + return err + } + if !settings.Enabled { + return nil + } + runRelease, acquired, lockErr := s.tryAcquireLeaderLock(ctx, upstreamBillingProbeLeaderLockKey) + if lockErr != nil { + return fmt.Errorf("acquire upstream billing probe leader lock: %w", lockErr) + } + if !acquired { + return nil + } + defer runRelease() + + lockNow := time.Now() + cadenceRelease, acquired, lockErr := s.tryAcquireLeaderLock(ctx, upstreamBillingProbeLeaderLockKeyAt(lockNow)) + if lockErr != nil { + return fmt.Errorf("acquire upstream billing probe cadence lock: %w", lockErr) + } + if !acquired { + return nil + } + defer releaseUpstreamBillingProbeLeaderLock(cadenceRelease, lockNow.Truncate(upstreamBillingProbeCycleInterval).Add(upstreamBillingProbeCycleInterval)) + + now := s.currentTime() + accounts, err := s.listDueAccounts(ctx, now) + if err != nil { + return fmt.Errorf("list enabled upstream billing probes: %w", err) + } + due := make([]Account, 0, len(accounts)) + for i := range accounts { + account := accounts[i] + if !isUpstreamBillingProbeAccount(&account) || !account.IsActive() || !upstreamBillingProbeEnabled(&account) { + continue + } + snapshot := decodeUpstreamBillingProbeSnapshot(account.Extra) + if snapshot != nil && !snapshot.NextProbeAt.IsZero() && now.Before(snapshot.NextProbeAt) { + continue + } + due = append(due, account) + } + sort.SliceStable(due, func(i, j int) bool { + left := decodeUpstreamBillingProbeSnapshot(due[i].Extra) + right := decodeUpstreamBillingProbeSnapshot(due[j].Extra) + leftUnset := left == nil || left.NextProbeAt.IsZero() + rightUnset := right == nil || right.NextProbeAt.IsZero() + if leftUnset && rightUnset { + return due[i].ID < due[j].ID + } + if leftUnset { + return true + } + if rightUnset { + return false + } + return left.NextProbeAt.Before(right.NextProbeAt) + }) + if len(due) > upstreamBillingProbeMaxPerCycle { + due = due[:upstreamBillingProbeMaxPerCycle] + } + + var group errgroup.Group + for i := range due { + accountID := due[i].ID + group.Go(func() error { + if _, probeErr := s.probeScheduledAccount(ctx, accountID, settings.IntervalMinutes); probeErr != nil { + logger.LegacyPrintf("service.upstream_billing_probe", "probe_due_failed: account_id=%d err=%v", accountID, probeErr) + } + return nil + }) + } + return group.Wait() +} + +func (s *UpstreamBillingProbeService) listDueAccounts(ctx context.Context, now time.Time) ([]Account, error) { + if lister, ok := s.accountRepo.(upstreamBillingProbeDueAccountLister); ok { + return lister.ListDueUpstreamBillingProbeAccounts(ctx, now, upstreamBillingProbeMaxPerCycle) + } + // Non-production repositories and older adapters keep the generic path. The + // runner still truncates before issuing network requests. + return s.accountRepo.FindByExtraField(ctx, UpstreamBillingProbeEnabledExtraKey, true) +} + +func (s *UpstreamBillingProbeService) getSettings(ctx context.Context) (*UpstreamBillingProbeSettings, error) { + if s.settingService == nil { + return defaultUpstreamBillingProbeSettings(), nil + } + return s.settingService.GetUpstreamBillingProbeSettings(ctx) +} + +func (s *UpstreamBillingProbeService) GetSettings(ctx context.Context) (*UpstreamBillingProbeSettings, error) { + return s.getSettings(ctx) +} + +func (s *UpstreamBillingProbeService) UpdateSettings(ctx context.Context, settings *UpstreamBillingProbeSettings) error { + if s == nil || s.settingService == nil { + return ErrUpstreamBillingProbeUnavailable + } + return s.settingService.SetUpstreamBillingProbeSettings(ctx, settings) +} + +// ProbeAccount performs one manual or scheduled probe. Manual calls ignore both switches. +func (s *UpstreamBillingProbeService) ProbeAccount(ctx context.Context, accountID int64) (*UpstreamBillingProbeSnapshot, error) { + if s == nil || s.accountRepo == nil { + return nil, ErrUpstreamBillingProbeUnavailable + } + settings, err := s.getSettings(ctx) + if err != nil { + return nil, err + } + return s.probeAccount(ctx, accountID, settings.IntervalMinutes) +} + +func (s *UpstreamBillingProbeService) probeAccount(ctx context.Context, accountID int64, intervalMinutes int) (*UpstreamBillingProbeSnapshot, error) { + return s.probeAccountWithMode(ctx, accountID, intervalMinutes, false) +} + +func (s *UpstreamBillingProbeService) probeScheduledAccount(ctx context.Context, accountID int64, intervalMinutes int) (*UpstreamBillingProbeSnapshot, error) { + return s.probeAccountWithMode(ctx, accountID, intervalMinutes, true) +} + +func (s *UpstreamBillingProbeService) probeAccountWithMode(ctx context.Context, accountID int64, intervalMinutes int, requireEnabled bool) (*UpstreamBillingProbeSnapshot, error) { + key := strconv.FormatInt(accountID, 10) + value, err, _ := s.probeGroup.Do(key, func() (any, error) { + select { + case s.probeSlots <- struct{}{}: + defer func() { <-s.probeSlots }() + case <-ctx.Done(): + return nil, ctx.Err() + } + account, loadErr := s.accountRepo.GetByID(ctx, accountID) + if loadErr != nil { + return nil, loadErr + } + if !isUpstreamBillingProbeAccount(account) { + return nil, ErrUpstreamBillingProbeAccountInvalid + } + if requireEnabled { + if !account.IsActive() || !upstreamBillingProbeEnabled(account) { + return nil, nil + } + if snapshot := decodeUpstreamBillingProbeSnapshot(account.Extra); snapshot != nil && + !snapshot.NextProbeAt.IsZero() && s.currentTime().Before(snapshot.NextProbeAt) { + return nil, nil + } + } + return s.probeLoadedAccount(ctx, account, intervalMinutes) + }) + if err != nil { + return nil, err + } + if value == nil { + return nil, nil + } + snapshot, ok := value.(*UpstreamBillingProbeSnapshot) + if !ok { + return nil, fmt.Errorf("invalid upstream billing probe result") + } + return snapshot, nil +} + +// ProbeAccounts performs a bounded manual batch with the same concurrency limit as the runner. +func (s *UpstreamBillingProbeService) ProbeAccounts(ctx context.Context, accountIDs []int64) []UpstreamBillingProbeResult { + if len(accountIDs) > upstreamBillingProbeMaxPerCycle { + accountIDs = accountIDs[:upstreamBillingProbeMaxPerCycle] + } + results := make([]UpstreamBillingProbeResult, len(accountIDs)) + if s == nil || s.accountRepo == nil { + for i, accountID := range accountIDs { + results[i] = UpstreamBillingProbeResult{AccountID: accountID, Error: ErrUpstreamBillingProbeUnavailable.Error()} + } + return results + } + settings, settingsErr := s.getSettings(ctx) + if settingsErr != nil { + for i, accountID := range accountIDs { + results[i] = UpstreamBillingProbeResult{AccountID: accountID, Error: safeProbeError(settingsErr)} + } + return results + } + var group errgroup.Group + for i, accountID := range accountIDs { + i, accountID := i, accountID + results[i].AccountID = accountID + group.Go(func() error { + snapshot, err := s.probeAccount(ctx, accountID, settings.IntervalMinutes) + if err != nil { + results[i].Error = safeProbeError(err) + return nil + } + results[i].Snapshot = snapshot + return nil + }) + } + _ = group.Wait() + return results +} + +func upstreamBillingProbeLeaderLockKeyAt(now time.Time) string { + return fmt.Sprintf("%s:%d", upstreamBillingProbeLeaderLockKey, now.Unix()/int64(upstreamBillingProbeCycleInterval/time.Second)) +} + +func (s *UpstreamBillingProbeService) tryAcquireLeaderLock(ctx context.Context, key string) (func(), bool, error) { + lockCtx, cancel := context.WithTimeout(ctx, 2*time.Second) + defer cancel() + if s.lockCache != nil { + acquired, err := s.lockCache.TryAcquireLeaderLock(lockCtx, key, s.instanceID, upstreamBillingProbeLeaderLockTTL) + if err != nil { + return nil, false, err + } + if !acquired { + return nil, false, nil + } + return func() { + releaseCtx, releaseCancel := context.WithTimeout(context.Background(), 2*time.Second) + defer releaseCancel() + _ = s.lockCache.ReleaseLeaderLock(releaseCtx, key, s.instanceID) + }, true, nil + } + if s.db != nil { + return tryAcquireDBAdvisoryLockWithError(lockCtx, s.db, hashAdvisoryLockID(key)) + } + return func() {}, true, nil +} + +func releaseUpstreamBillingProbeLeaderLock(release func(), releaseAt time.Time) { + delay := time.Until(releaseAt) + if delay <= 0 { + release() + return + } + time.AfterFunc(delay, release) +} + +func (s *UpstreamBillingProbeService) SetAccountEnabled(ctx context.Context, accountID int64, enabled bool) error { + if s == nil || s.accountRepo == nil { + return ErrUpstreamBillingProbeUnavailable + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return err + } + if !isUpstreamBillingProbeAccount(account) { + return ErrUpstreamBillingProbeAccountInvalid + } + return s.accountRepo.UpdateExtra(ctx, accountID, map[string]any{ + UpstreamBillingProbeEnabledExtraKey: enabled, + }) +} + +func (s *UpstreamBillingProbeService) probeLoadedAccount(ctx context.Context, account *Account, intervalMinutes int) (*UpstreamBillingProbeSnapshot, error) { + now := s.currentTime().UTC() + if s.accountTestService == nil || s.accountTestService.httpUpstream == nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "transport_unavailable", 0) + } + apiKey := account.GetOpenAIApiKey() + if apiKey == "" { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "missing_api_key", 0) + } + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + baseURL = "https://api.openai.com" + } + normalizedBaseURL, err := s.accountTestService.validateUpstreamBaseURL(baseURL) + if err != nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "invalid_base_url", 0) + } + proxyURL := "" + if account.ProxyID != nil { + if account.Proxy == nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "proxy_unavailable", 0) + } + if account.Proxy.ID != *account.ProxyID { + return nil, ErrUpstreamBillingProbeIdentityChanged + } + proxyURL = account.Proxy.URL() + } + probeURL := buildOpenAIEndpointURL(normalizedBaseURL, "/v1/sub2api/billing") + probeCtx, cancel := context.WithTimeout(ctx, upstreamBillingProbeRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(probeCtx, http.MethodGet, probeURL, bytes.NewReader(nil)) + if err != nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "request_build_failed", 0) + } + reqCtx := WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI) + req = req.WithContext(WithHTTPUpstreamRedirectsDisabled(reqCtx)) + req.Header.Set("Accept", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + account.ApplyHeaderOverrides(req.Header) + var tlsProfile *tlsfingerprint.Profile + if s.accountTestService.tlsFPProfileService != nil { + tlsProfile = s.accountTestService.tlsFPProfileService.ResolveTLSProfile(account) + } + resp, err := s.accountTestService.httpUpstream.DoWithTLS(req, proxyURL, account.ID, account.Concurrency, tlsProfile) + if err != nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "request_failed", 0) + } + if resp == nil || resp.Body == nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, 0, "empty_response", 0) + } + defer func() { _ = resp.Body.Close() }() + body, readErr := io.ReadAll(io.LimitReader(resp.Body, upstreamBillingProbeMaxBodyBytes+1)) + if readErr != nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "response_read_failed", retryAfter(resp.Header, now)) + } + if len(body) > upstreamBillingProbeMaxBodyBytes { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "response_too_large", retryAfter(resp.Header, now)) + } + if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "unsupported", retryAfter(resp.Header, now)) + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "http_error", retryAfter(resp.Header, now)) + } + data, err := parseUpstreamBillingProbeResponse(body) + if err != nil { + return s.persistProbeFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "invalid_response", retryAfter(resp.Header, now)) + } + snapshot := &UpstreamBillingProbeSnapshot{ + Status: UpstreamBillingProbeStatusOK, + Data: data, + ReceivedAt: probeTimePtr(now), + FreshUntil: probeTimePtr(now.Add(2 * time.Duration(intervalMinutes) * time.Minute)), + LastAttemptAt: now, + NextProbeAt: now.Add(nextProbeDelay(intervalMinutes, 0, 0)), + HTTPStatus: resp.StatusCode, + } + if err := s.updateSnapshot(ctx, account, snapshot); err != nil { + return nil, err + } + return snapshot, nil +} + +func (s *UpstreamBillingProbeService) persistProbeFailure( + ctx context.Context, + account *Account, + intervalMinutes int, + now time.Time, + statusCode int, + reason string, + retryAfterDuration time.Duration, +) (*UpstreamBillingProbeSnapshot, error) { + previous := decodeUpstreamBillingProbeSnapshot(account.Extra) + failureCount := 1 + if previous != nil { + failureCount = previous.FailureCount + 1 + } + status := UpstreamBillingProbeStatusFailed + if reason == "unsupported" { + status = UpstreamBillingProbeStatusUnsupported + } + snapshot := &UpstreamBillingProbeSnapshot{ + Status: status, + LastAttemptAt: now, + NextProbeAt: now.Add(nextProbeDelay(intervalMinutes, failureCount, retryAfterDuration)), + FailureCount: failureCount, + HTTPStatus: statusCode, + LastError: reason, + } + if previous != nil { + snapshot.Data = previous.Data + snapshot.ReceivedAt = previous.ReceivedAt + snapshot.FreshUntil = previous.FreshUntil + if snapshot.FreshUntil == nil && previous.Status == UpstreamBillingProbeStatusOK && previous.ReceivedAt != nil { + snapshot.FreshUntil = probeTimePtr(previous.ReceivedAt.Add(2 * time.Duration(intervalMinutes) * time.Minute)) + } + } + if err := s.updateSnapshot(ctx, account, snapshot); err != nil { + return nil, err + } + return snapshot, nil +} + +func (s *UpstreamBillingProbeService) updateSnapshot(ctx context.Context, account *Account, snapshot *UpstreamBillingProbeSnapshot) error { + writer, ok := s.accountRepo.(upstreamBillingProbeSnapshotWriter) + if !ok { + return ErrUpstreamBillingProbeUnavailable + } + return writer.UpdateUpstreamBillingProbeSnapshot(ctx, account, snapshot) +} + +func parseUpstreamBillingProbeResponse(body []byte) (map[string]any, error) { + var response upstreamBillingProbeResponse + if err := json.Unmarshal(body, &response); err != nil { + return nil, err + } + if response.Object != "sub2api.key_billing" || response.SchemaVersion != 1 || response.BillingScope != "token" { + return nil, fmt.Errorf("unexpected billing response schema") + } + if response.GroupRateMultiplier == nil || response.ResolvedRateMultiplier == nil || + response.PeakRateEnabled == nil || response.EffectiveRateMultiplier == nil { + return nil, fmt.Errorf("incomplete billing response") + } + for _, value := range []float64{ + *response.GroupRateMultiplier, + *response.ResolvedRateMultiplier, + *response.EffectiveRateMultiplier, + } { + if value < 0 || math.IsNaN(value) || math.IsInf(value, 0) { + return nil, fmt.Errorf("invalid billing multiplier") + } + } + if response.UserRateMultiplier != nil && (*response.UserRateMultiplier < 0 || math.IsNaN(*response.UserRateMultiplier) || math.IsInf(*response.UserRateMultiplier, 0)) { + return nil, fmt.Errorf("invalid user billing multiplier") + } + expectedResolved := *response.GroupRateMultiplier + if response.UserRateMultiplier != nil { + expectedResolved = *response.UserRateMultiplier + } + if !equalBillingMultiplier(*response.ResolvedRateMultiplier, expectedResolved) { + return nil, fmt.Errorf("inconsistent resolved billing multiplier") + } + observedAt, err := time.Parse(time.RFC3339Nano, response.ObservedAt) + if err != nil || observedAt.IsZero() { + return nil, fmt.Errorf("invalid observed_at") + } + data := map[string]any{ + "object": response.Object, + "schema_version": response.SchemaVersion, + "billing_scope": response.BillingScope, + "group_rate_multiplier": *response.GroupRateMultiplier, + "resolved_rate_multiplier": *response.ResolvedRateMultiplier, + "peak_rate_enabled": *response.PeakRateEnabled, + "effective_rate_multiplier": *response.EffectiveRateMultiplier, + "observed_at": observedAt.UTC().Format(time.RFC3339Nano), + } + if response.UserRateMultiplier != nil { + data["user_rate_multiplier"] = *response.UserRateMultiplier + } + if *response.PeakRateEnabled { + if response.PeakStart == nil || response.PeakEnd == nil || response.Timezone == nil || + response.PeakRateMultiplier == nil || response.AppliedPeakMultiplier == nil || + *response.PeakStart == "" || *response.PeakEnd == "" || *response.Timezone == "" || + *response.PeakRateMultiplier < 0 || *response.AppliedPeakMultiplier < 0 || + math.IsNaN(*response.PeakRateMultiplier) || math.IsInf(*response.PeakRateMultiplier, 0) || + math.IsNaN(*response.AppliedPeakMultiplier) || math.IsInf(*response.AppliedPeakMultiplier, 0) { + return nil, fmt.Errorf("incomplete peak billing response") + } + data["peak_start"] = *response.PeakStart + data["peak_end"] = *response.PeakEnd + data["peak_rate_multiplier"] = *response.PeakRateMultiplier + data["applied_peak_multiplier"] = *response.AppliedPeakMultiplier + data["timezone"] = *response.Timezone + } + appliedPeak, ok := upstreamBillingPeakMultiplierAt(data, observedAt) + if !ok { + return nil, fmt.Errorf("invalid peak billing response") + } + if response.PeakRateEnabled != nil && *response.PeakRateEnabled { + if !equalBillingMultiplier(*response.AppliedPeakMultiplier, appliedPeak) { + return nil, fmt.Errorf("inconsistent applied peak multiplier") + } + } else if response.AppliedPeakMultiplier != nil && !equalBillingMultiplier(*response.AppliedPeakMultiplier, 1) { + return nil, fmt.Errorf("inconsistent applied peak multiplier") + } + if !equalBillingMultiplier(*response.EffectiveRateMultiplier, *response.ResolvedRateMultiplier*appliedPeak) { + return nil, fmt.Errorf("inconsistent effective billing multiplier") + } + return data, nil +} + +func upstreamBillingRateAt(data map[string]any, now time.Time) (float64, bool) { + if scope, _ := data["billing_scope"].(string); scope != "token" { + return 0, false + } + base, ok := resolveAccountExtraNumber(data, "resolved_rate_multiplier") + if !ok || base < 0 || math.IsNaN(base) || math.IsInf(base, 0) { + return 0, false + } + appliedPeak, ok := upstreamBillingPeakMultiplierAt(data, now) + if !ok { + return 0, false + } + base *= appliedPeak + if math.IsNaN(base) || math.IsInf(base, 0) { + return 0, false + } + return base, true +} + +func upstreamBillingPeakMultiplierAt(data map[string]any, now time.Time) (float64, bool) { + peakEnabled, ok := data["peak_rate_enabled"].(bool) + if !ok { + return 0, false + } + if !peakEnabled { + return 1, true + } + + start, startOK := data["peak_start"].(string) + end, endOK := data["peak_end"].(string) + timezoneName, timezoneOK := data["timezone"].(string) + peakMultiplier, multiplierOK := resolveAccountExtraNumber(data, "peak_rate_multiplier") + startMinute, validStart := parseMinutes(start) + endMinute, validEnd := parseMinutes(end) + if !startOK || !endOK || !timezoneOK || !multiplierOK || !validStart || !validEnd || + startMinute >= endMinute || peakMultiplier < 0 || math.IsNaN(peakMultiplier) || math.IsInf(peakMultiplier, 0) { + return 0, false + } + location, err := time.LoadLocation(timezoneName) + if err != nil { + return 0, false + } + + local := now.In(location) + minute := local.Hour()*60 + local.Minute() + if minute >= startMinute && minute < endMinute { + return peakMultiplier, true + } + return 1, true +} + +func equalBillingMultiplier(left, right float64) bool { + if math.IsNaN(left) || math.IsNaN(right) || math.IsInf(left, 0) || math.IsInf(right, 0) { + return false + } + scale := math.Max(1, math.Max(math.Abs(left), math.Abs(right))) + return math.Abs(left-right) <= 1e-9*scale +} + +func decodeUpstreamBillingProbeSnapshot(extra map[string]any) *UpstreamBillingProbeSnapshot { + if extra == nil { + return nil + } + value, ok := extra[UpstreamBillingProbeExtraKey] + if !ok { + return nil + } + raw, err := json.Marshal(value) + if err != nil { + return nil + } + var snapshot UpstreamBillingProbeSnapshot + if err := json.Unmarshal(raw, &snapshot); err != nil || snapshot.Status == "" { + return nil + } + if snapshot.Status != UpstreamBillingProbeStatusOK && + snapshot.Status != UpstreamBillingProbeStatusUnsupported && + snapshot.Status != UpstreamBillingProbeStatusFailed { + return nil + } + return &snapshot +} + +func isUpstreamBillingProbeAccount(account *Account) bool { + return account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey +} + +func upstreamBillingProbeEnabled(account *Account) bool { + if account == nil || account.Extra == nil { + return false + } + enabled, ok := account.Extra[UpstreamBillingProbeEnabledExtraKey].(bool) + return ok && enabled +} + +func (s *UpstreamBillingProbeService) currentTime() time.Time { + if s != nil && s.now != nil { + return s.now() + } + return time.Now() +} + +func nextProbeDelay(intervalMinutes, failureCount int, retryAfterDuration time.Duration) time.Duration { + interval := time.Duration(intervalMinutes) * time.Minute + if interval < upstreamBillingProbeMinIntervalMinutes*time.Minute { + interval = upstreamBillingProbeMinIntervalMinutes * time.Minute + } + if failureCount > 0 { + shift := failureCount + if shift > 5 { + shift = 5 + } + interval *= time.Duration(1 << shift) + } + if interval > upstreamBillingProbeMaxBackoff { + interval = upstreamBillingProbeMaxBackoff + } + jitterRange := interval / 5 + if jitterRange > 5*time.Minute { + jitterRange = 5 * time.Minute + } + if jitterRange > 0 { + interval += time.Duration(rand.Int64N(int64(jitterRange)*2+1)) - jitterRange + } + if retryAfterDuration > interval { + // Retry-After is an explicit upstream instruction; do not shorten it + // with the local exponential-backoff ceiling. + return retryAfterDuration + } + if interval > upstreamBillingProbeMaxBackoff { + return upstreamBillingProbeMaxBackoff + } + return interval +} + +func retryAfter(header http.Header, now time.Time) time.Duration { + value := strings.TrimSpace(header.Get("Retry-After")) + if value == "" { + return 0 + } + if seconds, err := strconv.Atoi(value); err == nil && seconds > 0 { + return time.Duration(seconds) * time.Second + } + if at, err := http.ParseTime(value); err == nil { + if delay := at.Sub(now); delay > 0 { + return delay + } + } + return 0 +} + +func probeTimePtr(value time.Time) *time.Time { + return &value +} + +func safeProbeError(err error) string { + if err == nil { + return "" + } + if errors.Is(err, ErrUpstreamBillingProbeAccountInvalid) { + return ErrUpstreamBillingProbeAccountInvalid.Error() + } + if errors.Is(err, ErrUpstreamBillingProbeUnavailable) { + return ErrUpstreamBillingProbeUnavailable.Error() + } + return "probe_failed" +} diff --git a/backend/internal/service/upstream_billing_probe_test.go b/backend/internal/service/upstream_billing_probe_test.go new file mode 100644 index 0000000000..0ffcc6006e --- /dev/null +++ b/backend/internal/service/upstream_billing_probe_test.go @@ -0,0 +1,939 @@ +package service + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/stretchr/testify/require" +) + +type upstreamBillingProbeAccountRepo struct { + AccountRepository + mu sync.Mutex + accounts map[int64]*Account + updates map[int64][]map[string]any + bulkUpdates []AccountBulkUpdate +} + +type staleDueUpstreamBillingProbeAccountRepo struct { + *upstreamBillingProbeAccountRepo + due []Account +} + +func (r *staleDueUpstreamBillingProbeAccountRepo) ListDueUpstreamBillingProbeAccounts(_ context.Context, _ time.Time, limit int) ([]Account, error) { + if limit < len(r.due) { + return append([]Account(nil), r.due[:limit]...), nil + } + return append([]Account(nil), r.due...), nil +} + +func (r *upstreamBillingProbeAccountRepo) Create(_ context.Context, account *Account) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.accounts == nil { + r.accounts = make(map[int64]*Account) + } + if account.ID == 0 { + account.ID = int64(len(r.accounts) + 1) + } + r.accounts[account.ID] = account + return nil +} + +func (r *upstreamBillingProbeAccountRepo) Update(_ context.Context, account *Account) error { + r.mu.Lock() + defer r.mu.Unlock() + r.accounts[account.ID] = account + return nil +} + +func (r *upstreamBillingProbeAccountRepo) BulkUpdate(_ context.Context, ids []int64, updates AccountBulkUpdate) (int64, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.bulkUpdates = append(r.bulkUpdates, updates) + return int64(len(ids)), nil +} + +func (r *upstreamBillingProbeAccountRepo) GetByID(_ context.Context, id int64) (*Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + account := r.accounts[id] + if account == nil { + return nil, ErrAccountNotFound + } + clone := *account + clone.Credentials = mergeMap(nil, account.Credentials) + clone.Extra = mergeMap(nil, account.Extra) + return &clone, nil +} + +func (r *upstreamBillingProbeAccountRepo) GetByIDs(_ context.Context, ids []int64) ([]*Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + result := make([]*Account, 0, len(ids)) + for _, id := range ids { + if account := r.accounts[id]; account != nil { + result = append(result, account) + } + } + return result, nil +} + +func (r *upstreamBillingProbeAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + r.mu.Lock() + defer r.mu.Unlock() + account := r.accounts[id] + if account == nil { + return ErrAccountNotFound + } + if account.Extra == nil { + account.Extra = make(map[string]any) + } + for key, value := range updates { + account.Extra[key] = value + } + if r.updates == nil { + r.updates = make(map[int64][]map[string]any) + } + r.updates[id] = append(r.updates[id], updates) + return nil +} + +func (r *upstreamBillingProbeAccountRepo) UpdateUpstreamBillingProbeSnapshot(_ context.Context, expected *Account, snapshot *UpstreamBillingProbeSnapshot) error { + r.mu.Lock() + defer r.mu.Unlock() + account := r.accounts[expected.ID] + if account == nil || account.Platform != expected.Platform || account.Type != expected.Type || !reflect.DeepEqual(account.Credentials, expected.Credentials) { + return ErrUpstreamBillingProbeIdentityChanged + } + if account.Extra == nil { + account.Extra = make(map[string]any) + } + account.Extra[UpstreamBillingProbeExtraKey] = snapshot + return nil +} + +func (r *upstreamBillingProbeAccountRepo) FindByExtraField(_ context.Context, key string, value any) ([]Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + result := make([]Account, 0) + for _, account := range r.accounts { + if account.Extra != nil && account.Extra[key] == value { + result = append(result, *account) + } + } + return result, nil +} + +type upstreamBillingProbeSettingRepo struct { + SettingRepository + mu sync.Mutex + values map[string]string +} + +type upstreamBillingProbeHTTPStub struct { + calls atomic.Int64 + active atomic.Int64 + maxActive atomic.Int64 + beforeResponse func() +} + +func (u *upstreamBillingProbeHTTPStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + u.calls.Add(1) + active := u.active.Add(1) + defer u.active.Add(-1) + for { + peak := u.maxActive.Load() + if active <= peak || u.maxActive.CompareAndSwap(peak, active) { + break + } + } + if u.beforeResponse != nil { + u.beforeResponse() + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "object":"sub2api.key_billing", + "schema_version":1, + "billing_scope":"token", + "group_rate_multiplier":0.8, + "resolved_rate_multiplier":0.8, + "peak_rate_enabled":false, + "effective_rate_multiplier":0.8, + "observed_at":"2026-07-13T01:00:00Z" + }`)), + }, nil +} + +func (u *upstreamBillingProbeHTTPStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, profile *tlsfingerprint.Profile) (*http.Response, error) { + return u.Do(req, proxyURL, accountID, accountConcurrency) +} + +func (r *upstreamBillingProbeSettingRepo) GetValue(_ context.Context, key string) (string, error) { + r.mu.Lock() + defer r.mu.Unlock() + value, ok := r.values[key] + if !ok { + return "", ErrSettingNotFound + } + return value, nil +} + +func (r *upstreamBillingProbeSettingRepo) Set(_ context.Context, key, value string) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.values == nil { + r.values = make(map[string]string) + } + r.values[key] = value + return nil +} + +func newUpstreamBillingProbeTestService( + repo AccountRepository, + upstream HTTPUpstream, + settingRepo SettingRepository, +) *UpstreamBillingProbeService { + cfg := &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{ + Enabled: false, + AllowInsecureHTTP: true, + }}} + accountTestService := &AccountTestService{accountRepo: repo, httpUpstream: upstream, cfg: cfg} + return NewUpstreamBillingProbeService(repo, accountTestService, NewSettingService(settingRepo, cfg)) +} + +func TestUpstreamBillingProbeSettingsDefaultsAndValidation(t *testing.T) { + repo := &upstreamBillingProbeSettingRepo{} + settingsService := NewSettingService(repo, &config.Config{}) + + settings, err := settingsService.GetUpstreamBillingProbeSettings(context.Background()) + require.NoError(t, err) + require.True(t, settings.Enabled) + require.Equal(t, 30, settings.IntervalMinutes) + + err = settingsService.SetUpstreamBillingProbeSettings(context.Background(), &UpstreamBillingProbeSettings{ + Enabled: false, + IntervalMinutes: 4, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "interval_minutes must be between 5 and 1440") + + err = settingsService.SetUpstreamBillingProbeSettings(context.Background(), &UpstreamBillingProbeSettings{ + Enabled: false, + IntervalMinutes: 60, + }) + require.NoError(t, err) + settings, err = settingsService.GetUpstreamBillingProbeSettings(context.Background()) + require.NoError(t, err) + require.False(t, settings.Enabled) + require.Equal(t, 60, settings.IntervalMinutes) + + repo.values[SettingKeyUpstreamBillingProbeSettings] = `{"interval_minutes":45}` + settings, err = settingsService.GetUpstreamBillingProbeSettings(context.Background()) + require.NoError(t, err) + require.True(t, settings.Enabled) + require.Equal(t, 45, settings.IntervalMinutes) + repo.values[SettingKeyUpstreamBillingProbeSettings] = `{"enabled":false}` + settings, err = settingsService.GetUpstreamBillingProbeSettings(context.Background()) + require.NoError(t, err) + require.False(t, settings.Enabled) + require.Equal(t, 30, settings.IntervalMinutes) + + repo.values[SettingKeyUpstreamBillingProbeSettings] = `{"enabled":` + settings, err = settingsService.GetUpstreamBillingProbeSettings(context.Background()) + require.ErrorContains(t, err, "parse upstream billing probe settings") + require.Nil(t, settings) +} + +func TestUpstreamBillingProbeSuccessPersistsSanitizedSnapshot(t *testing.T) { + account := &Account{ + ID: 17, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 2, + Credentials: map[string]any{ + "api_key": "sk-sensitive", + "base_url": "https://upstream.example/v1", + }, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{ + "object":"sub2api.key_billing", + "schema_version":1, + "billing_scope":"token", + "group_rate_multiplier":0.8, + "user_rate_multiplier":0.6, + "resolved_rate_multiplier":0.6, + "peak_rate_enabled":true, + "peak_start":"09:00", + "peak_end":"18:00", + "peak_rate_multiplier":1.5, + "applied_peak_multiplier":1.5, + "effective_rate_multiplier":0.9, + "timezone":"Asia/Shanghai", + "observed_at":"2026-07-13T01:00:00Z", + "unexpected_secret":"must-not-persist" + }`)), + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + fixedNow := time.Date(2026, time.July, 13, 2, 0, 0, 0, time.UTC) + svc.now = func() time.Time { return fixedNow } + + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusOK, snapshot.Status) + require.Equal(t, 0.9, snapshot.Data["effective_rate_multiplier"]) + require.NotContains(t, snapshot.Data, "unexpected_secret") + require.NotNil(t, snapshot.ReceivedAt) + require.Equal(t, fixedNow, *snapshot.ReceivedAt) + require.NotNil(t, snapshot.FreshUntil) + require.Equal(t, fixedNow.Add(time.Hour), *snapshot.FreshUntil) + require.False(t, snapshot.NextProbeAt.Before(fixedNow.Add(24*time.Minute))) + require.False(t, snapshot.NextProbeAt.After(fixedNow.Add(36*time.Minute))) + require.Equal(t, "https://upstream.example/v1/sub2api/billing", upstream.lastReq.URL.String()) + require.Equal(t, http.MethodGet, upstream.lastReq.Method) + require.Equal(t, "Bearer sk-sensitive", upstream.lastReq.Header.Get("Authorization")) + require.True(t, HTTPUpstreamRedirectsDisabled(upstream.lastReq.Context())) + + persisted := decodeUpstreamBillingProbeSnapshot(account.Extra) + require.NotNil(t, persisted) + require.Equal(t, snapshot.Status, persisted.Status) +} + +func TestUpstreamBillingProbeRejectsMissingRequiredMultiplier(t *testing.T) { + _, err := parseUpstreamBillingProbeResponse([]byte(`{ + "object":"sub2api.key_billing", + "schema_version":1, + "billing_scope":"token", + "group_rate_multiplier":0.8, + "peak_rate_enabled":false, + "effective_rate_multiplier":0.8, + "observed_at":"2026-07-13T01:00:00Z" + }`)) + + require.ErrorContains(t, err, "incomplete billing response") +} + +func TestUpstreamBillingProbeDiscardsResultWhenIdentityChangesInFlight(t *testing.T) { + account := &Account{ + ID: 19, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-old", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &upstreamBillingProbeHTTPStub{beforeResponse: func() { + repo.mu.Lock() + defer repo.mu.Unlock() + repo.accounts[account.ID].Credentials = map[string]any{"api_key": "sk-new", "base_url": "https://new.example"} + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + + require.Nil(t, snapshot) + require.ErrorIs(t, err, ErrUpstreamBillingProbeIdentityChanged) + require.NotContains(t, repo.accounts[account.ID].Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpstreamBillingProbeRejectsInvalidPeakConfiguration(t *testing.T) { + tests := []struct { + name string + start string + end string + timezone string + }{ + {name: "invalid start", start: "25:00", end: "18:00", timezone: "UTC"}, + {name: "cross midnight", start: "22:00", end: "02:00", timezone: "UTC"}, + {name: "invalid timezone", start: "09:00", end: "18:00", timezone: "Mars/Olympus"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := fmt.Sprintf(`{ + "object":"sub2api.key_billing", + "schema_version":1, + "billing_scope":"token", + "group_rate_multiplier":0.8, + "resolved_rate_multiplier":0.8, + "peak_rate_enabled":true, + "peak_start":%q, + "peak_end":%q, + "peak_rate_multiplier":1.5, + "applied_peak_multiplier":1, + "effective_rate_multiplier":0.8, + "timezone":%q, + "observed_at":"2026-07-13T01:00:00Z" + }`, tt.start, tt.end, tt.timezone) + + _, err := parseUpstreamBillingProbeResponse([]byte(body)) + require.ErrorContains(t, err, "invalid peak billing response") + }) + } +} + +func TestUpstreamBillingProbeRejectsInconsistentMultipliers(t *testing.T) { + tests := []struct { + name string + body string + }{ + { + name: "resolved does not use user override", + body: `{ + "object":"sub2api.key_billing","schema_version":1,"billing_scope":"token", + "group_rate_multiplier":0.8,"user_rate_multiplier":0.5,"resolved_rate_multiplier":0.8, + "peak_rate_enabled":false,"effective_rate_multiplier":0.8,"observed_at":"2026-07-13T01:00:00Z" + }`, + }, + { + name: "effective rate does not match resolved rate", + body: `{ + "object":"sub2api.key_billing","schema_version":1,"billing_scope":"token", + "group_rate_multiplier":0.8,"resolved_rate_multiplier":0.8, + "peak_rate_enabled":false,"effective_rate_multiplier":1.2,"observed_at":"2026-07-13T01:00:00Z" + }`, + }, + { + name: "applied peak does not match observed window", + body: `{ + "object":"sub2api.key_billing","schema_version":1,"billing_scope":"token", + "group_rate_multiplier":0.8,"resolved_rate_multiplier":0.8, + "peak_rate_enabled":true,"peak_start":"09:00","peak_end":"18:00", + "peak_rate_multiplier":1.5,"applied_peak_multiplier":1, + "effective_rate_multiplier":0.8,"timezone":"Asia/Shanghai","observed_at":"2026-07-13T01:00:00Z" + }`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := parseUpstreamBillingProbeResponse([]byte(tt.body)) + require.ErrorContains(t, err, "inconsistent") + }) + } +} + +func TestUpstreamBillingRateAtHandlesDST(t *testing.T) { + data := map[string]any{ + "billing_scope": "token", + "resolved_rate_multiplier": 1.0, + "peak_rate_enabled": true, + "peak_start": "02:00", + "peak_end": "04:00", + "peak_rate_multiplier": 2.0, + "timezone": "America/New_York", + } + beforeJump := time.Date(2026, time.March, 8, 6, 30, 0, 0, time.UTC) + afterJump := time.Date(2026, time.March, 8, 7, 30, 0, 0, time.UTC) + + rate, ok := upstreamBillingRateAt(data, beforeJump) + require.True(t, ok) + require.Equal(t, 1.0, rate) + rate, ok = upstreamBillingRateAt(data, afterJump) + require.True(t, ok) + require.Equal(t, 2.0, rate) +} + +func TestUpstreamBillingProbeFailurePreservesLastSuccessAndRetryAfter(t *testing.T) { + receivedAt := time.Date(2026, time.July, 12, 12, 0, 0, 0, time.UTC) + previous := &UpstreamBillingProbeSnapshot{ + Status: UpstreamBillingProbeStatusOK, + Data: map[string]any{"effective_rate_multiplier": 0.5}, + ReceivedAt: &receivedAt, + FailureCount: 1, + } + account := &Account{ + ID: 18, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeExtraKey: previous}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"14400"}}, + Body: io.NopCloser(strings.NewReader(`{"error":"do not persist this"}`)), + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + fixedNow := time.Date(2026, time.July, 13, 2, 0, 0, 0, time.UTC) + svc.now = func() time.Time { return fixedNow } + + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusFailed, snapshot.Status) + require.Equal(t, previous.Data, snapshot.Data) + require.Equal(t, previous.ReceivedAt, snapshot.ReceivedAt) + require.NotNil(t, snapshot.FreshUntil) + require.Equal(t, receivedAt.Add(time.Hour), *snapshot.FreshUntil) + require.Equal(t, 2, snapshot.FailureCount) + require.Equal(t, "http_error", snapshot.LastError) + require.False(t, snapshot.NextProbeAt.Before(fixedNow.Add(4*time.Hour))) + require.NotContains(t, snapshot.LastError, "do not persist") +} + +func TestUpstreamBillingProbeRetryAfterIsNotShortened(t *testing.T) { + delay := nextProbeDelay(30, 1, 48*time.Hour) + require.Equal(t, 48*time.Hour, delay) +} + +func TestUpstreamBillingProbeEmptyResponseIsPersistedAsFailure(t *testing.T) { + account := &Account{ + ID: 21, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + svc := newUpstreamBillingProbeTestService(repo, &httpUpstreamRecorder{}, &upstreamBillingProbeSettingRepo{}) + + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusFailed, snapshot.Status) + require.Equal(t, "empty_response", snapshot.LastError) +} + +func TestUpstreamBillingProbeUnsupportedAndAccountToggle(t *testing.T) { + account := &Account{ + ID: 19, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusNotFound, + Header: http.Header{}, + Body: io.NopCloser(strings.NewReader("not found")), + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + + require.NoError(t, svc.SetAccountEnabled(context.Background(), account.ID, true)) + require.Equal(t, true, account.Extra[UpstreamBillingProbeEnabledExtraKey]) + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusUnsupported, snapshot.Status) + require.Equal(t, "unsupported", snapshot.LastError) + + invalid := &Account{ID: 20, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + repo.accounts[invalid.ID] = invalid + err = svc.SetAccountEnabled(context.Background(), invalid.ID, true) + require.True(t, errors.Is(err, ErrUpstreamBillingProbeAccountInvalid)) +} + +func TestUpstreamBillingProbeRunnerIsBoundedAndManualProbeIgnoresSwitches(t *testing.T) { + accounts := make(map[int64]*Account, 25) + for id := int64(1); id <= 25; id++ { + accounts[id] = &Account{ + ID: id, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + } + repo := &upstreamBillingProbeAccountRepo{accounts: accounts} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyUpstreamBillingProbeSettings: `{"enabled":true,"interval_minutes":30}`, + }} + upstream := &upstreamBillingProbeHTTPStub{} + svc := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + svc.now = func() time.Time { return time.Date(2026, time.July, 13, 2, 0, 0, 0, time.UTC) } + + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(20), upstream.calls.Load()) + + settingsRepo.mu.Lock() + settingsRepo.values[SettingKeyUpstreamBillingProbeSettings] = `{"enabled":false,"interval_minutes":30}` + settingsRepo.mu.Unlock() + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(20), upstream.calls.Load()) + + accounts[25].Extra[UpstreamBillingProbeEnabledExtraKey] = false + snapshot, err := svc.ProbeAccount(context.Background(), 25) + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusOK, snapshot.Status) + require.Equal(t, int64(21), upstream.calls.Load()) +} + +func TestUpstreamBillingProbeRunnerRechecksEnabledAfterDueSelection(t *testing.T) { + account := &Account{ + ID: 26, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: false}, + } + staleDue := *account + staleDue.Extra = map[string]any{UpstreamBillingProbeEnabledExtraKey: true} + baseRepo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + repo := &staleDueUpstreamBillingProbeAccountRepo{upstreamBillingProbeAccountRepo: baseRepo, due: []Account{staleDue}} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyUpstreamBillingProbeSettings: `{"enabled":true,"interval_minutes":30}`, + }} + upstream := &upstreamBillingProbeHTTPStub{} + svc := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + + require.NoError(t, svc.RunDue(context.Background())) + require.Zero(t, upstream.calls.Load()) + require.NotContains(t, account.Extra, UpstreamBillingProbeExtraKey) +} + +func TestUpstreamBillingProbeNeverDowngradesMissingConfiguredProxyToDirect(t *testing.T) { + proxyID := int64(7) + for _, tc := range []struct { + name string + proxy *Proxy + wantReason string + wantErr error + }{ + {name: "missing hydrated proxy", wantReason: "proxy_unavailable"}, + {name: "mismatched hydrated proxy", proxy: &Proxy{ID: 8, Protocol: "http", Host: "127.0.0.1", Port: 8080}, wantErr: ErrUpstreamBillingProbeIdentityChanged}, + } { + t.Run(tc.name, func(t *testing.T) { + account := &Account{ + ID: 27, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-sensitive", "base_url": "https://upstream.example"}, + ProxyID: &proxyID, + Proxy: tc.proxy, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &upstreamBillingProbeHTTPStub{} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + + snapshot, err := svc.ProbeAccount(context.Background(), account.ID) + if tc.wantErr != nil { + require.ErrorIs(t, err, tc.wantErr) + require.Nil(t, snapshot) + } else { + require.NoError(t, err) + require.Equal(t, UpstreamBillingProbeStatusFailed, snapshot.Status) + require.Equal(t, tc.wantReason, snapshot.LastError) + } + require.Zero(t, upstream.calls.Load()) + if tc.wantErr != nil { + require.NotContains(t, account.Extra, UpstreamBillingProbeExtraKey) + } + }) + } +} + +func TestUpstreamBillingProbeRunnerOnlyScansOnLeader(t *testing.T) { + account := &Account{ + ID: 31, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &upstreamBillingProbeHTTPStub{} + cache := &fakeLeaderLockCache{} + lockKey := upstreamBillingProbeLeaderLockKeyAt(time.Now()) + peer := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + peer.instanceID = "peer" + peer.SetLeaderLock(cache, nil) + _, acquired, err := peer.tryAcquireLeaderLock(context.Background(), lockKey) + require.NoError(t, err) + require.True(t, acquired) + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + svc.SetLeaderLock(cache, nil) + + require.NoError(t, svc.RunDue(context.Background())) + require.Zero(t, upstream.calls.Load()) + + require.NoError(t, cache.ReleaseLeaderLock(context.Background(), lockKey, "peer")) + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), upstream.calls.Load()) +} + +func TestUpstreamBillingProbeLeaderLockFailsClosedOnCacheError(t *testing.T) { + svc := newUpstreamBillingProbeTestService(&upstreamBillingProbeAccountRepo{}, &upstreamBillingProbeHTTPStub{}, &upstreamBillingProbeSettingRepo{}) + svc.SetLeaderLock(&fakeLeaderLockCache{acquireErr: context.DeadlineExceeded}, nil) + + release, acquired, err := svc.tryAcquireLeaderLock(context.Background(), upstreamBillingProbeLeaderLockKeyAt(time.Now())) + + require.ErrorIs(t, err, context.DeadlineExceeded) + require.False(t, acquired) + require.Nil(t, release) +} + +func TestUpstreamBillingProbeLeaderLockUsesCadenceBuckets(t *testing.T) { + cache := &fakeLeaderLockCache{} + first := newUpstreamBillingProbeTestService(&upstreamBillingProbeAccountRepo{}, &upstreamBillingProbeHTTPStub{}, &upstreamBillingProbeSettingRepo{}) + second := newUpstreamBillingProbeTestService(&upstreamBillingProbeAccountRepo{}, &upstreamBillingProbeHTTPStub{}, &upstreamBillingProbeSettingRepo{}) + first.SetLeaderLock(cache, nil) + second.SetLeaderLock(cache, nil) + beforeBoundary := time.Unix(59, 0) + afterBoundary := beforeBoundary.Add(time.Second) + + releaseFirst, acquired, err := first.tryAcquireLeaderLock(context.Background(), upstreamBillingProbeLeaderLockKeyAt(beforeBoundary)) + require.NoError(t, err) + require.True(t, acquired) + releaseSecond, acquired, err := second.tryAcquireLeaderLock(context.Background(), upstreamBillingProbeLeaderLockKeyAt(afterBoundary)) + require.NoError(t, err) + require.True(t, acquired, "the prior cadence lock must not suppress the next cadence") + releaseFirst() + releaseSecond() +} + +func TestUpstreamBillingProbeFiveInstancesRunOneConcurrentBatch(t *testing.T) { + account := &Account{ + ID: 32, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "http://127.0.0.1:8080"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyUpstreamBillingProbeSettings: `{"enabled":true,"interval_minutes":30}`, + }} + cache := &fakeLeaderLockCache{} + entered := make(chan struct{}) + unblock := make(chan struct{}) + var enteredOnce sync.Once + upstream := &upstreamBillingProbeHTTPStub{beforeResponse: func() { + enteredOnce.Do(func() { close(entered) }) + <-unblock + }} + + start := make(chan struct{}) + results := make(chan error, 5) + for range 5 { + svc := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + svc.SetLeaderLock(cache, nil) + go func() { + <-start + results <- svc.RunDue(context.Background()) + }() + } + close(start) + + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("leader did not start the probe batch") + } + for range 4 { + select { + case err := <-results: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("non-leader instance did not skip the active batch") + } + } + require.Equal(t, int64(1), upstream.calls.Load()) + close(unblock) + require.NoError(t, <-results) + require.Equal(t, int64(1), upstream.calls.Load()) +} + +func TestUpstreamBillingProbeManualBatchesShareConcurrencyLimit(t *testing.T) { + accounts := make(map[int64]*Account, 12) + for id := int64(1); id <= 12; id++ { + accounts[id] = &Account{ + ID: id, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "http://127.0.0.1:8080"}, + } + } + repo := &upstreamBillingProbeAccountRepo{accounts: accounts} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyUpstreamBillingProbeSettings: `{"enabled":true,"interval_minutes":30}`, + }} + entered := make(chan struct{}, len(accounts)) + unblock := make(chan struct{}) + var unblockOnce sync.Once + release := func() { unblockOnce.Do(func() { close(unblock) }) } + t.Cleanup(release) + upstream := &upstreamBillingProbeHTTPStub{beforeResponse: func() { + entered <- struct{}{} + <-unblock + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + + results := make(chan []UpstreamBillingProbeResult, 3) + for batch := 0; batch < 3; batch++ { + firstID := int64(batch*4 + 1) + ids := []int64{firstID, firstID + 1, firstID + 2, firstID + 3} + go func() { results <- svc.ProbeAccounts(context.Background(), ids) }() + } + for range upstreamBillingProbeConcurrency { + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("shared probe slots did not fill") + } + } + select { + case <-entered: + release() + t.Fatal("parallel manual batches exceeded the service-wide concurrency limit") + case <-time.After(100 * time.Millisecond): + } + release() + + for range 3 { + select { + case batchResults := <-results: + for _, result := range batchResults { + require.Empty(t, result.Error) + require.NotNil(t, result.Snapshot) + } + case <-time.After(time.Second): + t.Fatal("manual probe batch did not finish") + } + } + require.Equal(t, int64(upstreamBillingProbeConcurrency), upstream.maxActive.Load()) +} + +func TestUpstreamBillingProbeManualAndScheduledRequestsShareOneNetworkProbe(t *testing.T) { + account := &Account{ + ID: 46, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + started := make(chan struct{}) + unblock := make(chan struct{}) + var startedOnce sync.Once + upstream := &upstreamBillingProbeHTTPStub{beforeResponse: func() { + startedOnce.Do(func() { close(started) }) + <-unblock + }} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + + errs := make(chan error, 2) + go func() { + _, err := svc.probeScheduledAccount(context.Background(), account.ID, 30) + errs <- err + }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("scheduled probe did not reach the upstream") + } + manualStarted := make(chan struct{}) + go func() { + close(manualStarted) + _, err := svc.ProbeAccount(context.Background(), account.ID) + errs <- err + }() + <-manualStarted + time.Sleep(20 * time.Millisecond) + close(unblock) + require.NoError(t, <-errs) + require.NoError(t, <-errs) + require.Equal(t, int64(1), upstream.calls.Load()) +} + +func TestUpstreamBillingProbeScheduledRechecksAfterWaitingForSlot(t *testing.T) { + account := &Account{ + ID: 47, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "https://upstream.example"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}} + upstream := &upstreamBillingProbeHTTPStub{} + svc := newUpstreamBillingProbeTestService(repo, upstream, &upstreamBillingProbeSettingRepo{}) + for range upstreamBillingProbeConcurrency { + svc.probeSlots <- struct{}{} + } + result := make(chan error, 1) + go func() { + _, err := svc.probeScheduledAccount(context.Background(), account.ID, 30) + result <- err + }() + time.Sleep(20 * time.Millisecond) + repo.mu.Lock() + account.Extra[UpstreamBillingProbeEnabledExtraKey] = false + repo.mu.Unlock() + <-svc.probeSlots + + require.NoError(t, <-result) + require.Zero(t, upstream.calls.Load()) +} + +func TestUpstreamBillingProbeLeaderLockCoversStaggeredInstancesInCadenceWindow(t *testing.T) { + account := func(id int64) *Account { + return &Account{ + ID: id, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test", "base_url": "http://127.0.0.1:8080"}, + Extra: map[string]any{UpstreamBillingProbeEnabledExtraKey: true}, + } + } + repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{41: account(41)}} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyUpstreamBillingProbeSettings: `{"enabled":true,"interval_minutes":30}`, + }} + cache := &fakeLeaderLockCache{} + upstream := &upstreamBillingProbeHTTPStub{} + first := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + first.SetLeaderLock(cache, nil) + + require.NoError(t, first.RunDue(context.Background())) + require.Equal(t, int64(1), upstream.calls.Load()) + require.Equal(t, first.instanceID, cache.heldBy(upstreamBillingProbeLeaderLockKeyAt(time.Now()))) + + repo.mu.Lock() + repo.accounts[42] = account(42) + repo.mu.Unlock() + staggered := newUpstreamBillingProbeTestService(repo, upstream, settingsRepo) + staggered.SetLeaderLock(cache, nil) + require.NoError(t, staggered.RunDue(context.Background())) + require.Equal(t, int64(1), upstream.calls.Load(), "a staggered instance must not start a second batch inside the cadence window") +} diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index ace69e0c34..ba078308c1 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -664,6 +664,7 @@ var ProviderSet = wire.NewSet( ProvideRateLimitService, ProvideAccountUsageService, ProvideAccountTestService, + ProvideUpstreamBillingProbeService, ProvideSettingService, NewDataManagementService, ProvideBackupService, diff --git a/frontend/src/api/__tests__/admin.accounts.upstreamBillingProbe.spec.ts b/frontend/src/api/__tests__/admin.accounts.upstreamBillingProbe.spec.ts new file mode 100644 index 0000000000..cd536c3c39 --- /dev/null +++ b/frontend/src/api/__tests__/admin.accounts.upstreamBillingProbe.spec.ts @@ -0,0 +1,53 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { get, post, put } = vi.hoisted(() => ({ + get: vi.fn(), + post: vi.fn(), + put: vi.fn() +})) + +vi.mock('@/api/client', () => ({ + apiClient: { get, post, put } +})) + +import { + getUpstreamBillingProbeSettings, + probeUpstreamBilling, + probeUpstreamBillingBatch, + setUpstreamBillingProbeEnabled, + updateUpstreamBillingProbeSettings +} from '@/api/admin/accounts' + +describe('admin account upstream billing probe API', () => { + beforeEach(() => { + get.mockReset() + post.mockReset() + put.mockReset() + }) + + it('reads and updates global settings', async () => { + const settings = { enabled: true, interval_minutes: 30 } + get.mockResolvedValueOnce({ data: settings }) + put.mockResolvedValueOnce({ data: settings }) + + await expect(getUpstreamBillingProbeSettings()).resolves.toEqual(settings) + await expect(updateUpstreamBillingProbeSettings(settings)).resolves.toEqual(settings) + expect(get).toHaveBeenCalledWith('/admin/accounts/upstream-billing-probe/settings') + expect(put).toHaveBeenCalledWith('/admin/accounts/upstream-billing-probe/settings', settings) + }) + + it('uses dedicated account and batch endpoints', async () => { + const result = { account_id: 7, snapshot: { status: 'unsupported' } } + put.mockResolvedValueOnce({ data: {} }) + post.mockResolvedValueOnce({ data: result }) + post.mockResolvedValueOnce({ data: { results: [result] } }) + + await setUpstreamBillingProbeEnabled(7, true) + await expect(probeUpstreamBilling(7)).resolves.toEqual(result) + await expect(probeUpstreamBillingBatch([7])).resolves.toEqual([result]) + + expect(put).toHaveBeenCalledWith('/admin/accounts/7/upstream-billing-probe', { enabled: true }) + expect(post).toHaveBeenNthCalledWith(1, '/admin/accounts/7/upstream-billing-probe') + expect(post).toHaveBeenNthCalledWith(2, '/admin/accounts/upstream-billing-probe/batch', { account_ids: [7] }) + }) +}) diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index 64a47a50fb..5bc9752155 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -20,7 +20,9 @@ import type { CodexSessionImportResult, OpenAICodexPATCreateRequest, CheckMixedChannelRequest, - CheckMixedChannelResponse + CheckMixedChannelResponse, + UpstreamBillingProbeResult, + UpstreamBillingProbeSettings } from '@/types' /** @@ -848,6 +850,38 @@ export async function createSparkShadow(parentId: number, payload: SparkShadowCr return data } +export async function getUpstreamBillingProbeSettings(): Promise { + const { data } = await apiClient.get('/admin/accounts/upstream-billing-probe/settings') + return data +} + +export async function updateUpstreamBillingProbeSettings( + settings: UpstreamBillingProbeSettings +): Promise { + const { data } = await apiClient.put( + '/admin/accounts/upstream-billing-probe/settings', + settings + ) + return data +} + +export async function setUpstreamBillingProbeEnabled(id: number, enabled: boolean): Promise { + await apiClient.put(`/admin/accounts/${id}/upstream-billing-probe`, { enabled }) +} + +export async function probeUpstreamBilling(id: number): Promise { + const { data } = await apiClient.post(`/admin/accounts/${id}/upstream-billing-probe`) + return data +} + +export async function probeUpstreamBillingBatch(accountIds: number[]): Promise { + const { data } = await apiClient.post<{ results: UpstreamBillingProbeResult[] }>( + '/admin/accounts/upstream-billing-probe/batch', + { account_ids: accountIds } + ) + return data.results +} + export const accountsAPI = { list, listWithEtag, @@ -894,7 +928,12 @@ export const accountsAPI = { revertProxyFallback, queryOpenAIQuota, resetOpenAIQuota, - createSparkShadow + createSparkShadow, + getUpstreamBillingProbeSettings, + updateUpstreamBillingProbeSettings, + setUpstreamBillingProbeEnabled, + probeUpstreamBilling, + probeUpstreamBillingBatch } export default accountsAPI diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 418fb62d88..eed9920bb2 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1592,6 +1592,23 @@ +
+
+ +

+ {{ t('admin.accounts.upstreamBilling.autoProbeHint') }} +

+
+ +
+
(null) const autoPause7dThreshold = ref(null) const autoPause5hDisabled = ref(false) const autoPause7dDisabled = ref(false) +const upstreamBillingAutoProbeEnabled = ref(false) const mixedScheduling = ref(false) // For antigravity accounts: enable mixed scheduling const allowOverages = ref(false) // For antigravity accounts: enable AI Credits overages const antigravityProjectId = ref('') @@ -3210,6 +3229,7 @@ const syncFormFromAccount = (newAccount: Account | null) => { autoPause7dThreshold.value = typeof extra?.auto_pause_7d_threshold === 'number' ? extra.auto_pause_7d_threshold * 100 : null autoPause5hDisabled.value = extra?.auto_pause_5h_disabled === true autoPause7dDisabled.value = extra?.auto_pause_7d_disabled === true + upstreamBillingAutoProbeEnabled.value = extra?.upstream_billing_probe_enabled === true // Load OpenAI passthrough toggle (OpenAI OAuth/SetupToken/API Key) openaiPassthroughEnabled.value = false @@ -4463,6 +4483,7 @@ const handleSubmit = async () => { } else { newExtra.openai_responses_mode = openAIResponsesMode.value } + newExtra.upstream_billing_probe_enabled = upstreamBillingAutoProbeEnabled.value } if (autoPause5hThreshold.value != null && autoPause5hThreshold.value > 0) { newExtra.auto_pause_5h_threshold = autoPause5hThreshold.value / 100 diff --git a/frontend/src/components/account/UpstreamBillingRateCell.vue b/frontend/src/components/account/UpstreamBillingRateCell.vue new file mode 100644 index 0000000000..73db04d851 --- /dev/null +++ b/frontend/src/components/account/UpstreamBillingRateCell.vue @@ -0,0 +1,160 @@ + + + diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index b3a583d102..98de8f0c06 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -588,6 +588,24 @@ describe('EditAccountModal', () => { expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_responses_supported).toBe(false) }) + it('submits the account upstream billing auto-probe setting', async () => { + const account = buildAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="upstream-billing-auto-probe"]') + expect(toggle.attributes('aria-checked')).toBe('false') + + await toggle.trigger('click') + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.upstream_billing_probe_enabled).toBe(true) + }) + it('clears OpenAI APIKey Responses override when set back to auto', async () => { const account = buildAccount() account.extra = { diff --git a/frontend/src/components/account/__tests__/UpstreamBillingRateCell.spec.ts b/frontend/src/components/account/__tests__/UpstreamBillingRateCell.spec.ts new file mode 100644 index 0000000000..8246c70662 --- /dev/null +++ b/frontend/src/components/account/__tests__/UpstreamBillingRateCell.spec.ts @@ -0,0 +1,233 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { mount } from '@vue/test-utils' +import UpstreamBillingRateCell from '../UpstreamBillingRateCell.vue' +import type { Account } from '@/types' + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ + t: (key: string, params?: Record) => + params ? `${key}:${Object.values(params).join(',')}` : key + }) + } +}) + +const makeAccount = (overrides: Partial = {}): Account => ({ + id: 1, + name: 'upstream', + platform: 'openai', + type: 'apikey', + proxy_id: null, + concurrency: 1, + priority: 1, + status: 'active', + error_message: null, + last_used_at: null, + expires_at: null, + auto_pause_on_expired: false, + created_at: '2026-07-13T00:00:00Z', + updated_at: '2026-07-13T00:00:00Z', + schedulable: true, + rate_limited_at: null, + rate_limit_reset_at: null, + overload_until: null, + temp_unschedulable_until: null, + temp_unschedulable_reason: null, + session_window_start: null, + session_window_end: null, + session_window_status: null, + ...overrides +}) + +const billingData = { + object: 'sub2api.key_billing' as const, + schema_version: 1 as const, + billing_scope: 'token' as const, + group_rate_multiplier: 0.8, + resolved_rate_multiplier: 0.6, + peak_rate_enabled: true, + peak_start: '09:00', + peak_end: '18:00', + peak_rate_multiplier: 1.5, + applied_peak_multiplier: 1.5, + effective_rate_multiplier: 0.9, + timezone: 'Asia/Shanghai', + observed_at: '2026-07-13T00:00:00Z' +} + +describe('UpstreamBillingRateCell', () => { + beforeEach(() => { + vi.useFakeTimers() + vi.setSystemTime(new Date('2026-07-13T00:30:00Z')) + }) + + afterEach(() => { + vi.useRealTimers() + }) + + it('recomputes the current effective rate and keeps the icon-only probe action', async () => { + const wrapper = mount(UpstreamBillingRateCell, { + props: { + account: makeAccount({ + extra: { + upstream_billing_probe_enabled: true, + upstream_billing_probe: { + status: 'ok', + data: billingData, + received_at: '2026-07-13T00:00:00Z', + fresh_until: '2026-07-14T00:00:00Z', + last_attempt_at: '2026-07-13T00:00:00Z', + next_probe_at: '2026-07-13T00:30:00Z' + } + } + }), + intervalMinutes: 30, + now: Date.now() + } + }) + + expect(wrapper.text()).toContain('0.6x') + await wrapper.setProps({ now: Date.parse('2026-07-13T01:00:00Z') }) + expect(wrapper.text()).toContain('0.9x') + await wrapper.setProps({ now: Date.parse('2026-07-13T10:00:00Z') }) + expect(wrapper.text()).toContain('0.6x') + expect(wrapper.text()).not.toContain('admin.accounts.upstreamBilling.latest') + expect(wrapper.get('[data-testid="upstream-billing-probe"]').text()).toBe('') + expect(wrapper.get('[data-testid="upstream-billing-probe"]').attributes('aria-label')).toBe( + 'admin.accounts.upstreamBilling.manualProbe' + ) + }) + + it('uses retained failed data only while it is still fresh', async () => { + const account = makeAccount({ + extra: { + upstream_billing_probe: { + status: 'ok', + data: billingData, + received_at: '2026-07-12T22:00:00Z', + fresh_until: '2026-07-12T23:00:00Z', + last_attempt_at: '2026-07-12T22:00:00Z', + next_probe_at: '2026-07-12T22:30:00Z' + } + } + }) + const wrapper = mount(UpstreamBillingRateCell, { props: { account, intervalMinutes: 30, now: Date.now() } }) + expect(wrapper.text()).toContain('admin.accounts.upstreamBilling.stale') + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + + await wrapper.setProps({ + account: makeAccount({ + extra: { + upstream_billing_probe: { + status: 'failed', + data: billingData, + received_at: '2026-07-13T00:00:00Z', + fresh_until: '2026-07-13T01:00:00Z', + last_attempt_at: '2026-07-13T00:00:00Z', + next_probe_at: '2026-07-13T01:00:00Z', + last_error: 'http_error' + } + } + }) + }) + expect(wrapper.text()).toContain('0.6x') + expect(wrapper.text()).toContain('admin.accounts.upstreamBilling.failed') + + await wrapper.setProps({ now: Date.parse('2026-07-13T01:00:00Z') }) + expect(wrapper.text()).toContain('0.9x') + expect(wrapper.text()).not.toContain('admin.accounts.upstreamBilling.stale') + + await wrapper.setProps({ now: Date.parse('2026-07-13T01:00:00.001Z') }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + expect(wrapper.text()).toContain('admin.accounts.upstreamBilling.stale') + + await wrapper.setProps({ + now: Date.now(), + account: makeAccount({ + extra: { + upstream_billing_probe: { + status: 'failed', + data: billingData, + received_at: '2026-07-12T22:00:00Z', + fresh_until: '2026-07-12T23:00:00Z', + last_attempt_at: '2026-07-13T00:00:00Z', + next_probe_at: '2026-07-13T01:00:00Z', + last_error: 'http_error' + } + } + }) + }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + expect(wrapper.text()).toContain('admin.accounts.upstreamBilling.stale') + }) + + it('emits manual probe commands only for eligible accounts', async () => { + const wrapper = mount(UpstreamBillingRateCell, { + props: { account: makeAccount(), intervalMinutes: 30, now: Date.now() } + }) + await wrapper.get('[data-testid="upstream-billing-probe"]').trigger('click') + expect(wrapper.emitted('probe')).toHaveLength(1) + + await wrapper.setProps({ account: makeAccount({ type: 'oauth' }) }) + expect(wrapper.findAll('button')).toHaveLength(0) + expect(wrapper.text()).toBe('-') + }) + + it('fails neutral for malformed data and timestamps', async () => { + const malformedAccount = ( + dataOverrides: Partial = {}, + snapshotOverrides: Record = {} + ) => makeAccount({ + extra: { + upstream_billing_probe: { + status: 'ok', + data: { ...billingData, ...dataOverrides }, + received_at: '2026-07-13T00:00:00Z', + fresh_until: '2026-07-13T01:00:00Z', + last_attempt_at: '2026-07-13T00:00:00Z', + next_probe_at: '2026-07-13T01:00:00Z', + ...snapshotOverrides + } + } + }) + const wrapper = mount(UpstreamBillingRateCell, { + props: { + account: malformedAccount({ + resolved_rate_multiplier: -1, + peak_rate_enabled: false, + effective_rate_multiplier: -1 + }), + intervalMinutes: 30, + now: Date.now() + } + }) + + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + await wrapper.setProps({ account: malformedAccount({ billing_scope: 'request' as 'token' }) }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + await wrapper.setProps({ account: malformedAccount({}, { received_at: 'not-a-time' }) }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + await wrapper.setProps({ account: malformedAccount({}, { received_at: '2026-07-13T00:31:00Z' }) }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + await wrapper.setProps({ account: malformedAccount({}, { fresh_until: '2026-07-12T23:59:00Z' }) }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + + await wrapper.setProps({ + account: makeAccount({ + extra: { + upstream_billing_probe: { + status: 'failed', + last_attempt_at: '2026-07-13T00:00:00Z', + next_probe_at: '2026-07-13T01:00:00Z', + last_error: 'network_error' + } + } + }) + }) + expect(wrapper.get('[data-testid="upstream-billing-rate"]').text()).toBe('-') + expect(wrapper.text()).toContain('admin.accounts.upstreamBilling.failed') + expect(wrapper.text()).not.toContain('admin.accounts.upstreamBilling.stale') + }) +}) diff --git a/frontend/src/components/admin/account/AccountBulkActionsBar.vue b/frontend/src/components/admin/account/AccountBulkActionsBar.vue index a632bdd421..f3a062d0d6 100644 --- a/frontend/src/components/admin/account/AccountBulkActionsBar.vue +++ b/frontend/src/components/admin/account/AccountBulkActionsBar.vue @@ -28,6 +28,7 @@ + @@ -41,5 +42,19 @@ diff --git a/frontend/src/i18n/locales/en/admin/accounts.ts b/frontend/src/i18n/locales/en/admin/accounts.ts index d0be8bbc84..518137d12e 100644 --- a/frontend/src/i18n/locales/en/admin/accounts.ts +++ b/frontend/src/i18n/locales/en/admin/accounts.ts @@ -152,6 +152,7 @@ export default { notes: 'Notes', priority: 'Priority', billingRateMultiplier: 'Billing Rate', + upstreamBillingRate: 'Upstream Declared Rate', weight: 'Weight', schedulerScore: 'Scheduler Score', status: 'Status', @@ -172,6 +173,31 @@ export default { hint: 'Displayed as "group / base score / sticky bonus". The base score is computed within the current filtered candidate set and includes priority, load, queue depth, error rate, first-token latency, reset window, quota headroom, and related factors. The sticky bonus applies only when sticky weighting is enabled for previous_response_id or session_hash. Higher scores are preferred.' }, usageWindowsHint: '"5h / 7d" are the upstream account\'s official rolling usage windows (e.g. OpenAI ChatGPT, Claude). They are imposed by the upstream provider on the account itself — not configured by sub2api, and unrelated to the models you map. Usage resets automatically once each window rolls over, and the limit cannot be lifted from within sub2api.', + upstreamBilling: { + trustWarning: 'This rate is declared by the upstream site for the current API key. Sub2API cannot verify that it matches actual charges. The upstream site or an intermediary may return forged, stale, or modified data. Verify it against bills, balance changes, and actual usage.', + autoProbeSettings: 'Upstream rate auto probe', + intervalMinutes: 'Probe interval (minutes)', + autoProbe: 'Auto probe', + autoProbeHint: 'Probe this account on the global interval when global probing is enabled.', + manualProbe: 'Probe upstream rate now', + stale: 'Stale', + unsupported: 'Unsupported', + failed: 'Failed', + notProbed: 'Not probed', + groupRate: 'Group default: {value}x', + userRate: 'User rate: {value}x', + peakRate: 'Peak: {start}-{end}, {value}x ({timezone})', + noPeakRate: 'Peak rate: disabled', + effectiveRate: 'Current rate: {value}x', + updatedAt: 'Updated: {value}', + settingsSaved: 'Upstream rate probe settings saved', + settingsFailed: 'Failed to save upstream rate probe settings', + probeFailed: 'Failed to probe upstream rate', + noEligibleAccounts: 'Select OpenAI API key accounts', + batchLimit: 'A batch can probe at most 20 accounts', + batchCompleted: 'Probed {count} account(s)', + batchPartial: 'Probe partially completed: {success} succeeded, {failed} failed' + }, allPrivacyModes: 'All Privacy States', privacyUnset: 'Unset', privacyTrainingOff: 'Training data sharing disabled', @@ -313,6 +339,7 @@ export default { disableScheduling: 'Disable Scheduling', resetStatus: 'Reset Status', refreshToken: 'Refresh Token', + probeUpstreamBilling: 'Probe Upstream Rate', resetStatusSuccess: 'Successfully reset {count} account(s) status', refreshTokenSuccess: 'Successfully refreshed {count} account(s) token', partialSuccess: 'Partially completed: {success} succeeded, {failed} failed' diff --git a/frontend/src/i18n/locales/zh/admin/accounts.ts b/frontend/src/i18n/locales/zh/admin/accounts.ts index 1feb1030c0..8364dd0ffa 100644 --- a/frontend/src/i18n/locales/zh/admin/accounts.ts +++ b/frontend/src/i18n/locales/zh/admin/accounts.ts @@ -108,6 +108,7 @@ export default { notes: '备注', priority: '优先级', billingRateMultiplier: '账号倍率', + upstreamBillingRate: '上游声明倍率', weight: '权重', schedulerScore: '调度权值', status: '状态', @@ -128,6 +129,31 @@ export default { hint: '显示格式为“分组名 / 基础分 / 粘性加分”。基础分按当前筛选条件限定的候选账号计算,包含优先级、负载、排队、错误率、首包延迟、重置窗口、额度余量等因子;粘性加分只在开启粘性加权时用于 previous_response_id 或 session_hash。分数越大越优先。' }, usageWindowsHint: '“5h / 7d”是上游账号(如 OpenAI ChatGPT、Claude)官方的滚动用量窗口限制,由上游对账号设定,并非 sub2api 配置,也与你映射的模型无关。窗口滚动到期后用量会自动重置,无法在 sub2api 端解除该限制。', + upstreamBilling: { + trustWarning: '此倍率由上游站点针对当前 API Key 自行声明。Sub2API 无法验证该值是否与实际扣费一致;上游站点或中间代理可能返回伪造、过期或被篡改的数据。请结合账单、余额变化和实际用量自行核验。', + autoProbeSettings: '上游倍率自动探测', + intervalMinutes: '探测周期(分钟)', + autoProbe: '自动探测', + autoProbeHint: '启用后按全局探测周期查询此账号;全局探测关闭时不会执行。', + manualProbe: '立即探测上游倍率', + stale: '已过期', + unsupported: '不支持', + failed: '失败', + notProbed: '未探测', + groupRate: '分组默认:{value}x', + userRate: '用户专属倍率:{value}x', + peakRate: '高峰:{start}-{end},{value}x({timezone})', + noPeakRate: '高峰倍率:未启用', + effectiveRate: '当前倍率:{value}x', + updatedAt: '更新时间:{value}', + settingsSaved: '上游倍率探测设置已保存', + settingsFailed: '保存上游倍率探测设置失败', + probeFailed: '探测上游倍率失败', + noEligibleAccounts: '请选择 OpenAI API Key 账号', + batchLimit: '每次最多探测 20 个账号', + batchCompleted: '已完成 {count} 个账号的倍率探测', + batchPartial: '倍率探测部分完成:成功 {success} 个,失败 {failed} 个' + }, allPrivacyModes: '全部Privacy状态', privacyUnset: '未设置', privacyTrainingOff: '已关闭训练数据共享', @@ -417,6 +443,7 @@ export default { disableScheduling: '批量停止调度', resetStatus: '批量重置状态', refreshToken: '批量刷新令牌', + probeUpstreamBilling: '探测上游倍率', resetStatusSuccess: '已成功重置 {count} 个账号状态', refreshTokenSuccess: '已成功刷新 {count} 个账号令牌', partialSuccess: '操作部分完成:{success} 成功,{failed} 失败' diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 1d981097ef..f1d15a7db3 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -865,6 +865,48 @@ export interface TempUnschedulableStatus { state?: TempUnschedulableState } +export interface UpstreamBillingData { + object: 'sub2api.key_billing' + schema_version: 1 + billing_scope: 'token' + group_rate_multiplier: number + user_rate_multiplier?: number + resolved_rate_multiplier: number + peak_rate_enabled: boolean + peak_start?: string + peak_end?: string + peak_rate_multiplier?: number + applied_peak_multiplier?: number + effective_rate_multiplier: number + timezone?: string + observed_at: string +} + +export type UpstreamBillingProbeStatus = 'ok' | 'unsupported' | 'failed' + +export interface UpstreamBillingProbeSnapshot { + status: UpstreamBillingProbeStatus + data?: UpstreamBillingData + received_at?: string + fresh_until?: string + last_attempt_at: string + next_probe_at: string + failure_count?: number + http_status?: number + last_error?: string +} + +export interface UpstreamBillingProbeSettings { + enabled: boolean + interval_minutes: number +} + +export interface UpstreamBillingProbeResult { + account_id: number + snapshot?: UpstreamBillingProbeSnapshot + error?: string +} + export interface Account { id: number name: string @@ -881,6 +923,8 @@ export interface Account { extra?: (CodexUsageSnapshot & OpenAICompactState & { model_rate_limits?: Record antigravity_credits_overages?: Record + upstream_billing_probe_enabled?: boolean + upstream_billing_probe?: UpstreamBillingProbeSnapshot } & Record) proxy_id: number | null proxy_fallback_origin_id?: number | null diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index 11e68a1ac3..688ca50027 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -132,6 +132,41 @@ {{ t('admin.tlsFingerprintProfiles.title') }} +
+
+
+ + {{ t('admin.accounts.upstreamBilling.autoProbeSettings') }} + + +
+
+ + + +
+
+
@@ -177,6 +212,7 @@ @delete="handleBulkDelete" @reset-status="handleBulkResetStatus" @refresh-token="handleBulkRefreshToken" + @probe-upstream-billing="handleBulkProbeUpstreamBilling" @edit-selected="openBulkEditSelected" @edit-filtered="openBulkEditFiltered" @clear="clearSelection" @@ -321,6 +357,21 @@ {{ (row.rate_multiplier ?? 1).toFixed(2) }}x + + @@ -441,6 +492,7 @@ import AppLayout from '@/components/layout/AppLayout.vue' import TablePageLayout from '@/components/layout/TablePageLayout.vue' import DataTable from '@/components/common/DataTable.vue' import HelpTooltip from '@/components/common/HelpTooltip.vue' +import Toggle from '@/components/common/Toggle.vue' import Pagination from '@/components/common/Pagination.vue' import ConfirmDialog from '@/components/common/ConfirmDialog.vue' import { CreateAccountModal, EditAccountModal, BulkEditAccountModal, SyncFromCrsModal, TempUnschedStatusModal } from '@/components/account' @@ -459,6 +511,7 @@ import AccountUsageCell from '@/components/account/AccountUsageCell.vue' import AccountTodayStatsCell from '@/components/account/AccountTodayStatsCell.vue' import AccountGroupsCell from '@/components/account/AccountGroupsCell.vue' import AccountCapacityCell from '@/components/account/AccountCapacityCell.vue' +import UpstreamBillingRateCell from '@/components/account/UpstreamBillingRateCell.vue' import PlatformTypeBadge from '@/components/common/PlatformTypeBadge.vue' import Icon from '@/components/icons/Icon.vue' import ErrorPassthroughRulesModal from '@/components/admin/ErrorPassthroughRulesModal.vue' @@ -466,7 +519,8 @@ import TLSFingerprintProfilesModal from '@/components/admin/TLSFingerprintProfil import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { formatDateTime, formatRelativeTime } from '@/utils/format' import { proxyExpiryBadgeClass, proxyExpiryLabelKey } from '@/utils/proxyExpiry' -import type { Account, AccountPlatform, AccountSchedulerGroupScore, AccountType, Proxy as AccountProxy, AdminGroup, WindowStats, ClaudeModel } from '@/types' +import { extractApiErrorMessage } from '@/utils/apiError' +import type { Account, AccountPlatform, AccountSchedulerGroupScore, AccountType, Proxy as AccountProxy, AdminGroup, WindowStats, ClaudeModel, UpstreamBillingProbeSettings, UpstreamBillingProbeSnapshot } from '@/types' const { t } = useI18n() const appStore = useAppStore() @@ -544,6 +598,15 @@ const scheduleModelOptions = ref([]) const togglingSchedulable = ref(null) const menu = reactive<{show:boolean, acc:Account|null, pos:{top:number, left:number}|null}>({ show: false, acc: null, pos: null }) const exportingData = ref(false) +const upstreamBillingProbeSettings = reactive({ + enabled: true, + interval_minutes: 30 +}) +const upstreamBillingSettingsLoading = ref(false) +const upstreamBillingSettingsSaving = ref(false) +const probingUpstreamBilling = reactive(new Set()) +const upstreamBillingNow = ref(Date.now()) +useIntervalFn(() => { upstreamBillingNow.value = Date.now() }, 60_000) // Account tools dropdown const showAccountToolsDropdown = ref(false) @@ -1113,6 +1176,31 @@ const openTLSFingerprintProfiles = () => { showTLSFingerprintProfiles.value = true } +const loadUpstreamBillingProbeSettings = async () => { + upstreamBillingSettingsLoading.value = true + try { + Object.assign(upstreamBillingProbeSettings, await adminAPI.accounts.getUpstreamBillingProbeSettings()) + } catch (error) { + console.error('Failed to load upstream billing probe settings:', error) + } finally { + upstreamBillingSettingsLoading.value = false + } +} + +const saveUpstreamBillingProbeSettings = async () => { + upstreamBillingSettingsSaving.value = true + try { + const saved = await adminAPI.accounts.updateUpstreamBillingProbeSettings({ ...upstreamBillingProbeSettings }) + Object.assign(upstreamBillingProbeSettings, saved) + appStore.showSuccess(t('admin.accounts.upstreamBilling.settingsSaved')) + } catch (error) { + console.error('Failed to save upstream billing probe settings:', error) + appStore.showError(extractApiErrorMessage(error, t('admin.accounts.upstreamBilling.settingsFailed'))) + } finally { + upstreamBillingSettingsSaving.value = false + } +} + const syncPendingListChanges = async () => { hasPendingListSync.value = false await load() @@ -1282,6 +1370,7 @@ const allColumns = computed(() => { { key: 'priority', label: t('admin.accounts.columns.priority'), sortable: true }, { key: 'scheduler_score', label: t('admin.accounts.columns.schedulerScore'), sortable: false }, { key: 'rate_multiplier', label: t('admin.accounts.columns.billingRateMultiplier'), sortable: true }, + { key: 'upstream_billing_rate', label: t('admin.accounts.columns.upstreamBillingRate'), sortable: false }, { key: 'last_used_at', label: t('admin.accounts.columns.lastUsed'), sortable: true }, { key: 'created_at', label: t('admin.accounts.columns.createdAt'), sortable: true }, { key: 'expires_at', label: t('admin.accounts.columns.expiresAt'), sortable: true }, @@ -1392,6 +1481,35 @@ const handleBulkRefreshToken = async () => { appStore.showError(String(error)) } } +const handleBulkProbeUpstreamBilling = async () => { + const accountIDs = [...selIds.value] + if (accountIDs.length === 0) { + appStore.showError(t('admin.accounts.upstreamBilling.noEligibleAccounts')) + return + } + if (accountIDs.length > 20) { + appStore.showError(t('admin.accounts.upstreamBilling.batchLimit')) + return + } + accountIDs.forEach(id => probingUpstreamBilling.add(id)) + try { + const results = await adminAPI.accounts.probeUpstreamBillingBatch(accountIDs) + results.forEach(result => { + if (result.snapshot) patchUpstreamBillingSnapshot(result.account_id, result.snapshot) + }) + const failed = results.filter(result => result.error).length + if (failed > 0) { + appStore.showError(t('admin.accounts.upstreamBilling.batchPartial', { success: results.length - failed, failed })) + } else { + appStore.showSuccess(t('admin.accounts.upstreamBilling.batchCompleted', { count: results.length })) + } + } catch (error) { + console.error('Failed to probe upstream billing in batch:', error) + appStore.showError(extractApiErrorMessage(error, t('admin.accounts.upstreamBilling.probeFailed'))) + } finally { + accountIDs.forEach(id => probingUpstreamBilling.delete(id)) + } +} const updateSchedulableInList = (accountIds: number[], schedulable: boolean) => { if (accountIds.length === 0) return const idSet = new Set(accountIds) @@ -1641,6 +1759,27 @@ const patchAccountInList = (updatedAccount: Account) => { accounts.value = nextAccounts syncAccountRefs(mergedAccount) } +const patchUpstreamBillingSnapshot = (accountID: number, snapshot: UpstreamBillingProbeSnapshot) => { + const account = accounts.value.find(item => item.id === accountID) + if (!account) return + patchAccountInList({ + ...account, + extra: { ...account.extra, upstream_billing_probe: snapshot } + }) +} +const handleProbeUpstreamBilling = async (account: Account) => { + if (probingUpstreamBilling.has(account.id)) return + probingUpstreamBilling.add(account.id) + try { + const result = await adminAPI.accounts.probeUpstreamBilling(account.id) + if (result.snapshot) patchUpstreamBillingSnapshot(account.id, result.snapshot) + } catch (error) { + console.error('Failed to probe upstream billing:', error) + appStore.showError(extractApiErrorMessage(error, t('admin.accounts.upstreamBilling.probeFailed'))) + } finally { + probingUpstreamBilling.delete(account.id) + } +} const handleAccountUpdated = (updatedAccount: Account) => { patchAccountInList(updatedAccount) enterAutoRefreshSilentWindow() @@ -1885,6 +2024,7 @@ const handleClickOutside = (event: MouseEvent) => { onMounted(async () => { load() + loadUpstreamBillingProbeSettings() try { const [p, g] = await Promise.all([adminAPI.proxies.getAll(), adminAPI.groups.getAll()]) proxies.value = p diff --git a/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts index ed53a9e5df..89ccdf316d 100644 --- a/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts +++ b/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts @@ -8,13 +8,15 @@ const { listWithEtag, getBatchTodayStats, getAllProxies, - getAllGroups + getAllGroups, + probeUpstreamBillingBatch } = vi.hoisted(() => ({ listAccounts: vi.fn(), listWithEtag: vi.fn(), getBatchTodayStats: vi.fn(), getAllProxies: vi.fn(), - getAllGroups: vi.fn() + getAllGroups: vi.fn(), + probeUpstreamBillingBatch: vi.fn() })) vi.mock('@/api/admin', () => ({ @@ -23,9 +25,11 @@ vi.mock('@/api/admin', () => ({ list: listAccounts, listWithEtag, getBatchTodayStats, + getUpstreamBillingProbeSettings: vi.fn().mockResolvedValue({ enabled: true, interval_minutes: 30 }), delete: vi.fn(), batchClearError: vi.fn(), batchRefresh: vi.fn(), + probeUpstreamBillingBatch, toggleSchedulable: vi.fn() }, proxies: { @@ -67,6 +71,7 @@ const DataTableStub = {
{{ column.key }}
+
@@ -75,8 +80,18 @@ const DataTableStub = { const AccountBulkActionsBarStub = { props: ['selectedIds'], - emits: ['edit-filtered'], - template: '' + emits: ['edit-filtered', 'probe-upstream-billing'], + template: ` +
+ + +
+ ` +} + +const PaginationStub = { + emits: ['update:page'], + template: '' } const BulkEditAccountModalStub = { @@ -93,6 +108,7 @@ describe('admin AccountsView bulk edit scope', () => { getBatchTodayStats.mockReset() getAllProxies.mockReset() getAllGroups.mockReset() + probeUpstreamBillingBatch.mockReset() listAccounts.mockResolvedValue({ items: [], @@ -109,6 +125,7 @@ describe('admin AccountsView bulk edit scope', () => { getBatchTodayStats.mockResolvedValue({ stats: {} }) getAllProxies.mockResolvedValue([]) getAllGroups.mockResolvedValue([]) + probeUpstreamBillingBatch.mockResolvedValue([]) }) it('opens bulk edit in filtered-results mode from the bulk actions dropdown', async () => { @@ -224,4 +241,65 @@ describe('admin AccountsView bulk edit scope', () => { sortable: true }) }) + + it('submits selected account IDs from every page for backend eligibility checks', async () => { + const account = (id: number) => ({ + id, + name: `account-${id}`, + platform: 'openai', + type: 'apikey', + status: 'active', + schedulable: true, + created_at: '2026-07-13T00:00:00Z', + updated_at: '2026-07-13T00:00:00Z' + }) + listAccounts + .mockResolvedValueOnce({ items: [account(7)], total: 2, page: 1, page_size: 1, pages: 2 }) + .mockResolvedValueOnce({ items: [account(11)], total: 2, page: 2, page_size: 1, pages: 2 }) + + const wrapper = mount(AccountsView, { + global: { + stubs: { + AppLayout: { template: '
' }, + TablePageLayout: { template: '
' }, + DataTable: DataTableStub, + Pagination: PaginationStub, + ConfirmDialog: true, + AccountTableActions: true, + AccountTableFilters: true, + AccountBulkActionsBar: AccountBulkActionsBarStub, + AccountActionMenu: true, + ImportDataModal: true, + ReAuthAccountModal: true, + AccountTestModal: true, + AccountStatsModal: true, + ScheduledTestsPanel: true, + SyncFromCrsModal: true, + TempUnschedStatusModal: true, + ErrorPassthroughRulesModal: true, + TLSFingerprintProfilesModal: true, + CreateAccountModal: true, + EditAccountModal: true, + BulkEditAccountModal: BulkEditAccountModalStub, + PlatformTypeBadge: true, + AccountCapacityCell: true, + AccountStatusIndicator: true, + AccountTodayStatsCell: true, + AccountGroupsCell: true, + AccountUsageCell: true, + Icon: true + } + } + }) + + await flushPromises() + await wrapper.get('[data-test="select-row"] input').trigger('change') + await wrapper.get('[data-test="next-page"]').trigger('click') + await flushPromises() + await wrapper.get('[data-test="select-row"] input').trigger('change') + await wrapper.get('[data-test="probe-upstream-billing"]').trigger('click') + await flushPromises() + + expect(probeUpstreamBillingBatch).toHaveBeenCalledWith([7, 11]) + }) }) diff --git a/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts index e74af50182..c63087b393 100644 --- a/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts +++ b/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts @@ -23,6 +23,7 @@ vi.mock('@/api/admin', () => ({ list: listAccounts, listWithEtag, getBatchTodayStats, + getUpstreamBillingProbeSettings: vi.fn().mockResolvedValue({ enabled: true, interval_minutes: 30 }), delete: vi.fn(), batchClearError: vi.fn(), batchRefresh: vi.fn(), diff --git a/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts index 4eaa665b72..aedf74ce96 100644 --- a/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts +++ b/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts @@ -37,6 +37,7 @@ vi.mock('@/api/admin', () => ({ listWithEtag, getBatchTodayStats, duplicate: duplicateAccount, + getUpstreamBillingProbeSettings: vi.fn().mockResolvedValue({ enabled: true, interval_minutes: 30 }), createSparkShadow, delete: vi.fn(), batchClearError: vi.fn(), diff --git a/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts index 81e7d87e0e..0edd426ac1 100644 --- a/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts +++ b/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts @@ -23,6 +23,7 @@ vi.mock('@/api/admin', () => ({ list: listAccounts, listWithEtag, getBatchTodayStats, + getUpstreamBillingProbeSettings: vi.fn().mockResolvedValue({ enabled: true, interval_minutes: 30 }), delete: vi.fn(), batchClearError: vi.fn(), batchRefresh: vi.fn(), @@ -70,6 +71,9 @@ const DataTableStub = {
+
+ +
` @@ -161,4 +165,16 @@ describe('admin AccountsView usage windows hint', () => { expect(hint.exists()).toBe(true) expect(hint.text()).toBe('admin.accounts.usageWindowsHint') }) + + it('renders the upstream billing trust warning next to the declared-rate column', async () => { + const wrapper = mountView() + await flushPromises() + + const header = wrapper.find('[data-test="upstream-billing-header"]') + expect(header.exists()).toBe(true) + expect(header.text()).toContain('admin.accounts.columns.upstreamBillingRate') + expect(wrapper.findAll('[data-test="usage-windows-hint"]').some(node => + node.text() === 'admin.accounts.upstreamBilling.trustWarning' + )).toBe(true) + }) })