feat: 增加上游 Sub2API 计费倍率探测与账号展示

This commit is contained in:
Tian Lee
2026-07-15 23:58:46 +08:00
parent eb2b8632de
commit 0765d10c1d
49 changed files with 5747 additions and 74 deletions
+7
View File
@@ -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{
+10 -2
View File
@@ -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{
+1
View File
@@ -83,6 +83,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
nil, // paymentOrderExpiry
nil, // channelMonitorRunner
nil, // quotaFlusher
nil, // upstreamBillingProbe
)
require.NotPanics(t, func() {
@@ -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
@@ -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)})
}
@@ -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)
}
+2
View File
@@ -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,
+499 -28
View File
@@ -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, &currentEnabled, &currentSnapshot); 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(&current.protocol, &current.host, &current.port, &current.username, &current.password, &current.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"')`
@@ -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())
}
@@ -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)
}
@@ -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)
}
@@ -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,
}))
}
@@ -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,
)
}
+23 -3
View File
@@ -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)
@@ -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", "")
+163 -9
View File
@@ -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)
}
@@ -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())
}
@@ -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
}
+5
View File
@@ -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)
+99 -3
View File
@@ -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
@@ -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)
}
@@ -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", "")
@@ -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)
})
}
}
+64 -10
View File
@@ -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))
}
@@ -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)
// =========================
@@ -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
}
@@ -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")
}
}
@@ -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 {
@@ -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))
})
}
}
+11 -5
View File
@@ -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
}
@@ -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 &copy, nil
}
func (s *updatingProxyRepoStub) Update(_ context.Context, proxy *Proxy) error {
s.updateCalls++
copy := *proxy
s.proxy = &copy
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)
})
}
@@ -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"
}
@@ -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")
}
+1
View File
@@ -664,6 +664,7 @@ var ProviderSet = wire.NewSet(
ProvideRateLimitService,
ProvideAccountUsageService,
ProvideAccountTestService,
ProvideUpstreamBillingProbeService,
ProvideSettingService,
NewDataManagementService,
ProvideBackupService,
@@ -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] })
})
})
+41 -2
View File
@@ -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<UpstreamBillingProbeSettings> {
const { data } = await apiClient.get<UpstreamBillingProbeSettings>('/admin/accounts/upstream-billing-probe/settings')
return data
}
export async function updateUpstreamBillingProbeSettings(
settings: UpstreamBillingProbeSettings
): Promise<UpstreamBillingProbeSettings> {
const { data } = await apiClient.put<UpstreamBillingProbeSettings>(
'/admin/accounts/upstream-billing-probe/settings',
settings
)
return data
}
export async function setUpstreamBillingProbeEnabled(id: number, enabled: boolean): Promise<void> {
await apiClient.put(`/admin/accounts/${id}/upstream-billing-probe`, { enabled })
}
export async function probeUpstreamBilling(id: number): Promise<UpstreamBillingProbeResult> {
const { data } = await apiClient.post<UpstreamBillingProbeResult>(`/admin/accounts/${id}/upstream-billing-probe`)
return data
}
export async function probeUpstreamBillingBatch(accountIds: number[]): Promise<UpstreamBillingProbeResult[]> {
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
@@ -1592,6 +1592,23 @@
</div>
</div>
<div
v-if="account?.platform === 'openai' && account?.type === 'apikey'"
class="flex items-center justify-between gap-4 border-t border-gray-200 pt-4 dark:border-dark-600"
>
<div>
<label class="input-label mb-0">{{ t('admin.accounts.upstreamBilling.autoProbe') }}</label>
<p class="mt-1 text-xs text-gray-500 dark:text-gray-400">
{{ t('admin.accounts.upstreamBilling.autoProbeHint') }}
</p>
</div>
<Toggle
v-model="upstreamBillingAutoProbeEnabled"
data-testid="upstream-billing-auto-probe"
:aria-label="t('admin.accounts.upstreamBilling.autoProbe')"
/>
</div>
<!-- Anthropic API Key 自动透传开关 -->
<div
v-if="account?.platform === 'anthropic' && account?.type === 'apikey'"
@@ -2560,6 +2577,7 @@ import type {
import BaseDialog from '@/components/common/BaseDialog.vue'
import ConfirmDialog from '@/components/common/ConfirmDialog.vue'
import Select from '@/components/common/Select.vue'
import Toggle from '@/components/common/Toggle.vue'
import Icon from '@/components/icons/Icon.vue'
import ProxySelector from '@/components/common/ProxySelector.vue'
import ProxyAdBanner from '@/components/common/ProxyAdBanner.vue'
@@ -2729,6 +2747,7 @@ const autoPause5hThreshold = ref<number | null>(null)
const autoPause7dThreshold = ref<number | null>(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
@@ -0,0 +1,160 @@
<template>
<div v-if="eligible" class="flex h-6 min-w-[7rem] items-center gap-1">
<HelpTooltip class="-ml-1" width-class="w-64" data-testid="upstream-billing-details">
<template #trigger>
<span
class="cursor-help border-b border-dotted border-gray-300 font-mono text-sm font-medium text-gray-800 dark:border-gray-600 dark:text-gray-200"
data-testid="upstream-billing-rate"
>
{{ effectiveRate }}
</span>
</template>
<div v-if="data" class="space-y-1">
<p>{{ t('admin.accounts.upstreamBilling.groupRate', { value: data.group_rate_multiplier }) }}</p>
<p v-if="data.user_rate_multiplier != null">
{{ t('admin.accounts.upstreamBilling.userRate', { value: data.user_rate_multiplier }) }}
</p>
<p>
{{
data.peak_rate_enabled
? t('admin.accounts.upstreamBilling.peakRate', {
start: data.peak_start,
end: data.peak_end,
value: data.peak_rate_multiplier,
timezone: data.timezone
})
: t('admin.accounts.upstreamBilling.noPeakRate')
}}
</p>
<p>{{ t('admin.accounts.upstreamBilling.effectiveRate', { value: currentEffectiveRate ?? '-' }) }}</p>
<p>{{ t('admin.accounts.upstreamBilling.updatedAt', { value: formatDate(snapshot?.received_at) }) }}</p>
</div>
<span v-else>{{ statusLabel }}</span>
</HelpTooltip>
<span v-if="statusLabel" :class="statusClass" class="whitespace-nowrap text-[10px] font-medium">
{{ statusLabel }}
</span>
<button
type="button"
class="inline-flex h-6 w-6 flex-shrink-0 items-center justify-center rounded text-blue-600 transition-colors hover:bg-blue-50 disabled:cursor-not-allowed disabled:opacity-50 dark:text-blue-400 dark:hover:bg-blue-900/30"
:disabled="probing"
:aria-label="t('admin.accounts.upstreamBilling.manualProbe')"
:title="t('admin.accounts.upstreamBilling.manualProbe')"
data-testid="upstream-billing-probe"
@click="$emit('probe')"
>
<Icon name="refresh" size="xs" :class="{ 'animate-spin': probing }" />
</button>
</div>
<span v-else class="text-sm text-gray-400 dark:text-dark-500">-</span>
</template>
<script setup lang="ts">
import { computed } from 'vue'
import { useI18n } from 'vue-i18n'
import HelpTooltip from '@/components/common/HelpTooltip.vue'
import Icon from '@/components/icons/Icon.vue'
import type { Account, UpstreamBillingProbeSnapshot } from '@/types'
const props = defineProps<{
account: Account
intervalMinutes: number
now: number
probing?: boolean
}>()
defineEmits<{
(event: 'probe'): void
}>()
const { t } = useI18n()
const eligible = computed(() => props.account.platform === 'openai' && props.account.type === 'apikey')
const snapshot = computed<UpstreamBillingProbeSnapshot | undefined>(() => props.account.extra?.upstream_billing_probe)
const data = computed(() => snapshot.value?.data)
const receivedAt = computed(() => typeof snapshot.value?.received_at === 'string' ? Date.parse(snapshot.value.received_at) : Number.NaN)
const freshUntil = computed(() => {
if (typeof snapshot.value?.fresh_until === 'string') return Date.parse(snapshot.value.fresh_until)
if (snapshot.value?.status !== 'ok' || typeof snapshot.value.next_probe_at !== 'string') return Number.NaN
const nextProbeAt = Date.parse(snapshot.value.next_probe_at)
return Number.isFinite(nextProbeAt) && nextProbeAt > receivedAt.value
? receivedAt.value + 2 * (nextProbeAt - receivedAt.value)
: Number.NaN
})
const validTimestamps = computed(() => {
if (!Number.isFinite(receivedAt.value) || receivedAt.value > props.now) return false
return Number.isFinite(freshUntil.value) && freshUntil.value > receivedAt.value
})
const stale = computed(() => {
if (!snapshot.value) return false
if (!Number.isFinite(receivedAt.value)) return snapshot.value.status === 'ok'
if (!validTimestamps.value) return true
return props.now > freshUntil.value
})
const parseMinute = (value?: string) => {
if (typeof value !== 'string') return null
const match = /^(\d{2}):(\d{2})$/.exec(value)
if (!match) return null
const hour = Number(match[1])
const minute = Number(match[2])
return hour < 24 && minute < 60 ? hour * 60 + minute : null
}
const minuteInTimeZone = (timestamp: number, timeZone?: string) => {
if (!timeZone) return null
try {
const parts = new Intl.DateTimeFormat('en-GB', {
timeZone,
hour: '2-digit',
minute: '2-digit',
hourCycle: 'h23'
}).formatToParts(new Date(timestamp))
const hour = Number(parts.find(part => part.type === 'hour')?.value)
const minute = Number(parts.find(part => part.type === 'minute')?.value)
return Number.isInteger(hour) && Number.isInteger(minute) ? hour * 60 + minute : null
} catch {
return null
}
}
const currentEffectiveRate = computed(() => {
const billing = data.value
if (!billing) return null
if (billing.billing_scope !== 'token') return null
const base = billing.resolved_rate_multiplier
if (typeof base !== 'number' || !Number.isFinite(base) || base < 0) return null
if (typeof billing.peak_rate_enabled !== 'boolean') return null
if (!billing.peak_rate_enabled) return base
const start = parseMinute(billing.peak_start)
const end = parseMinute(billing.peak_end)
const minute = minuteInTimeZone(props.now, billing.timezone)
const peak = billing.peak_rate_multiplier
if (start == null || end == null || minute == null || start >= end || typeof peak !== 'number' || !Number.isFinite(peak) || peak < 0) return null
const value = minute >= start && minute < end ? base * peak : base
return Number.isFinite(value) ? value : null
})
const effectiveRate = computed(() => {
if (!validTimestamps.value || stale.value || !['ok', 'failed'].includes(snapshot.value?.status ?? '')) return '-'
const value = currentEffectiveRate.value
return value == null ? '-' : `${Number(value.toPrecision(12))}x`
})
const statusLabel = computed(() => {
if (!snapshot.value) return t('admin.accounts.upstreamBilling.notProbed')
if (snapshot.value.status === 'unsupported') return t('admin.accounts.upstreamBilling.unsupported')
if (stale.value) return t('admin.accounts.upstreamBilling.stale')
if (snapshot.value.status === 'failed') return t('admin.accounts.upstreamBilling.failed')
return ''
})
const statusClass = computed(() => {
if (!snapshot.value) return 'text-gray-400 dark:text-gray-500'
if (snapshot.value.status === 'unsupported') return 'text-gray-500 dark:text-gray-400'
if (stale.value) return 'text-amber-600 dark:text-amber-400'
if (snapshot.value.status === 'failed') return 'text-red-600 dark:text-red-400'
return ''
})
const formatDate = (value?: string) => value
? new Date(value).toLocaleString(undefined, {
month: '2-digit',
day: '2-digit',
hour: '2-digit',
minute: '2-digit'
})
: '-'
</script>
@@ -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 = {
@@ -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<typeof import('vue-i18n')>('vue-i18n')
return {
...actual,
useI18n: () => ({
t: (key: string, params?: Record<string, unknown>) =>
params ? `${key}:${Object.values(params).join(',')}` : key
})
}
})
const makeAccount = (overrides: Partial<Account> = {}): 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<typeof billingData> = {},
snapshotOverrides: Record<string, unknown> = {}
) => 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')
})
})
@@ -28,6 +28,7 @@
<button @click="$emit('delete')" class="btn btn-danger btn-sm">{{ t('admin.accounts.bulkActions.delete') }}</button>
<button @click="$emit('reset-status')" class="btn btn-secondary btn-sm">{{ t('admin.accounts.bulkActions.resetStatus') }}</button>
<button @click="$emit('refresh-token')" class="btn btn-secondary btn-sm">{{ t('admin.accounts.bulkActions.refreshToken') }}</button>
<button @click="$emit('probe-upstream-billing')" class="btn btn-secondary btn-sm">{{ t('admin.accounts.bulkActions.probeUpstreamBilling') }}</button>
<button @click="$emit('toggle-schedulable', true)" class="btn btn-success btn-sm">{{ t('admin.accounts.bulkActions.enableScheduling') }}</button>
<button @click="$emit('toggle-schedulable', false)" class="btn btn-warning btn-sm">{{ t('admin.accounts.bulkActions.disableScheduling') }}</button>
<button @click="$emit('edit-selected')" class="btn btn-primary btn-sm">{{ t('admin.accounts.bulkActions.edit') }}</button>
@@ -41,5 +42,19 @@
<script setup lang="ts">
import { useI18n } from 'vue-i18n'
defineProps(['selectedIds']); defineEmits(['delete', 'edit-selected', 'edit-filtered', 'clear', 'select-page', 'toggle-schedulable', 'reset-status', 'refresh-token']); const { t } = useI18n()
defineProps<{ selectedIds: number[] }>()
defineEmits([
'delete',
'edit-selected',
'edit-filtered',
'clear',
'select-page',
'toggle-schedulable',
'reset-status',
'refresh-token',
'probe-upstream-billing'
])
const { t } = useI18n()
</script>
@@ -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'
@@ -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} 失败'
+44
View File
@@ -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<string, { rate_limited_at: string; rate_limit_reset_at: string }>
antigravity_credits_overages?: Record<string, { activated_at: string; active_until: string }>
upstream_billing_probe_enabled?: boolean
upstream_billing_probe?: UpstreamBillingProbeSnapshot
} & Record<string, unknown>)
proxy_id: number | null
proxy_fallback_origin_id?: number | null
+141 -1
View File
@@ -132,6 +132,41 @@
<span class="flex-1 text-left">{{ t('admin.tlsFingerprintProfiles.title') }}</span>
</button>
<div class="my-2 border-t border-gray-100 dark:border-gray-700"></div>
<div class="space-y-2 px-3 py-2">
<div class="flex items-center justify-between gap-3">
<span class="text-sm font-medium text-gray-700 dark:text-gray-200">
{{ t('admin.accounts.upstreamBilling.autoProbeSettings') }}
</span>
<Toggle
v-model="upstreamBillingProbeSettings.enabled"
:aria-label="t('admin.accounts.upstreamBilling.autoProbeSettings')"
/>
</div>
<div class="flex items-center gap-2">
<label class="flex-1 text-xs text-gray-500 dark:text-gray-400" for="upstream-billing-probe-interval">
{{ t('admin.accounts.upstreamBilling.intervalMinutes') }}
</label>
<input
id="upstream-billing-probe-interval"
v-model.number="upstreamBillingProbeSettings.interval_minutes"
type="number"
min="5"
max="1440"
class="input h-8 w-20 px-2 text-sm"
/>
<button
type="button"
class="btn btn-secondary h-8 px-2"
:disabled="upstreamBillingSettingsLoading || upstreamBillingSettingsSaving"
:title="t('common.save')"
@click="saveUpstreamBillingProbeSettings"
>
<Icon name="check" size="sm" />
</button>
</div>
</div>
<div class="my-2 border-t border-gray-100 dark:border-gray-700"></div>
<div class="px-2 py-2">
<div class="flex items-center justify-between gap-3">
@@ -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
</span>
</template>
<template #header-upstream_billing_rate="{ column }">
<div class="flex items-center">
<span>{{ column.label }}</span>
<HelpTooltip :content="t('admin.accounts.upstreamBilling.trustWarning')" width-class="w-80" />
</div>
</template>
<template #cell-upstream_billing_rate="{ row }">
<UpstreamBillingRateCell
:account="row"
:interval-minutes="upstreamBillingProbeSettings.interval_minutes"
:now="upstreamBillingNow"
:probing="probingUpstreamBilling.has(row.id)"
@probe="handleProbeUpstreamBilling(row)"
/>
</template>
<template #cell-priority="{ value }">
<span class="text-sm text-gray-700 dark:text-gray-300">{{ value }}</span>
</template>
@@ -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<SelectOption[]>([])
const togglingSchedulable = ref<number | null>(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<UpstreamBillingProbeSettings>({
enabled: true,
interval_minutes: 30
})
const upstreamBillingSettingsLoading = ref(false)
const upstreamBillingSettingsSaving = ref(false)
const probingUpstreamBilling = reactive(new Set<number>())
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
@@ -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 = {
<div data-test="data-table">
<span v-for="column in columns" :key="column.key" data-test="column-key">{{ column.key }}</span>
<div v-for="row in data" :key="row.id">
<div data-test="select-row"><slot name="cell-select" :row="row" /></div>
<slot name="cell-created_at" :value="row.created_at" :row="row" />
</div>
</div>
@@ -75,8 +80,18 @@ const DataTableStub = {
const AccountBulkActionsBarStub = {
props: ['selectedIds'],
emits: ['edit-filtered'],
template: '<button data-test="edit-filtered" @click="$emit(\'edit-filtered\')">edit filtered</button>'
emits: ['edit-filtered', 'probe-upstream-billing'],
template: `
<div>
<button data-test="edit-filtered" @click="$emit('edit-filtered')">edit filtered</button>
<button data-test="probe-upstream-billing" @click="$emit('probe-upstream-billing')">probe</button>
</div>
`
}
const PaginationStub = {
emits: ['update:page'],
template: '<button data-test="next-page" @click="$emit(\'update:page\', 2)">next</button>'
}
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: '<div><slot /></div>' },
TablePageLayout: { template: '<div><slot name="table" /><slot name="pagination" /></div>' },
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])
})
})
@@ -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(),
@@ -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(),
@@ -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 = {
<div v-if="column.key === 'usage'" data-test="usage-header">
<slot :name="'header-' + column.key" :column="column" />
</div>
<div v-if="column.key === 'upstream_billing_rate'" data-test="upstream-billing-header">
<slot :name="'header-' + column.key" :column="column" />
</div>
</template>
</div>
`
@@ -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)
})
})