mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 14:58:23 +08:00
feat: 增加上游 Sub2API 计费倍率探测与账号展示
This commit is contained in:
@@ -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{
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -57,6 +57,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{
|
||||
"codex_5h_",
|
||||
"codex_7d_",
|
||||
"passive_usage_",
|
||||
"upstream_billing_probe",
|
||||
}
|
||||
|
||||
var schedulerNeutralExtraKeys = map[string]struct{}{
|
||||
@@ -395,21 +396,80 @@ func (r *accountRepository) ListCRSAccountIDs(ctx context.Context) (map[string]i
|
||||
}
|
||||
|
||||
func (r *accountRepository) Update(ctx context.Context, account *service.Account) error {
|
||||
return r.updateAccount(ctx, account, nil)
|
||||
}
|
||||
|
||||
// UpdateWithUpstreamBillingProbeEnabled applies an explicit probe switch in the
|
||||
// same row-lock transaction as the rest of an admin account edit.
|
||||
func (r *accountRepository) UpdateWithUpstreamBillingProbeEnabled(ctx context.Context, account *service.Account, enabled bool) error {
|
||||
return r.updateAccount(ctx, account, &enabled)
|
||||
}
|
||||
|
||||
func (r *accountRepository) updateAccount(ctx context.Context, account *service.Account, explicitProbeEnabled *bool) error {
|
||||
if account == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
baseCtx := ctx
|
||||
contextTx := dbent.TxFromContext(ctx)
|
||||
client := r.client
|
||||
var tx *dbent.Tx
|
||||
if contextTx != nil {
|
||||
client = contextTx.Client()
|
||||
} else {
|
||||
var err error
|
||||
tx, err = r.client.Tx(ctx)
|
||||
if err != nil && !errors.Is(err, dbent.ErrTxStarted) {
|
||||
return err
|
||||
}
|
||||
if tx != nil {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
ctx = dbent.NewTxContext(ctx, tx)
|
||||
client = tx.Client()
|
||||
}
|
||||
}
|
||||
|
||||
updated, err := r.updateLockedAccount(ctx, client, account, explicitProbeEnabled)
|
||||
if err != nil {
|
||||
return translatePersistenceError(err, service.ErrAccountNotFound, nil)
|
||||
}
|
||||
if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, buildSchedulerGroupPayload(account.GroupIDs)); err != nil {
|
||||
return err
|
||||
}
|
||||
if tx != nil {
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
account.UpdatedAt = updated.UpdatedAt
|
||||
// 普通账号编辑(如 model_mapping / credentials)也需要立即刷新单账号快照,
|
||||
// 否则网关在 outbox worker 延迟或异常时仍可能读到旧配置。
|
||||
if contextTx == nil {
|
||||
r.syncSchedulerAccountSnapshot(baseCtx, account.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *accountRepository) updateLockedAccount(ctx context.Context, client *dbent.Client, account *service.Account, explicitProbeEnabled *bool) (*dbent.Account, error) {
|
||||
extra, err := lockAndMergeAccountProbeExtra(ctx, client, account, explicitProbeEnabled)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
account.Extra = extra
|
||||
|
||||
schedulable := account.Schedulable
|
||||
if account.Status == service.StatusError {
|
||||
schedulable = false
|
||||
}
|
||||
|
||||
builder := r.client.Account.UpdateOneID(account.ID).
|
||||
builder := client.Account.UpdateOneID(account.ID).
|
||||
SetName(account.Name).
|
||||
SetNillableNotes(account.Notes).
|
||||
SetPlatform(account.Platform).
|
||||
SetType(account.Type).
|
||||
SetCredentials(normalizeJSONMap(account.Credentials)).
|
||||
SetExtra(normalizeJSONMap(account.Extra)).
|
||||
SetExtra(extra).
|
||||
SetConcurrency(account.Concurrency).
|
||||
SetPriority(account.Priority).
|
||||
SetStatus(account.Status).
|
||||
@@ -478,31 +538,140 @@ func (r *accountRepository) Update(ctx context.Context, account *service.Account
|
||||
builder.SetQuotaDimension(dbaccount.QuotaDimension(account.QuotaDimensionOrDefault()))
|
||||
builder.SetNillableParentAccountID(account.ParentAccountID)
|
||||
|
||||
updated, err := builder.Save(ctx)
|
||||
return builder.Save(ctx)
|
||||
}
|
||||
|
||||
func lockAndMergeAccountProbeExtra(ctx context.Context, client *dbent.Client, account *service.Account, explicitProbeEnabled *bool) (map[string]any, error) {
|
||||
credentials, err := json.Marshal(normalizeJSONMap(account.Credentials))
|
||||
if err != nil {
|
||||
return translatePersistenceError(err, service.ErrAccountNotFound, nil)
|
||||
return nil, err
|
||||
}
|
||||
account.UpdatedAt = updated.UpdatedAt
|
||||
if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, buildSchedulerGroupPayload(account.GroupIDs)); err != nil {
|
||||
logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue account update failed: account=%d err=%v", account.ID, err)
|
||||
var proxyID any
|
||||
if account.ProxyID != nil {
|
||||
proxyID = *account.ProxyID
|
||||
}
|
||||
// 普通账号编辑(如 model_mapping / credentials)也需要立即刷新单账号快照,
|
||||
// 否则网关在 outbox worker 延迟或异常时仍可能读到旧配置。
|
||||
r.syncSchedulerAccountSnapshot(ctx, account.ID)
|
||||
return nil
|
||||
rows, err := client.QueryContext(ctx, `
|
||||
SELECT
|
||||
platform = $2
|
||||
AND type = $3
|
||||
AND credentials = $4::jsonb
|
||||
AND proxy_id IS NOT DISTINCT FROM $5,
|
||||
extra -> 'upstream_billing_probe_enabled',
|
||||
extra -> 'upstream_billing_probe'
|
||||
FROM accounts
|
||||
WHERE id = $1 AND deleted_at IS NULL
|
||||
FOR NO KEY UPDATE
|
||||
`, account.ID, account.Platform, account.Type, string(credentials), proxyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, service.ErrAccountNotFound
|
||||
}
|
||||
|
||||
var (
|
||||
identityUnchanged bool
|
||||
currentEnabled []byte
|
||||
currentSnapshot []byte
|
||||
)
|
||||
if err := rows.Scan(&identityUnchanged, ¤tEnabled, ¤tSnapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
extra := copyJSONMap(normalizeJSONMap(account.Extra))
|
||||
delete(extra, service.UpstreamBillingProbeEnabledExtraKey)
|
||||
delete(extra, service.UpstreamBillingProbeExtraKey)
|
||||
probeExplicitlyDisabled := false
|
||||
probeAccount := account.Platform == service.PlatformOpenAI && account.Type == service.AccountTypeAPIKey
|
||||
if probeAccount && explicitProbeEnabled != nil {
|
||||
extra[service.UpstreamBillingProbeEnabledExtraKey] = *explicitProbeEnabled
|
||||
probeExplicitlyDisabled = !*explicitProbeEnabled
|
||||
} else if probeAccount && len(currentEnabled) > 0 && string(currentEnabled) != "null" {
|
||||
var enabled any
|
||||
if err := json.Unmarshal(currentEnabled, &enabled); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
extra[service.UpstreamBillingProbeEnabledExtraKey] = enabled
|
||||
if value, ok := enabled.(bool); ok && !value {
|
||||
probeExplicitlyDisabled = true
|
||||
}
|
||||
}
|
||||
if !identityUnchanged || probeExplicitlyDisabled || len(currentSnapshot) == 0 || string(currentSnapshot) == "null" {
|
||||
return extra, nil
|
||||
}
|
||||
var snapshot any
|
||||
if err := json.Unmarshal(currentSnapshot, &snapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
extra[service.UpstreamBillingProbeExtraKey] = snapshot
|
||||
return extra, nil
|
||||
}
|
||||
|
||||
func (r *accountRepository) UpdateCredentials(ctx context.Context, id int64, credentials map[string]any) error {
|
||||
_, err := r.client.Account.UpdateOneID(id).
|
||||
SetCredentials(normalizeJSONMap(credentials)).
|
||||
Save(ctx)
|
||||
payload, err := json.Marshal(normalizeJSONMap(credentials))
|
||||
if err != nil {
|
||||
return translatePersistenceError(err, service.ErrAccountNotFound, nil)
|
||||
return err
|
||||
}
|
||||
if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil {
|
||||
logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue credentials update failed: account=%d err=%v", id, err)
|
||||
baseCtx := ctx
|
||||
contextTx := dbent.TxFromContext(ctx)
|
||||
client := r.client
|
||||
var tx *dbent.Tx
|
||||
if contextTx != nil {
|
||||
client = contextTx.Client()
|
||||
} else if r.client != nil {
|
||||
var txErr error
|
||||
tx, txErr = r.client.Tx(ctx)
|
||||
if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) {
|
||||
return txErr
|
||||
}
|
||||
if tx != nil {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
ctx = dbent.NewTxContext(ctx, tx)
|
||||
client = tx.Client()
|
||||
}
|
||||
}
|
||||
result, err := client.ExecContext(ctx, `
|
||||
UPDATE accounts
|
||||
SET
|
||||
credentials = $1::jsonb,
|
||||
extra = CASE
|
||||
WHEN platform = 'openai'
|
||||
AND type = 'apikey'
|
||||
AND credentials IS DISTINCT FROM $1::jsonb
|
||||
THEN COALESCE(extra, '{}'::jsonb) - 'upstream_billing_probe'
|
||||
ELSE extra
|
||||
END,
|
||||
updated_at = NOW()
|
||||
WHERE id = $2 AND deleted_at IS NULL
|
||||
`, string(payload), id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
affected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if affected == 0 {
|
||||
return service.ErrAccountNotFound
|
||||
}
|
||||
if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if tx != nil {
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if contextTx == nil {
|
||||
r.syncSchedulerAccountSnapshot(baseCtx, id)
|
||||
}
|
||||
r.syncSchedulerAccountSnapshot(ctx, id)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -2117,10 +2286,31 @@ func (r *accountRepository) UpdateExtra(ctx context.Context, id int64, updates m
|
||||
return err
|
||||
}
|
||||
|
||||
clearProbeSnapshot := upstreamBillingProbeExplicitlyDisabled(updates) || upstreamBillingProbeSnapshotClearRequested(updates)
|
||||
durableSchedulerChange := shouldEnqueueSchedulerOutboxForExtraUpdates(updates) || clearProbeSnapshot
|
||||
baseCtx := ctx
|
||||
contextTx := dbent.TxFromContext(ctx)
|
||||
client := clientFromContext(ctx, r.client)
|
||||
var tx *dbent.Tx
|
||||
if durableSchedulerChange && contextTx == nil {
|
||||
var txErr error
|
||||
tx, txErr = r.client.Tx(ctx)
|
||||
if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) {
|
||||
return txErr
|
||||
}
|
||||
if tx != nil {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
ctx = dbent.NewTxContext(ctx, tx)
|
||||
client = tx.Client()
|
||||
}
|
||||
}
|
||||
extraExpression := "COALESCE(extra, '{}'::jsonb) || $1::jsonb"
|
||||
if clearProbeSnapshot {
|
||||
extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'"
|
||||
}
|
||||
result, err := client.ExecContext(
|
||||
ctx,
|
||||
"UPDATE accounts SET extra = COALESCE(extra, '{}'::jsonb) || $1::jsonb, updated_at = NOW() WHERE id = $2 AND deleted_at IS NULL",
|
||||
"UPDATE accounts SET extra = "+extraExpression+", updated_at = NOW() WHERE id = $2 AND deleted_at IS NULL",
|
||||
string(payload), id,
|
||||
)
|
||||
|
||||
@@ -2135,19 +2325,159 @@ func (r *accountRepository) UpdateExtra(ctx context.Context, id int64, updates m
|
||||
if affected == 0 {
|
||||
return service.ErrAccountNotFound
|
||||
}
|
||||
if shouldEnqueueSchedulerOutboxForExtraUpdates(updates) {
|
||||
if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil {
|
||||
logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue extra update failed: account=%d err=%v", id, err)
|
||||
if durableSchedulerChange {
|
||||
if err := enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if tx != nil {
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if contextTx == nil {
|
||||
r.syncSchedulerAccountSnapshot(baseCtx, id)
|
||||
}
|
||||
} else {
|
||||
// 观测型 extra 字段不需要触发 bucket 重建,但仍同步单账号快照,
|
||||
// 让 sticky session / GetAccount 命中缓存时也能读到最新数据,
|
||||
// 同时避免缓存局部 patch 覆盖掉并发写入的其它账号字段。
|
||||
r.syncSchedulerAccountSnapshot(ctx, id)
|
||||
if dbent.TxFromContext(ctx) == nil {
|
||||
r.syncSchedulerAccountSnapshot(ctx, id)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateUpstreamBillingProbeSnapshot stores a probe result only while the
|
||||
// network identity used by that probe is still current.
|
||||
func (r *accountRepository) UpdateUpstreamBillingProbeSnapshot(
|
||||
ctx context.Context,
|
||||
account *service.Account,
|
||||
snapshot *service.UpstreamBillingProbeSnapshot,
|
||||
) error {
|
||||
if account == nil || snapshot == nil {
|
||||
return service.ErrAccountNilInput
|
||||
}
|
||||
if dbent.TxFromContext(ctx) == nil {
|
||||
tx, err := r.client.Tx(ctx)
|
||||
if errors.Is(err, dbent.ErrTxStarted) {
|
||||
return r.updateUpstreamBillingProbeSnapshotInTx(ctx, account, snapshot)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
|
||||
if err := r.updateUpstreamBillingProbeSnapshotInTx(dbent.NewTxContext(ctx, tx), account, snapshot); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
// The durable outbox event is committed with the snapshot. This direct
|
||||
// cache write only reduces visibility latency on the current instance.
|
||||
r.syncSchedulerAccountSnapshot(ctx, account.ID)
|
||||
return nil
|
||||
}
|
||||
return r.updateUpstreamBillingProbeSnapshotInTx(ctx, account, snapshot)
|
||||
}
|
||||
|
||||
func (r *accountRepository) updateUpstreamBillingProbeSnapshotInTx(
|
||||
ctx context.Context,
|
||||
account *service.Account,
|
||||
snapshot *service.UpstreamBillingProbeSnapshot,
|
||||
) error {
|
||||
payload, err := json.Marshal(map[string]any{service.UpstreamBillingProbeExtraKey: snapshot})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
credentials, err := json.Marshal(account.Credentials)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var expectedSnapshot any
|
||||
if account.Extra != nil {
|
||||
expectedSnapshot = account.Extra[service.UpstreamBillingProbeExtraKey]
|
||||
}
|
||||
expectedSnapshotJSON, err := json.Marshal(expectedSnapshot)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var expectedEnabled any
|
||||
if account.Extra != nil {
|
||||
expectedEnabled = account.Extra[service.UpstreamBillingProbeEnabledExtraKey]
|
||||
}
|
||||
expectedEnabledJSON, err := json.Marshal(expectedEnabled)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client := clientFromContext(ctx, r.client)
|
||||
proxyMatches, err := lockAndMatchProbeProxyIdentity(ctx, client, account)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !proxyMatches {
|
||||
return service.ErrUpstreamBillingProbeIdentityChanged
|
||||
}
|
||||
var proxyID any
|
||||
if account.ProxyID != nil {
|
||||
proxyID = *account.ProxyID
|
||||
}
|
||||
result, err := client.ExecContext(ctx, `
|
||||
UPDATE accounts
|
||||
SET extra = COALESCE(extra, '{}'::jsonb) || $1::jsonb, updated_at = NOW()
|
||||
WHERE id = $2
|
||||
AND platform = $3
|
||||
AND type = $4
|
||||
AND credentials = $5::jsonb
|
||||
AND proxy_id IS NOT DISTINCT FROM $6
|
||||
AND COALESCE(extra -> 'upstream_billing_probe', 'null'::jsonb) = $7::jsonb
|
||||
AND COALESCE(extra -> 'upstream_billing_probe_enabled', 'null'::jsonb) = $8::jsonb
|
||||
AND deleted_at IS NULL
|
||||
`, string(payload), account.ID, account.Platform, account.Type, string(credentials), proxyID, string(expectedSnapshotJSON), string(expectedEnabledJSON))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
affected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if affected == 0 {
|
||||
return service.ErrUpstreamBillingProbeIdentityChanged
|
||||
}
|
||||
return enqueueSchedulerOutbox(ctx, client, service.SchedulerOutboxEventAccountChanged, &account.ID, nil, nil)
|
||||
}
|
||||
|
||||
func lockAndMatchProbeProxyIdentity(ctx context.Context, client *dbent.Client, account *service.Account) (bool, error) {
|
||||
if account.ProxyID == nil {
|
||||
return true, nil
|
||||
}
|
||||
rows, err := client.QueryContext(ctx, `
|
||||
SELECT protocol, host, port, COALESCE(username, ''), COALESCE(password, ''), status
|
||||
FROM proxies
|
||||
WHERE id = $1 AND deleted_at IS NULL
|
||||
FOR SHARE
|
||||
`, *account.ProxyID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return account.Proxy == nil, nil
|
||||
}
|
||||
if account.Proxy == nil || account.Proxy.ID != *account.ProxyID {
|
||||
return false, nil
|
||||
}
|
||||
var current proxyProbeIdentity
|
||||
if err := rows.Scan(¤t.protocol, ¤t.host, ¤t.port, ¤t.username, ¤t.password, ¤t.status); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return current == proxyProbeIdentityFromService(account.Proxy), rows.Err()
|
||||
}
|
||||
|
||||
func shouldEnqueueSchedulerOutboxForExtraUpdates(updates map[string]any) bool {
|
||||
if len(updates) == 0 {
|
||||
return false
|
||||
@@ -2177,6 +2507,16 @@ func isSchedulerNeutralExtraKey(key string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func upstreamBillingProbeExplicitlyDisabled(extra map[string]any) bool {
|
||||
enabled, ok := extra[service.UpstreamBillingProbeEnabledExtraKey].(bool)
|
||||
return ok && !enabled
|
||||
}
|
||||
|
||||
func upstreamBillingProbeSnapshotClearRequested(extra map[string]any) bool {
|
||||
value, ok := extra[service.UpstreamBillingProbeExtraKey]
|
||||
return ok && value == nil
|
||||
}
|
||||
|
||||
func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates service.AccountBulkUpdate) (int64, error) {
|
||||
if len(ids) == 0 {
|
||||
return 0, nil
|
||||
@@ -2250,7 +2590,11 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
setClauses = append(setClauses, "extra = COALESCE(extra, '{}'::jsonb) || $"+itoa(idx)+"::jsonb")
|
||||
extraExpression := "COALESCE(extra, '{}'::jsonb) || $" + itoa(idx) + "::jsonb"
|
||||
if upstreamBillingProbeExplicitlyDisabled(updates.Extra) || upstreamBillingProbeSnapshotClearRequested(updates.Extra) {
|
||||
extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'"
|
||||
}
|
||||
setClauses = append(setClauses, "extra = "+extraExpression)
|
||||
args = append(args, payload)
|
||||
idx++
|
||||
}
|
||||
@@ -2264,7 +2608,26 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
|
||||
query := "UPDATE accounts SET " + joinClauses(setClauses, ", ") + " WHERE id = ANY($" + itoa(idx) + ") AND deleted_at IS NULL"
|
||||
args = append(args, pq.Array(ids))
|
||||
|
||||
result, err := r.sql.ExecContext(ctx, query, args...)
|
||||
baseCtx := ctx
|
||||
contextTx := dbent.TxFromContext(ctx)
|
||||
exec := r.sql
|
||||
var tx *dbent.Tx
|
||||
if contextTx != nil {
|
||||
exec = contextTx.Client()
|
||||
} else if r.client != nil {
|
||||
var txErr error
|
||||
tx, txErr = r.client.Tx(ctx)
|
||||
if txErr != nil && !errors.Is(txErr, dbent.ErrTxStarted) {
|
||||
return 0, txErr
|
||||
}
|
||||
if tx != nil {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
ctx = dbent.NewTxContext(ctx, tx)
|
||||
exec = tx.Client()
|
||||
}
|
||||
}
|
||||
|
||||
result, err := exec.ExecContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
@@ -2274,9 +2637,16 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
|
||||
}
|
||||
if rows > 0 {
|
||||
payload := map[string]any{"account_ids": ids}
|
||||
if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil {
|
||||
logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue bulk update failed: err=%v", err)
|
||||
if err := enqueueSchedulerOutbox(ctx, exec, service.SchedulerOutboxEventAccountBulkChanged, nil, nil, payload); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
if tx != nil {
|
||||
if err := tx.Commit(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
if rows > 0 && contextTx == nil {
|
||||
shouldSync := false
|
||||
if updates.Status != nil && (*updates.Status == service.StatusError || *updates.Status == service.StatusDisabled) {
|
||||
shouldSync = true
|
||||
@@ -2285,7 +2655,7 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates
|
||||
shouldSync = true
|
||||
}
|
||||
if shouldSync {
|
||||
r.syncSchedulerAccountSnapshots(ctx, ids)
|
||||
r.syncSchedulerAccountSnapshots(baseCtx, ids)
|
||||
}
|
||||
}
|
||||
return rows, nil
|
||||
@@ -2733,6 +3103,107 @@ func (r *accountRepository) FindByExtraField(ctx context.Context, key string, va
|
||||
return r.accountsToService(ctx, accounts)
|
||||
}
|
||||
|
||||
// ListDueUpstreamBillingProbeAccounts bounds result hydration and network work
|
||||
// to limit. PostgreSQL must still filter and order all enabled candidates;
|
||||
// MATERIALIZED avoids repeating the defensive timestamp parse expression.
|
||||
func (r *accountRepository) ListDueUpstreamBillingProbeAccounts(ctx context.Context, now time.Time, limit int) ([]service.Account, error) {
|
||||
if limit <= 0 {
|
||||
return []service.Account{}, nil
|
||||
}
|
||||
if r.sql == nil {
|
||||
return nil, errors.New("account repository SQL executor not configured")
|
||||
}
|
||||
|
||||
rows, err := r.sql.QueryContext(ctx, `
|
||||
WITH candidates AS (
|
||||
SELECT
|
||||
id,
|
||||
extra #>> '{upstream_billing_probe,status}' AS probe_status,
|
||||
extra #>> '{upstream_billing_probe,next_probe_at}' AS next_probe_at
|
||||
FROM accounts
|
||||
WHERE deleted_at IS NULL
|
||||
AND status = 'active'
|
||||
AND platform = 'openai'
|
||||
AND type = 'apikey'
|
||||
AND extra @> '{"upstream_billing_probe_enabled": true}'::jsonb
|
||||
), parsed AS MATERIALIZED (
|
||||
SELECT
|
||||
id,
|
||||
probe_status,
|
||||
next_probe_at,
|
||||
next_probe_at ~ '^[0-9]{4}-[0-9]{2}-[0-9]{2}T[0-9]{2}:[0-9]{2}:[0-9]{2}(\.[0-9]+)?(Z|[+-][0-9]{2}:[0-9]{2})$' AS rfc3339_shape,
|
||||
jsonb_path_query_first_tz(
|
||||
jsonb_build_object(
|
||||
'value',
|
||||
replace(regexp_replace(next_probe_at, 'Z$', '+00:00'), 'T', ' ')
|
||||
),
|
||||
'$.value.datetime()',
|
||||
'{}'::jsonb,
|
||||
true
|
||||
) #>> '{}' AS parsed_next_probe_at
|
||||
FROM candidates
|
||||
), normalized AS (
|
||||
SELECT
|
||||
id,
|
||||
probe_status,
|
||||
next_probe_at,
|
||||
parsed_next_probe_at,
|
||||
rfc3339_shape AND parsed_next_probe_at IS NOT NULL AS valid_next_probe_at
|
||||
FROM parsed
|
||||
)
|
||||
SELECT id
|
||||
FROM normalized
|
||||
WHERE probe_status NOT IN ('ok', 'unsupported', 'failed')
|
||||
OR probe_status IS NULL
|
||||
OR next_probe_at IS NULL
|
||||
OR NOT valid_next_probe_at
|
||||
OR CASE WHEN valid_next_probe_at THEN parsed_next_probe_at::timestamptz <= $1 ELSE FALSE END
|
||||
ORDER BY
|
||||
CASE
|
||||
WHEN probe_status NOT IN ('ok', 'unsupported', 'failed')
|
||||
OR probe_status IS NULL
|
||||
OR next_probe_at IS NULL
|
||||
OR NOT valid_next_probe_at
|
||||
THEN 0
|
||||
ELSE 1
|
||||
END ASC,
|
||||
CASE WHEN valid_next_probe_at THEN parsed_next_probe_at::timestamptz END ASC NULLS FIRST,
|
||||
id ASC
|
||||
LIMIT $2
|
||||
`, now.UTC(), limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
ids := make([]int64, 0, limit)
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
if err := rows.Scan(&id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return []service.Account{}, nil
|
||||
}
|
||||
|
||||
accounts, err := r.GetByIDs(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]service.Account, 0, len(accounts))
|
||||
for _, account := range accounts {
|
||||
if account != nil {
|
||||
out = append(out, *account)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// nowUTC is a SQL expression to generate a UTC RFC3339 timestamp string.
|
||||
const nowUTC = `to_char(NOW() AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.US"Z"')`
|
||||
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
+51
@@ -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,
|
||||
)
|
||||
}
|
||||
@@ -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", "")
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 ©, nil
|
||||
}
|
||||
|
||||
func (s *updatingProxyRepoStub) Update(_ context.Context, proxy *Proxy) error {
|
||||
s.updateCalls++
|
||||
copy := *proxy
|
||||
s.proxy = ©
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestBothProxyUpdateServicesUseRepositoryUpdateBoundary(t *testing.T) {
|
||||
t.Run("ProxyService", func(t *testing.T) {
|
||||
repo := &updatingProxyRepoStub{
|
||||
proxyRepoStub: &proxyRepoStub{},
|
||||
proxy: &Proxy{ID: 9, Protocol: "http", Host: "old.example", Port: 8080, Status: StatusActive},
|
||||
}
|
||||
svc := NewProxyService(repo)
|
||||
host := "new.example"
|
||||
|
||||
_, err := svc.Update(context.Background(), 9, UpdateProxyRequest{Host: &host})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, repo.updateCalls)
|
||||
require.Equal(t, host, repo.proxy.Host)
|
||||
})
|
||||
|
||||
t.Run("adminService", func(t *testing.T) {
|
||||
repo := &updatingProxyRepoStub{
|
||||
proxyRepoStub: &proxyRepoStub{},
|
||||
proxy: &Proxy{
|
||||
ID: 9,
|
||||
Protocol: "http",
|
||||
Host: "old.example",
|
||||
Port: 8080,
|
||||
Status: StatusActive,
|
||||
FallbackMode: FallbackModeNone,
|
||||
ExpiryWarnDays: 7,
|
||||
},
|
||||
}
|
||||
svc := &adminServiceImpl{proxyRepo: repo}
|
||||
|
||||
_, err := svc.UpdateProxy(context.Background(), 9, &UpdateProxyInput{
|
||||
Host: "new.example",
|
||||
FallbackMode: FallbackModeNone,
|
||||
ExpiryWarnDays: 7,
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, repo.updateCalls)
|
||||
require.Equal(t, "new.example", repo.proxy.Host)
|
||||
})
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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] })
|
||||
})
|
||||
})
|
||||
@@ -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} 失败'
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user