diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index 22d7322afa..47d3b9ec33 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -111,6 +111,7 @@ func provideCleanup( channelMonitorRunner *service.ChannelMonitorRunner, quotaFlusher *service.UserPlatformQuotaUsageFlusher, upstreamBillingProbe *service.UpstreamBillingProbeService, + ollamaCloudUsage *service.OllamaCloudUsageService, auditLog *service.AuditLogService, promptAudit *securityaudit.PromptService, ) func() { @@ -331,6 +332,12 @@ func provideCleanup( } return nil }}, + {"OllamaCloudUsageService", func() error { + if ollamaCloudUsage != nil { + ollamaCloudUsage.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index f9d7b6fe9e..43c108db4d 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -269,7 +269,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { auditLogService := service.ProvideAuditLogService(auditLogRepository, settingService) auditLogHandler := admin.NewAuditLogHandler(auditLogService, totpService) 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, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService) + ollamaCloudUsageService := service.ProvideOllamaCloudUsageService(accountRepository, httpUpstream, settingService, secretEncryptor, configConfig, 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, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService, ollamaCloudUsageService) usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig) userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) @@ -317,7 +318,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, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, opsService, opsIngressRejectAggregator, apiKeyService, authCacheInvalidationWorker, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, ollamaCloudUsageService, auditLogService, promptService) application := &Application{ Server: httpServer, PromptAudit: promptService, @@ -384,6 +385,7 @@ func provideCleanup( channelMonitorRunner *service.ChannelMonitorRunner, quotaFlusher *service.UserPlatformQuotaUsageFlusher, upstreamBillingProbe *service.UpstreamBillingProbeService, + ollamaCloudUsage *service.OllamaCloudUsageService, auditLog *service.AuditLogService, promptAudit *securityaudit.PromptService, ) func() { @@ -603,6 +605,12 @@ func provideCleanup( } return nil }}, + {"OllamaCloudUsageService", func() error { + if ollamaCloudUsage != nil { + ollamaCloudUsage.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 3f159af7fe..592177be07 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -88,6 +88,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // channelMonitorRunner nil, // quotaFlusher nil, // upstreamBillingProbe + nil, // ollamaCloudUsage nil, // auditLog nil, // promptAudit ) diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 6f585fcfc7..7cdedad2ed 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -63,6 +63,7 @@ type AccountHandler struct { tokenCacheInvalidator service.TokenCacheInvalidator grokImportProber grokImportProber upstreamBillingProbe *service.UpstreamBillingProbeService + ollamaCloudUsage *service.OllamaCloudUsageService } // SetUpstreamBillingProbeService attaches the optional remote billing probe service. @@ -70,6 +71,10 @@ func (h *AccountHandler) SetUpstreamBillingProbeService(probe *service.UpstreamB h.upstreamBillingProbe = probe } +func (h *AccountHandler) SetOllamaCloudUsageService(usage *service.OllamaCloudUsageService) { + h.ollamaCloudUsage = usage +} + // NewAccountHandler creates a new admin account handler func NewAccountHandler( adminService service.AdminService, @@ -208,9 +213,17 @@ type AccountSchedulerGroupScore struct { const accountListGroupUngroupedQueryValue = "ungrouped" +func (h *AccountHandler) accountResponseFromService(account *service.Account) *dto.Account { + out := dto.AccountFromService(account) + if h != nil && h.ollamaCloudUsage != nil && out != nil { + h.ollamaCloudUsage.EnrichState(out.OllamaCloudUsage) + } + return out +} + func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, account *service.Account) AccountWithConcurrency { item := AccountWithConcurrency{ - Account: dto.AccountFromService(account), + Account: h.accountResponseFromService(account), CurrentConcurrency: 0, } if account == nil { @@ -523,6 +536,16 @@ func (h *AccountHandler) List(c *gin.Context) { response.ErrorFrom(c, err) return } + if h.ollamaCloudUsage != nil && len(accounts) > 0 { + accountPointers := make([]*service.Account, len(accounts)) + for index := range accounts { + accountPointers[index] = &accounts[index] + } + if err := h.ollamaCloudUsage.ResolveAccounts(c.Request.Context(), accountPointers); err != nil { + response.ErrorFrom(c, err) + return + } + } // Get current concurrency counts for all accounts accountIDs := make([]int64, len(accounts)) @@ -626,7 +649,7 @@ func (h *AccountHandler) List(c *gin.Context) { for i := range accounts { acc := &accounts[i] item := AccountWithConcurrency{ - Account: dto.AccountFromService(acc), + Account: h.accountResponseFromService(acc), CurrentConcurrency: concurrencyCounts[acc.ID], SchedulerScore: schedulerScores[acc.ID], SchedulerScores: schedulerGroupScores[acc.ID], @@ -740,6 +763,12 @@ func (h *AccountHandler) GetByID(c *gin.Context) { response.ErrorFrom(c, err) return } + if h.ollamaCloudUsage != nil { + if err := h.ollamaCloudUsage.ResolveAccounts(c.Request.Context(), []*service.Account{account}); err != nil { + response.ErrorFrom(c, err) + return + } + } response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } diff --git a/backend/internal/handler/admin/account_ollama_cloud_usage.go b/backend/internal/handler/admin/account_ollama_cloud_usage.go new file mode 100644 index 0000000000..1c0df30a33 --- /dev/null +++ b/backend/internal/handler/admin/account_ollama_cloud_usage.go @@ -0,0 +1,162 @@ +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 ollamaCloudUsageSessionRequest struct { + Session string `json:"session" binding:"required"` +} + +type ollamaCloudUsageAutoRefreshRequest struct { + Enabled *bool `json:"enabled" binding:"required"` +} + +func (h *AccountHandler) GetOllamaCloudUsageSettings(c *gin.Context) { + if h.ollamaCloudUsage == nil { + response.ErrorFrom(c, service.ErrOllamaCloudUsageUnavailable) + return + } + settings, err := h.ollamaCloudUsage.GetSettings(c.Request.Context()) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, settings) +} + +func (h *AccountHandler) UpdateOllamaCloudUsageSettings(c *gin.Context) { + if h.ollamaCloudUsage == nil { + response.ErrorFrom(c, service.ErrOllamaCloudUsageUnavailable) + return + } + var req service.OllamaCloudUsageSettings + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + if err := h.ollamaCloudUsage.UpdateSettings(c.Request.Context(), &req); err != nil { + response.ErrorFrom(c, err) + return + } + settings, err := h.ollamaCloudUsage.GetSettings(c.Request.Context()) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, settings) +} + +func (h *AccountHandler) GetOllamaCloudUsage(c *gin.Context) { + if !h.requireOllamaCloudUsage(c) { + return + } + accountID, ok := ollamaCloudUsageAccountID(c) + if !ok { + return + } + state, err := h.ollamaCloudUsage.GetState(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, state) +} + +func (h *AccountHandler) SaveOllamaCloudUsageSession(c *gin.Context) { + if !h.requireOllamaCloudUsage(c) { + return + } + accountID, ok := ollamaCloudUsageAccountID(c) + if !ok { + return + } + var req ollamaCloudUsageSessionRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + state, err := h.ollamaCloudUsage.SaveSession(c.Request.Context(), accountID, req.Session) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, state) +} + +func (h *AccountHandler) DeleteOllamaCloudUsageSession(c *gin.Context) { + if !h.requireOllamaCloudUsage(c) { + return + } + accountID, ok := ollamaCloudUsageAccountID(c) + if !ok { + return + } + state, err := h.ollamaCloudUsage.DeleteSession(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, state) +} + +func (h *AccountHandler) SetOllamaCloudUsageAutoRefresh(c *gin.Context) { + if !h.requireOllamaCloudUsage(c) { + return + } + accountID, ok := ollamaCloudUsageAccountID(c) + if !ok { + return + } + var req ollamaCloudUsageAutoRefreshRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + state, err := h.ollamaCloudUsage.SetAutoRefresh(c.Request.Context(), accountID, *req.Enabled) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, state) +} + +func (h *AccountHandler) RefreshOllamaCloudUsage(c *gin.Context) { + if !h.requireOllamaCloudUsage(c) { + return + } + accountID, ok := ollamaCloudUsageAccountID(c) + if !ok { + return + } + state, err := h.ollamaCloudUsage.Refresh(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, state) +} + +func (h *AccountHandler) requireOllamaCloudUsage(c *gin.Context) bool { + if h != nil && h.ollamaCloudUsage != nil { + return true + } + response.ErrorFrom(c, service.ErrOllamaCloudUsageUnavailable) + return false +} + +func ollamaCloudUsageAccountID(c *gin.Context) (int64, bool) { + if c == nil { + return 0, false + } + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil || accountID <= 0 { + response.BadRequest(c, "Invalid account ID") + return 0, false + } + return accountID, true +} diff --git a/backend/internal/handler/admin/account_ollama_cloud_usage_test.go b/backend/internal/handler/admin/account_ollama_cloud_usage_test.go new file mode 100644 index 0000000000..89f2aa44ab --- /dev/null +++ b/backend/internal/handler/admin/account_ollama_cloud_usage_test.go @@ -0,0 +1,276 @@ +package admin + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type ollamaCloudUsageHandlerTestRepo struct { + service.AccountRepository + account *service.Account + accounts []*service.Account + groupResolveCalls int +} + +func (r *ollamaCloudUsageHandlerTestRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + if r.account != nil && r.account.ID == id { + return r.account, nil + } + for _, account := range r.accounts { + if account.ID == id { + return account, nil + } + } + return nil, service.ErrAccountNotFound +} + +func (r *ollamaCloudUsageHandlerTestRepo) ListOllamaCloudUsageGroupAccounts(_ context.Context, _ []*service.Account) ([]service.Account, error) { + r.groupResolveCalls++ + result := make([]service.Account, 0, len(r.accounts)+1) + if r.account != nil { + result = append(result, *r.account) + } + for _, account := range r.accounts { + result = append(result, *account) + } + return result, nil +} + +func (r *ollamaCloudUsageHandlerTestRepo) SaveOllamaCloudUsageSession(context.Context, *service.Account, string, bool) error { + return nil +} +func (r *ollamaCloudUsageHandlerTestRepo) DeleteOllamaCloudUsageSession(context.Context, *service.Account) error { + return nil +} +func (r *ollamaCloudUsageHandlerTestRepo) SetOllamaCloudUsageAutoRefresh(context.Context, *service.Account, bool) error { + return nil +} +func (r *ollamaCloudUsageHandlerTestRepo) UpdateOllamaCloudUsageSnapshot(context.Context, *service.Account, *service.OllamaCloudUsageSnapshot) error { + return nil +} +func (r *ollamaCloudUsageHandlerTestRepo) DisableOllamaCloudUsageAutoRefresh(context.Context, *service.Account) error { + return nil +} +func (r *ollamaCloudUsageHandlerTestRepo) ListDueOllamaCloudUsageAccounts(context.Context, time.Time, int) ([]service.Account, error) { + return nil, nil +} + +func newOllamaCloudUsageHandlerTestService(t *testing.T) *service.OllamaCloudUsageService { + t.Helper() + svc := service.NewOllamaCloudUsageService(nil, nil, nil, nil, false) + t.Cleanup(svc.Stop) + return svc +} + +func newOllamaCloudUsageHandlerContext(method, target, body, id string) (*gin.Context, *httptest.ResponseRecorder) { + recorder := httptest.NewRecorder() + request := httptest.NewRequest(method, target, bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = request + if id != "" { + ctx.Params = gin.Params{{Key: "id", Value: id}} + } + return ctx, recorder +} + +func TestOllamaCloudUsageHandlersValidateRequestsAndDependencies(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := newOllamaCloudUsageHandlerTestService(t) + + t.Run("invalid account id", func(t *testing.T) { + ctx, recorder := newOllamaCloudUsageHandlerContext(http.MethodGet, "/admin/accounts/not-an-id/ollama-cloud-usage", "", "not-an-id") + (&AccountHandler{ollamaCloudUsage: svc}).GetOllamaCloudUsage(ctx) + require.Equal(t, http.StatusBadRequest, recorder.Code) + }) + + t.Run("empty session", func(t *testing.T) { + ctx, recorder := newOllamaCloudUsageHandlerContext(http.MethodPut, "/admin/accounts/7/ollama-cloud-usage/session", `{"session":""}`, "7") + (&AccountHandler{ollamaCloudUsage: svc}).SaveOllamaCloudUsageSession(ctx) + require.Equal(t, http.StatusBadRequest, recorder.Code) + }) + + t.Run("missing enabled", func(t *testing.T) { + ctx, recorder := newOllamaCloudUsageHandlerContext(http.MethodPut, "/admin/accounts/7/ollama-cloud-usage/auto-refresh", `{}`, "7") + (&AccountHandler{ollamaCloudUsage: svc}).SetOllamaCloudUsageAutoRefresh(ctx) + require.Equal(t, http.StatusBadRequest, recorder.Code) + }) + + t.Run("service unavailable", func(t *testing.T) { + ctx, recorder := newOllamaCloudUsageHandlerContext(http.MethodGet, "/admin/accounts/7/ollama-cloud-usage", "", "7") + (&AccountHandler{}).GetOllamaCloudUsage(ctx) + require.Equal(t, http.StatusServiceUnavailable, recorder.Code) + require.Contains(t, recorder.Body.String(), "OLLAMA_CLOUD_USAGE_UNAVAILABLE") + }) +} + +func TestOllamaCloudUsageEncryptionKeyStateConsistentAcrossAccountResponses(t *testing.T) { + gin.SetMode(gin.TestMode) + + for _, configured := range []bool{false, true} { + t.Run("configured="+strconv.FormatBool(configured), func(t *testing.T) { + account := &service.Account{ + ID: 7, + Name: "ollama", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ollama.com", "api_key": "test-key"}, + Extra: map[string]any{}, + Status: service.StatusActive, + } + adminService := newStubAdminService() + adminService.accounts = []service.Account{*account} + adminService.getAccountResult = account + usageService := service.NewOllamaCloudUsageService( + &ollamaCloudUsageHandlerTestRepo{account: account}, nil, nil, nil, configured, + ) + t.Cleanup(usageService.Stop) + + handler := NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + handler.SetOllamaCloudUsageService(usageService) + router := gin.New() + router.GET("/accounts", handler.List) + router.GET("/accounts/:id", handler.GetByID) + router.GET("/accounts/:id/ollama-cloud-usage", handler.GetOllamaCloudUsage) + + listRecorder := httptest.NewRecorder() + router.ServeHTTP(listRecorder, httptest.NewRequest(http.MethodGet, "/accounts?page=1&page_size=20", nil)) + require.Equal(t, http.StatusOK, listRecorder.Code) + var listPayload struct { + Data struct { + Items []struct { + OllamaCloudUsage *service.OllamaCloudUsageState `json:"ollama_cloud_usage"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(listRecorder.Body.Bytes(), &listPayload)) + require.Len(t, listPayload.Data.Items, 1) + require.NotNil(t, listPayload.Data.Items[0].OllamaCloudUsage) + + detailRecorder := httptest.NewRecorder() + router.ServeHTTP(detailRecorder, httptest.NewRequest(http.MethodGet, "/accounts/7", nil)) + require.Equal(t, http.StatusOK, detailRecorder.Code) + var detailPayload struct { + Data struct { + OllamaCloudUsage *service.OllamaCloudUsageState `json:"ollama_cloud_usage"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(detailRecorder.Body.Bytes(), &detailPayload)) + require.NotNil(t, detailPayload.Data.OllamaCloudUsage) + + stateRecorder := httptest.NewRecorder() + router.ServeHTTP(stateRecorder, httptest.NewRequest(http.MethodGet, "/accounts/7/ollama-cloud-usage", nil)) + require.Equal(t, http.StatusOK, stateRecorder.Code) + var statePayload struct { + Data service.OllamaCloudUsageState `json:"data"` + } + require.NoError(t, json.Unmarshal(stateRecorder.Body.Bytes(), &statePayload)) + + listConfigured := listPayload.Data.Items[0].OllamaCloudUsage.EncryptionKeyConfigured + detailConfigured := detailPayload.Data.OllamaCloudUsage.EncryptionKeyConfigured + require.Equal(t, configured, listConfigured) + require.Equal(t, statePayload.Data.EncryptionKeyConfigured, listConfigured) + require.Equal(t, statePayload.Data.EncryptionKeyConfigured, detailConfigured) + }) + } +} + +func TestOllamaCloudUsageSharedStateMatchesListDetailAndSpecialEndpointWithoutListNPlusOne(t *testing.T) { + gin.SetMode(gin.TestMode) + now := time.Now().UTC() + source := &service.Account{ + ID: 7, Name: "source", Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ollama.com", "api_key": "shared-secret-key"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "ciphertext-secret", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + service.OllamaCloudUsageSnapshotExtraKey: &service.OllamaCloudUsageSnapshot{ + Status: service.OllamaCloudUsageStatusOK, Data: &service.OllamaCloudUsageData{Plan: "pro"}, + LastAttemptAt: now, NextRefreshAt: now.Add(time.Hour), + }, + }, + Status: service.StatusActive, + } + sibling := &service.Account{ + ID: 8, Name: "sibling", Platform: service.PlatformAnthropic, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "HTTPS://WWW.OLLAMA.COM:443/v1", "api_key": "shared-secret-key"}, + Extra: map[string]any{}, Status: service.StatusActive, + } + repo := &ollamaCloudUsageHandlerTestRepo{accounts: []*service.Account{source, sibling}} + adminService := newStubAdminService() + adminService.accounts = []service.Account{*source, *sibling} + adminService.getAccountResult = sibling + usageService := service.NewOllamaCloudUsageService(repo, nil, nil, nil, true) + t.Cleanup(usageService.Stop) + handler := NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + handler.SetOllamaCloudUsageService(usageService) + router := gin.New() + router.GET("/accounts", handler.List) + router.GET("/accounts/:id", handler.GetByID) + router.GET("/accounts/:id/ollama-cloud-usage", handler.GetOllamaCloudUsage) + + listRecorder := httptest.NewRecorder() + router.ServeHTTP(listRecorder, httptest.NewRequest(http.MethodGet, "/accounts?page=1&page_size=20", nil)) + require.Equal(t, http.StatusOK, listRecorder.Code) + require.Equal(t, 1, repo.groupResolveCalls, "the full list page must use one group-resolution batch") + var listPayload struct { + Data struct { + Items []struct { + ID int64 `json:"id"` + OllamaCloudUsage *service.OllamaCloudUsageState `json:"ollama_cloud_usage"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(listRecorder.Body.Bytes(), &listPayload)) + require.Len(t, listPayload.Data.Items, 2) + for _, item := range listPayload.Data.Items { + require.True(t, item.OllamaCloudUsage.Configured) + require.Equal(t, "pro", item.OllamaCloudUsage.Snapshot.Data.Plan) + } + + detailRecorder := httptest.NewRecorder() + router.ServeHTTP(detailRecorder, httptest.NewRequest(http.MethodGet, "/accounts/8", nil)) + require.Equal(t, http.StatusOK, detailRecorder.Code) + var detailPayload struct { + Data struct { + OllamaCloudUsage *service.OllamaCloudUsageState `json:"ollama_cloud_usage"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(detailRecorder.Body.Bytes(), &detailPayload)) + + stateRecorder := httptest.NewRecorder() + router.ServeHTTP(stateRecorder, httptest.NewRequest(http.MethodGet, "/accounts/8/ollama-cloud-usage", nil)) + require.Equal(t, http.StatusOK, stateRecorder.Code) + var statePayload struct { + Data service.OllamaCloudUsageState `json:"data"` + } + require.NoError(t, json.Unmarshal(stateRecorder.Body.Bytes(), &statePayload)) + require.Equal(t, statePayload.Data.Configured, detailPayload.Data.OllamaCloudUsage.Configured) + require.Equal(t, statePayload.Data.Snapshot, detailPayload.Data.OllamaCloudUsage.Snapshot) + for _, body := range []string{listRecorder.Body.String(), detailRecorder.Body.String(), stateRecorder.Body.String()} { + require.NotContains(t, body, "shared-secret-key") + require.NotContains(t, body, "ciphertext-secret") + } +} + +func TestGetOllamaCloudUsageSettingsHandlerSuccess(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, recorder := newOllamaCloudUsageHandlerContext(http.MethodGet, "/admin/accounts/ollama-cloud-usage/settings", "", "") + handler := &AccountHandler{ollamaCloudUsage: newOllamaCloudUsageHandlerTestService(t)} + + handler.GetOllamaCloudUsageSettings(ctx) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Contains(t, recorder.Body.String(), `"enabled":false`) + require.Contains(t, recorder.Body.String(), `"interval_minutes":60`) +} diff --git a/backend/internal/handler/dto/account_mapper_redact_test.go b/backend/internal/handler/dto/account_mapper_redact_test.go index bd584e1123..9a979a66b6 100644 --- a/backend/internal/handler/dto/account_mapper_redact_test.go +++ b/backend/internal/handler/dto/account_mapper_redact_test.go @@ -58,6 +58,41 @@ func TestAccountFromServiceShallow_RedactsSensitiveCredentials(t *testing.T) { require.Equal(t, "rt-secret", src.Credentials["refresh_token"]) } +func TestAccountFromServiceShallow_RedactsOllamaCloudManagedExtra(t *testing.T) { + snapshot := map[string]any{ + "status": service.OllamaCloudUsageStatusOK, + "last_attempt_at": "2026-07-22T12:00:00Z", + "next_refresh_at": "2026-07-22T13:00:00Z", + "data": map[string]any{"plan": "Pro"}, + } + src := &service.Account{ + ID: 9, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ollama.com", "api_key": "secret-key"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "ciphertext-secret", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + service.OllamaCloudUsageSnapshotExtraKey: snapshot, + "ordinary": "kept", + }, + } + + got := AccountFromServiceShallow(src) + require.NotContains(t, got.Extra, service.OllamaCloudUsageSessionExtraKey) + require.NotContains(t, got.Extra, service.OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, got.Extra, service.OllamaCloudUsageSnapshotExtraKey) + require.Equal(t, "kept", got.Extra["ordinary"]) + require.NotNil(t, got.OllamaCloudUsage) + require.True(t, got.OllamaCloudUsage.Configured) + require.True(t, got.OllamaCloudUsage.AutoRefreshEnabled) + require.Equal(t, "Pro", got.OllamaCloudUsage.Snapshot.Data.Plan) + + raw, err := json.Marshal(got) + require.NoError(t, err) + require.NotContains(t, string(raw), "ciphertext-secret") + require.NotContains(t, string(raw), "secret-key") + require.Contains(t, src.Extra, service.OllamaCloudUsageSessionExtraKey) +} + func TestAccountFromServiceShallow_NilCredentialsOmitsStatus(t *testing.T) { src := &service.Account{ID: 1, Name: "n", Platform: "anthropic", Type: "oauth"} got := AccountFromServiceShallow(src) diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 6f82e29c06..1ae9a5e790 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -219,6 +219,11 @@ func AccountFromServiceShallow(a *service.Account) *Account { return nil } redactedCreds, credsStatus := RedactCredentials(a.Credentials) + extra := redactAccountManagedExtra(a.Extra) + var ollamaCloudUsage *service.OllamaCloudUsageState + if state := service.OllamaCloudUsageStateFromAccount(a); state.Eligible { + ollamaCloudUsage = state + } out := &Account{ ID: a.ID, Name: a.Name, @@ -227,7 +232,8 @@ func AccountFromServiceShallow(a *service.Account) *Account { Type: a.Type, Credentials: redactedCreds, CredentialsStatus: credsStatus, - Extra: a.Extra, + Extra: extra, + OllamaCloudUsage: ollamaCloudUsage, ProxyID: a.ProxyID, ProxyFallbackOriginID: a.ProxyFallbackOriginID, ProxyFallbackOriginName: a.ProxyFallbackOriginName, @@ -385,6 +391,24 @@ func AccountFromServiceShallow(a *service.Account) *Account { return out } +func redactAccountManagedExtra(extra map[string]any) map[string]any { + if extra == nil { + return nil + } + redacted := make(map[string]any, len(extra)) + for key, value := range extra { + switch key { + case service.OllamaCloudUsageSessionExtraKey, + service.OllamaCloudUsageAutoRefreshExtraKey, + service.OllamaCloudUsageSnapshotExtraKey: + continue + default: + redacted[key] = value + } + } + return redacted +} + func AccountFromService(a *service.Account) *Account { if a == nil { return nil diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 117091f410..38133cff3b 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -6,6 +6,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/domain" + "github.com/Wei-Shaw/sub2api/internal/service" ) type User struct { @@ -183,23 +184,24 @@ type Account struct { Type string `json:"type"` // Credentials 经 RedactCredentials 处理后只含非敏感子键;敏感 token / api_key / 私钥 // 的存在性通过 CredentialsStatus(has_)暴露,原始值不返回前端。 - Credentials map[string]any `json:"credentials"` - CredentialsStatus map[string]bool `json:"credentials_status,omitempty"` - Extra map[string]any `json:"extra"` - ProxyID *int64 `json:"proxy_id"` - ProxyFallbackOriginID *int64 `json:"proxy_fallback_origin_id"` - ProxyFallbackOriginName *string `json:"proxy_fallback_origin_name,omitempty"` - Concurrency int `json:"concurrency"` - LoadFactor *int `json:"load_factor,omitempty"` - Priority int `json:"priority"` - RateMultiplier float64 `json:"rate_multiplier"` - Status string `json:"status"` - ErrorMessage string `json:"error_message"` - LastUsedAt *time.Time `json:"last_used_at"` - ExpiresAt *int64 `json:"expires_at"` - AutoPauseOnExpired bool `json:"auto_pause_on_expired"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + Credentials map[string]any `json:"credentials"` + CredentialsStatus map[string]bool `json:"credentials_status,omitempty"` + Extra map[string]any `json:"extra"` + OllamaCloudUsage *service.OllamaCloudUsageState `json:"ollama_cloud_usage,omitempty"` + ProxyID *int64 `json:"proxy_id"` + ProxyFallbackOriginID *int64 `json:"proxy_fallback_origin_id"` + ProxyFallbackOriginName *string `json:"proxy_fallback_origin_name,omitempty"` + Concurrency int `json:"concurrency"` + LoadFactor *int `json:"load_factor,omitempty"` + Priority int `json:"priority"` + RateMultiplier float64 `json:"rate_multiplier"` + Status string `json:"status"` + ErrorMessage string `json:"error_message"` + LastUsedAt *time.Time `json:"last_used_at"` + ExpiresAt *int64 `json:"expires_at"` + AutoPauseOnExpired bool `json:"auto_pause_on_expired"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` Schedulable bool `json:"schedulable"` diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 54d82a69e9..073c19d375 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -46,8 +46,10 @@ func ProvideAdminHandlers( complianceHandler *admin.ComplianceHandler, auditLogHandler *admin.AuditLogHandler, upstreamBillingProbe *service.UpstreamBillingProbeService, + ollamaCloudUsage *service.OllamaCloudUsageService, ) *AdminHandlers { accountHandler.SetUpstreamBillingProbeService(upstreamBillingProbe) + accountHandler.SetOllamaCloudUsageService(ollamaCloudUsage) return &AdminHandlers{ Dashboard: dashboardHandler, User: userHandler, diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 87f91777b0..a668f578cb 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -58,6 +58,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{ "codex_7d_", "passive_usage_", "upstream_billing_probe", + "ollama_cloud_usage", } var schedulerNeutralExtraKeys = map[string]struct{}{ @@ -556,8 +557,22 @@ func lockAndMergeAccountProbeExtra(ctx context.Context, client *dbent.Client, ac AND type = $3 AND credentials = $4::jsonb AND proxy_id IS NOT DISTINCT FROM $5, + COALESCE( + platform IN ('openai', 'anthropic') + AND $2 IN ('openai', 'anthropic') + AND type = 'apikey' + AND $3 = 'apikey' + AND credentials -> 'api_key' IS NOT DISTINCT FROM $4::jsonb -> 'api_key' + AND `+ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'")+` + AND `+ollamaCloudBaseURLMatchesSQL("$4::jsonb ->> 'base_url'")+`, + false + ), + proxy_id IS NOT DISTINCT FROM $5, extra -> 'upstream_billing_probe_enabled', - extra -> 'upstream_billing_probe' + extra -> 'upstream_billing_probe', + extra -> 'ollama_cloud_usage_session', + extra -> 'ollama_cloud_usage_auto_refresh', + extra -> 'ollama_cloud_usage_snapshot' FROM accounts WHERE id = $1 AND deleted_at IS NULL FOR NO KEY UPDATE @@ -574,11 +589,25 @@ func lockAndMergeAccountProbeExtra(ctx context.Context, client *dbent.Client, ac } var ( - identityUnchanged bool - currentEnabled []byte - currentSnapshot []byte + identityUnchanged bool + ollamaGroupIdentityUnchanged bool + ollamaProxyIdentityUnchanged bool + currentEnabled []byte + currentSnapshot []byte + currentOllamaSession []byte + currentOllamaAutoRefresh []byte + currentOllamaSnapshot []byte ) - if err := rows.Scan(&identityUnchanged, ¤tEnabled, ¤tSnapshot); err != nil { + if err := rows.Scan( + &identityUnchanged, + &ollamaGroupIdentityUnchanged, + &ollamaProxyIdentityUnchanged, + ¤tEnabled, + ¤tSnapshot, + ¤tOllamaSession, + ¤tOllamaAutoRefresh, + ¤tOllamaSnapshot, + ); err != nil { return nil, err } if err := rows.Err(); err != nil { @@ -586,34 +615,71 @@ func lockAndMergeAccountProbeExtra(ctx context.Context, client *dbent.Client, ac } extra := copyJSONMap(normalizeJSONMap(account.Extra)) - delete(extra, service.UpstreamBillingProbeEnabledExtraKey) - delete(extra, service.UpstreamBillingProbeExtraKey) + for _, key := range []string{ + service.UpstreamBillingProbeEnabledExtraKey, + service.UpstreamBillingProbeExtraKey, + service.OllamaCloudUsageSessionExtraKey, + service.OllamaCloudUsageAutoRefreshExtraKey, + service.OllamaCloudUsageSnapshotExtraKey, + } { + delete(extra, key) + } 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 { + } else if probeAccount { + if enabled, ok, err := decodeAccountExtraJSON(currentEnabled); err != nil { return nil, err - } - extra[service.UpstreamBillingProbeEnabledExtraKey] = enabled - if value, ok := enabled.(bool); ok && !value { - probeExplicitlyDisabled = true + } else if ok { + extra[service.UpstreamBillingProbeEnabledExtraKey] = enabled + if value, isBool := enabled.(bool); isBool && !value { + probeExplicitlyDisabled = true + } } } - if !identityUnchanged || probeExplicitlyDisabled || len(currentSnapshot) == 0 || string(currentSnapshot) == "null" { - return extra, nil + if identityUnchanged && !probeExplicitlyDisabled { + if snapshot, ok, err := decodeAccountExtraJSON(currentSnapshot); err != nil { + return nil, err + } else if ok { + extra[service.UpstreamBillingProbeExtraKey] = snapshot + } } - var snapshot any - if err := json.Unmarshal(currentSnapshot, &snapshot); err != nil { - return nil, err + + if service.IsOllamaCloudUsageAccount(account) && ollamaGroupIdentityUnchanged { + for key, raw := range map[string][]byte{ + service.OllamaCloudUsageSessionExtraKey: currentOllamaSession, + service.OllamaCloudUsageAutoRefreshExtraKey: currentOllamaAutoRefresh, + } { + if value, ok, err := decodeAccountExtraJSON(raw); err != nil { + return nil, err + } else if ok { + extra[key] = value + } + } + if ollamaProxyIdentityUnchanged { + if snapshot, ok, err := decodeAccountExtraJSON(currentOllamaSnapshot); err != nil { + return nil, err + } else if ok { + extra[service.OllamaCloudUsageSnapshotExtraKey] = snapshot + } + } } - extra[service.UpstreamBillingProbeExtraKey] = snapshot return extra, nil } +func decodeAccountExtraJSON(raw []byte) (any, bool, error) { + if len(raw) == 0 || string(raw) == "null" { + return nil, false, nil + } + var value any + if err := json.Unmarshal(raw, &value); err != nil { + return nil, false, err + } + return value, true, nil +} + func (r *accountRepository) UpdateCredentials(ctx context.Context, id int64, credentials map[string]any) error { payload, err := json.Marshal(normalizeJSONMap(credentials)) if err != nil { @@ -642,6 +708,22 @@ func (r *accountRepository) UpdateCredentials(ctx context.Context, id int64, cre SET credentials = $1::jsonb, extra = CASE + WHEN platform IN ('openai', 'anthropic') + AND type = 'apikey' + AND ( + credentials -> 'api_key' IS DISTINCT FROM $1::jsonb -> 'api_key' + OR NOT ( + `+ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'")+` + AND `+ollamaCloudBaseURLMatchesSQL("$1::jsonb ->> 'base_url'")+` + ) + ) + THEN (CASE + WHEN platform = 'openai' THEN COALESCE(extra, '{}'::jsonb) - 'upstream_billing_probe' + ELSE COALESCE(extra, '{}'::jsonb) + END) + - 'ollama_cloud_usage_session' + - 'ollama_cloud_usage_auto_refresh' + - 'ollama_cloud_usage_snapshot' WHEN platform = 'openai' AND type = 'apikey' AND credentials IS DISTINCT FROM $1::jsonb @@ -2605,6 +2687,11 @@ func upstreamBillingProbeSnapshotClearRequested(extra map[string]any) bool { return ok && value == nil } +func ollamaCloudUsageSnapshotClearRequested(extra map[string]any) bool { + value, ok := extra[service.OllamaCloudUsageSnapshotExtraKey] + 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 @@ -2614,6 +2701,7 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates args := make([]any, 0, 8) idx := 1 + ollamaProxyIdentityChanged := "" if updates.Name != nil { setClauses = append(setClauses, "name = $"+itoa(idx)) args = append(args, *updates.Name) @@ -2623,8 +2711,11 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates // 0 表示清除代理(前端发送 0 而不是 null 来表达清除意图) if *updates.ProxyID == 0 { setClauses = append(setClauses, "proxy_id = NULL") + ollamaProxyIdentityChanged = "proxy_id IS NOT NULL" } else { - setClauses = append(setClauses, "proxy_id = $"+itoa(idx)) + proxyPlaceholder := "$" + itoa(idx) + setClauses = append(setClauses, "proxy_id = "+proxyPlaceholder) + ollamaProxyIdentityChanged = "proxy_id IS DISTINCT FROM " + proxyPlaceholder args = append(args, *updates.ProxyID) idx++ } @@ -2670,27 +2761,68 @@ func (r *accountRepository) BulkUpdate(ctx context.Context, ids []int64, updates updates.Extra[service.UpstreamBillingProbeEnabledExtraKey] = *updates.ProbeEnabled } // JSONB 需要合并而非覆盖,使用 raw SQL 保持旧行为。 + credentialPlaceholder := "" if len(updates.Credentials) > 0 { payload, err := json.Marshal(updates.Credentials) if err != nil { return 0, err } - setClauses = append(setClauses, "credentials = COALESCE(credentials, '{}'::jsonb) || $"+itoa(idx)+"::jsonb") + credentialPlaceholder = "$" + itoa(idx) + setClauses = append(setClauses, "credentials = COALESCE(credentials, '{}'::jsonb) || "+credentialPlaceholder+"::jsonb") args = append(args, payload) idx++ } - if len(updates.Extra) > 0 { - payload, err := json.Marshal(updates.Extra) - if err != nil { - return 0, err + + ollamaGroupIdentityChanges := make([]string, 0, 2) + if _, ok := updates.Credentials["api_key"]; ok { + ollamaGroupIdentityChanges = append(ollamaGroupIdentityChanges, "credentials -> 'api_key' IS DISTINCT FROM "+credentialPlaceholder+"::jsonb -> 'api_key'") + } + if _, ok := updates.Credentials["base_url"]; ok { + ollamaGroupIdentityChanges = append(ollamaGroupIdentityChanges, + "NOT ("+ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'")+ + " AND "+ollamaCloudBaseURLMatchesSQL(credentialPlaceholder+"::jsonb ->> 'base_url'")+")") + } + + if len(updates.Extra) > 0 || len(ollamaGroupIdentityChanges) > 0 || ollamaProxyIdentityChanged != "" { + extraExpression := "COALESCE(extra, '{}'::jsonb)" + if len(updates.Extra) > 0 { + payload, err := json.Marshal(updates.Extra) + if err != nil { + return 0, err + } + extraExpression += " || $" + itoa(idx) + "::jsonb" + args = append(args, payload) + idx++ + if upstreamBillingProbeExplicitlyDisabled(updates.Extra) || upstreamBillingProbeSnapshotClearRequested(updates.Extra) { + extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'" + } + if ollamaCloudUsageSnapshotClearRequested(updates.Extra) { + extraExpression = "(" + extraExpression + ") - 'ollama_cloud_usage_snapshot'" + } } - extraExpression := "COALESCE(extra, '{}'::jsonb) || $" + itoa(idx) + "::jsonb" - if upstreamBillingProbeExplicitlyDisabled(updates.Extra) || upstreamBillingProbeSnapshotClearRequested(updates.Extra) { - extraExpression = "(" + extraExpression + ") - 'upstream_billing_probe'" + eligibleAccount := "platform IN ('openai', 'anthropic') AND type = 'apikey'" + groupIdentityChanged := "" + if len(ollamaGroupIdentityChanges) > 0 { + groupIdentityChanged = "(" + eligibleAccount + " AND (" + joinClauses(ollamaGroupIdentityChanges, " OR ") + "))" + } + snapshotIdentityChanged := groupIdentityChanged + if ollamaProxyIdentityChanged != "" { + proxyChanged := "(" + eligibleAccount + " AND " + ollamaProxyIdentityChanged + ")" + if snapshotIdentityChanged == "" { + snapshotIdentityChanged = proxyChanged + } else { + snapshotIdentityChanged = "(" + snapshotIdentityChanged + " OR " + proxyChanged + ")" + } + } + if groupIdentityChanged != "" { + extraExpression = "CASE" + + " WHEN " + groupIdentityChanged + " THEN (" + extraExpression + ") - 'ollama_cloud_usage_session' - 'ollama_cloud_usage_auto_refresh' - 'ollama_cloud_usage_snapshot'" + + " WHEN " + snapshotIdentityChanged + " THEN (" + extraExpression + ") - 'ollama_cloud_usage_snapshot'" + + " ELSE " + extraExpression + " END" + } else if snapshotIdentityChanged != "" { + extraExpression = "CASE WHEN " + snapshotIdentityChanged + " THEN (" + extraExpression + ") - 'ollama_cloud_usage_snapshot' ELSE " + extraExpression + " END" } setClauses = append(setClauses, "extra = "+extraExpression) - args = append(args, payload) - idx++ } if len(setClauses) == 0 { diff --git a/backend/internal/repository/account_repo_ollama_cloud_usage.go b/backend/internal/repository/account_repo_ollama_cloud_usage.go new file mode 100644 index 0000000000..78d14d2097 --- /dev/null +++ b/backend/internal/repository/account_repo_ollama_cloud_usage.go @@ -0,0 +1,440 @@ +package repository + +import ( + "context" + "encoding/json" + "errors" + "time" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/lib/pq" +) + +const ( + ollamaCloudBaseURLRegexSQL = `^[hH][tT][tT][pP][sS]://([wW][wW][wW]\.)?[oO][lL][lL][aA][mM][aA]\.[cC][oO][mM](:443)?(/v1)?$` + ollamaCloudBaseURLMatchSQLPrefix = "btrim(" + ollamaCloudBaseURLMatchSQLSuffix = ") ~ '" + ollamaCloudBaseURLRegexSQL + "'" + ollamaCloudUsageEligibleSQL = ` + platform IN ('openai', 'anthropic') + AND type = 'apikey' + AND ` + ollamaCloudBaseURLMatchSQLPrefix + `credentials ->> 'base_url'` + ollamaCloudBaseURLMatchSQLSuffix + ` + AND jsonb_typeof(credentials -> 'api_key') = 'string' +` +) + +func ollamaCloudBaseURLMatchesSQL(expression string) string { + return ollamaCloudBaseURLMatchSQLPrefix + expression + ollamaCloudBaseURLMatchSQLSuffix +} + +// ListOllamaCloudUsageGroupAccounts resolves every sibling for all supplied +// identities with one ID query and one batch hydration. API keys are query +// parameters only; no derived shared key is persisted. +func (r *accountRepository) ListOllamaCloudUsageGroupAccounts(ctx context.Context, accounts []*service.Account) ([]service.Account, error) { + if r == nil || r.sql == nil { + return nil, service.ErrOllamaCloudUsageUnavailable + } + keys := make([]string, 0, len(accounts)) + seen := make(map[string]struct{}, len(accounts)) + for _, account := range accounts { + if !service.IsOllamaCloudUsageAccount(account) || account.Credentials == nil { + continue + } + apiKey, ok := account.Credentials["api_key"].(string) + if !ok || apiKey == "" { + continue + } + if _, duplicate := seen[apiKey]; duplicate { + continue + } + seen[apiKey] = struct{}{} + keys = append(keys, apiKey) + } + if len(keys) == 0 { + return []service.Account{}, nil + } + rows, err := r.sql.QueryContext(ctx, ` + SELECT id + FROM accounts + WHERE deleted_at IS NULL + AND `+ollamaCloudUsageEligibleSQL+` + AND credentials ->> 'api_key' = ANY($1) + ORDER BY id + `, pq.Array(keys)) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + ids := make([]int64, 0, len(keys)) + 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 + } + hydrated, err := r.GetByIDs(ctx, ids) + if err != nil { + return nil, err + } + result := make([]service.Account, 0, len(hydrated)) + for _, account := range hydrated { + if account != nil { + result = append(result, *account) + } + } + return result, nil +} + +func (r *accountRepository) SaveOllamaCloudUsageSession(ctx context.Context, account *service.Account, ciphertext string, autoRefresh bool) error { + return r.updateOllamaCloudUsageGroup(ctx, account, map[string]any{ + service.OllamaCloudUsageSessionExtraKey: ciphertext, + service.OllamaCloudUsageAutoRefreshExtraKey: autoRefresh, + }, false) +} + +func (r *accountRepository) DeleteOllamaCloudUsageSession(ctx context.Context, account *service.Account) error { + return r.updateOllamaCloudUsageGroup(ctx, account, map[string]any{}, false) +} + +func (r *accountRepository) SetOllamaCloudUsageAutoRefresh(ctx context.Context, account *service.Account, enabled bool) error { + if !ollamaCloudUsageAccountHasSession(account) { + return service.ErrOllamaCloudUsageSessionRequired + } + payload := ollamaCloudUsageManagedPayload(account) + payload[service.OllamaCloudUsageAutoRefreshExtraKey] = enabled + return r.updateOllamaCloudUsageGroup(ctx, account, payload, true) +} + +func (r *accountRepository) UpdateOllamaCloudUsageSnapshot(ctx context.Context, account *service.Account, snapshot *service.OllamaCloudUsageSnapshot) error { + if account == nil || snapshot == nil { + return service.ErrAccountNilInput + } + if !ollamaCloudUsageAccountHasSession(account) { + return service.ErrOllamaCloudUsageSessionRequired + } + payload := ollamaCloudUsageManagedPayload(account) + payload[service.OllamaCloudUsageSnapshotExtraKey] = snapshot + return r.updateOllamaCloudUsageGroup(ctx, account, payload, true) +} + +// DisableOllamaCloudUsageAutoRefresh is group-scoped and retains the loaded +// identity CAS. It cannot disable a new group after the account changes key. +func (r *accountRepository) DisableOllamaCloudUsageAutoRefresh(ctx context.Context, account *service.Account) error { + if !ollamaCloudUsageAccountHasSession(account) { + return service.ErrOllamaCloudUsageSessionRequired + } + payload := ollamaCloudUsageManagedPayload(account) + payload[service.OllamaCloudUsageAutoRefreshExtraKey] = false + delete(payload, service.OllamaCloudUsageSnapshotExtraKey) + return r.updateOllamaCloudUsageGroup(ctx, account, payload, true) +} + +func ollamaCloudUsageManagedPayload(account *service.Account) map[string]any { + payload := make(map[string]any, 3) + if account == nil || account.Extra == nil { + return payload + } + for _, key := range []string{ + service.OllamaCloudUsageSessionExtraKey, + service.OllamaCloudUsageAutoRefreshExtraKey, + service.OllamaCloudUsageSnapshotExtraKey, + } { + if value, ok := account.Extra[key]; ok { + payload[key] = value + } + } + return payload +} + +func ollamaCloudUsageAccountHasSession(account *service.Account) bool { + if account == nil || account.Extra == nil { + return false + } + value, ok := account.Extra[service.OllamaCloudUsageSessionExtraKey].(string) + return ok && value != "" +} + +type lockedOllamaCloudUsageMember struct { + id int64 + anchorMatches bool + sessionJSON string + autoJSON string + snapshotJSON string +} + +func (r *accountRepository) updateOllamaCloudUsageGroup( + ctx context.Context, + account *service.Account, + payload map[string]any, + requireExpectedState bool, +) error { + if account == nil { + return service.ErrAccountNilInput + } + if r == nil || r.client == nil || !service.IsOllamaCloudUsageAccount(account) { + return service.ErrOllamaCloudUsageUnavailable + } + apiKey, ok := account.Credentials["api_key"].(string) + if !ok || apiKey == "" { + return service.ErrOllamaCloudUsageAccountInvalid + } + apply := func(txCtx context.Context, client *dbent.Client) error { + matchesProxy, err := lockAndMatchProbeProxyIdentity(txCtx, client, account) + if err != nil { + return err + } + if !matchesProxy { + return service.ErrOllamaCloudUsageIdentityChanged + } + members, err := lockOllamaCloudUsageGroup(txCtx, client, account, apiKey) + if err != nil { + return err + } + anchorMatches := false + for _, member := range members { + anchorMatches = anchorMatches || member.anchorMatches + } + if !anchorMatches { + return service.ErrOllamaCloudUsageIdentityChanged + } + if requireExpectedState { + expectedSession, err := canonicalAccountExtraJSON(account, service.OllamaCloudUsageSessionExtraKey) + if err != nil { + return err + } + expectedAuto, err := canonicalAccountExtraJSON(account, service.OllamaCloudUsageAutoRefreshExtraKey) + if err != nil { + return err + } + expectedSnapshot, err := canonicalAccountExtraJSON(account, service.OllamaCloudUsageSnapshotExtraKey) + if err != nil { + return err + } + stateMatches := false + for _, member := range members { + if canonicalJSON(member.sessionJSON) == expectedSession && + canonicalJSON(member.autoJSON) == expectedAuto && + canonicalJSON(member.snapshotJSON) == expectedSnapshot { + stateMatches = true + break + } + } + if !stateMatches { + return service.ErrOllamaCloudUsageIdentityChanged + } + } + encoded, err := json.Marshal(payload) + if err != nil { + return err + } + memberIDs := make([]int64, len(members)) + for index := range members { + memberIDs[index] = members[index].id + } + result, err := client.ExecContext(txCtx, ` + UPDATE accounts + SET extra = (COALESCE(extra, '{}'::jsonb) + - 'ollama_cloud_usage_session' + - 'ollama_cloud_usage_auto_refresh' + - 'ollama_cloud_usage_snapshot') || $1::jsonb, + updated_at = NOW() + WHERE deleted_at IS NULL + AND `+ollamaCloudUsageEligibleSQL+` + AND credentials ->> 'api_key' = $2 + AND id = ANY($3) + `, string(encoded), apiKey, pq.Array(memberIDs)) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected != int64(len(members)) { + return service.ErrOllamaCloudUsageIdentityChanged + } + return nil + } + if dbent.TxFromContext(ctx) != nil { + return apply(ctx, clientFromContext(ctx, r.client)) + } + tx, err := r.client.Tx(ctx) + if errors.Is(err, dbent.ErrTxStarted) { + return apply(ctx, r.client) + } + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + txCtx := dbent.NewTxContext(ctx, tx) + if err := apply(txCtx, tx.Client()); err != nil { + return err + } + return tx.Commit() +} + +func lockOllamaCloudUsageGroup( + ctx context.Context, + client *dbent.Client, + account *service.Account, + apiKey string, +) ([]lockedOllamaCloudUsageMember, error) { + credentials, err := json.Marshal(normalizeJSONMap(account.Credentials)) + if err != nil { + return nil, err + } + var proxyID any + if account.ProxyID != nil { + proxyID = *account.ProxyID + } + rows, err := client.QueryContext(ctx, ` + SELECT + id, + id = $2 + AND platform = $3 + AND type = $4 + AND credentials = $5::jsonb + AND proxy_id IS NOT DISTINCT FROM $6, + COALESCE((extra -> 'ollama_cloud_usage_session')::text, 'null'), + COALESCE((extra -> 'ollama_cloud_usage_auto_refresh')::text, 'null'), + COALESCE((extra -> 'ollama_cloud_usage_snapshot')::text, 'null') + FROM accounts + WHERE deleted_at IS NULL + AND `+ollamaCloudUsageEligibleSQL+` + AND credentials ->> 'api_key' = $1 + ORDER BY id + FOR NO KEY UPDATE + `, apiKey, account.ID, account.Platform, account.Type, string(credentials), proxyID) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + members := make([]lockedOllamaCloudUsageMember, 0, 1) + for rows.Next() { + var member lockedOllamaCloudUsageMember + if err := rows.Scan(&member.id, &member.anchorMatches, &member.sessionJSON, &member.autoJSON, &member.snapshotJSON); err != nil { + return nil, err + } + members = append(members, member) + } + if err := rows.Err(); err != nil { + return nil, err + } + if len(members) == 0 { + return nil, service.ErrOllamaCloudUsageIdentityChanged + } + return members, nil +} + +func canonicalAccountExtraJSON(account *service.Account, key string) (string, error) { + var value any + if account != nil && account.Extra != nil { + value = account.Extra[key] + } + raw, err := json.Marshal(value) + if err != nil { + return "", err + } + return canonicalJSON(string(raw)), nil +} + +func canonicalJSON(raw string) string { + var value any + if err := json.Unmarshal([]byte(raw), &value); err != nil { + return "" + } + encoded, err := json.Marshal(value) + if err != nil { + return "" + } + return string(encoded) +} + +// ListDueOllamaCloudUsageAccounts returns at most one due representative per +// exact API key before hydration, preventing one shared group from consuming a +// whole runner cycle. +func (r *accountRepository) ListDueOllamaCloudUsageAccounts(ctx context.Context, now time.Time, limit int) ([]service.Account, error) { + if limit <= 0 { + return []service.Account{}, nil + } + if r == nil || r.sql == nil { + return nil, errors.New("account repository SQL executor not configured") + } + rows, err := r.sql.QueryContext(ctx, ` + WITH candidates AS ( + SELECT id, credentials ->> 'api_key' AS api_key, + extra #>> '{ollama_cloud_usage_snapshot,next_refresh_at}' AS next_refresh_at + FROM accounts + WHERE deleted_at IS NULL + AND status = 'active' + AND `+ollamaCloudUsageEligibleSQL+` + AND jsonb_typeof(extra -> 'ollama_cloud_usage_session') = 'string' + AND extra @> '{"ollama_cloud_usage_auto_refresh": true}'::jsonb + ), parsed AS MATERIALIZED ( + SELECT id, api_key, next_refresh_at, + next_refresh_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( + to_jsonb(regexp_replace( + next_refresh_at, + '(\.[0-9]{6})[0-9]+(Z|[+-][0-9]{2}:[0-9]{2})$', + '\1\2' + )), + '$.datetime()', '{}'::jsonb, true + ) #>> '{}' AS parsed_next_refresh_at + FROM candidates + ), due AS ( + SELECT *, + CASE WHEN next_refresh_at IS NULL OR NOT rfc3339_shape OR parsed_next_refresh_at IS NULL THEN 0 ELSE 1 END AS due_class + FROM parsed + WHERE next_refresh_at IS NULL + OR NOT rfc3339_shape + OR parsed_next_refresh_at IS NULL + OR parsed_next_refresh_at::timestamptz <= $1 + ), ranked AS ( + SELECT *, row_number() OVER ( + PARTITION BY api_key + ORDER BY due_class, + CASE WHEN rfc3339_shape AND parsed_next_refresh_at IS NOT NULL THEN parsed_next_refresh_at::timestamptz END NULLS FIRST, + id + ) AS group_rank + FROM due + ) + SELECT id + FROM ranked + WHERE group_rank = 1 + ORDER BY due_class, + CASE WHEN rfc3339_shape AND parsed_next_refresh_at IS NOT NULL THEN parsed_next_refresh_at::timestamptz END NULLS FIRST, + id + 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 + } + accounts, err := r.GetByIDs(ctx, ids) + if err != nil { + return nil, err + } + result := make([]service.Account, 0, len(accounts)) + for _, account := range accounts { + if account != nil { + result = append(result, *account) + } + } + return result, nil +} diff --git a/backend/internal/repository/account_repo_ollama_cloud_usage_integration_test.go b/backend/internal/repository/account_repo_ollama_cloud_usage_integration_test.go new file mode 100644 index 0000000000..f261323439 --- /dev/null +++ b/backend/internal/repository/account_repo_ollama_cloud_usage_integration_test.go @@ -0,0 +1,362 @@ +//go:build integration + +package repository + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestListDueOllamaCloudUsageAccountsOrderingLimitAndProxyHydration(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + now := time.Date(2026, time.July, 22, 12, 0, 0, 0, time.UTC) + proxy := mustCreateProxy(t, tx.Client(), &service.Proxy{ + Name: "ollama-due-proxy", Protocol: "http", Host: "127.0.0.1", Port: 3128, + Username: "user", Password: "pass", Status: service.StatusActive, + }) + + createAccount := func(name, baseURL string, proxyID *int64, nextRefreshAt *time.Time) *service.Account { + t.Helper() + extra := map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:wos-session=fixture", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + } + if nextRefreshAt != nil { + extra[service.OllamaCloudUsageSnapshotExtraKey] = map[string]any{ + "status": service.OllamaCloudUsageStatusOK, "next_refresh_at": nextRefreshAt.UTC().Format(time.RFC3339Nano), + } + } + return mustCreateAccount(t, tx.Client(), &service.Account{ + Name: name, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": name, "base_url": baseURL}, + Extra: extra, ProxyID: proxyID, + }) + } + + uppercasePath := createAccount("ollama-uppercase-path", "https://ollama.com/V1", nil, nil) + missingSnapshot := createAccount("ollama-due-missing", "HTTPS://WWW.OLLAMA.COM:443/v1", &proxy.ID, nil) + oldest := now.Add(-2 * time.Hour) + due := createAccount("ollama-due-oldest", "https://ollama.com", nil, &oldest) + future := now.Add(time.Minute) + _ = createAccount("ollama-not-due", "https://ollama.com", nil, &future) + _ = createAccount("ollama-ineligible", "https://ollama.com.evil.test", nil, nil) + + accounts, err := repo.ListDueOllamaCloudUsageAccounts(ctx, now, 2) + + require.NoError(t, err) + require.Len(t, accounts, 2) + require.Equal(t, missingSnapshot.ID, accounts[0].ID) + require.Equal(t, due.ID, accounts[1].ID) + require.NotContains(t, accountIDs(accounts), uppercasePath.ID) + require.NotNil(t, accounts[0].Proxy) + require.Equal(t, proxy.ID, accounts[0].Proxy.ID) + require.Equal(t, proxy.URL(), accounts[0].Proxy.URL()) +} + +func TestListDueOllamaCloudUsageAccountsParsesRFC3339NanoAndFailsOpen(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + now := time.Date(2026, time.July, 22, 14, 0, 0, 0, time.UTC) + + create := func(name, nextRefreshAt string) *service.Account { + t.Helper() + return mustCreateAccount(t, tx.Client(), &service.Account{ + Name: name, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": name, "base_url": "https://ollama.com"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:wos-session=fixture", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + service.OllamaCloudUsageSnapshotExtraKey: map[string]any{ + "status": service.OllamaCloudUsageStatusOK, "next_refresh_at": nextRefreshAt, + }, + }, + }) + } + + sevenDigitsOffset := create("ollama-nano-seven", "2026-07-22T11:00:00.1234567-02:00") + eightDigitsOffset := create("ollama-nano-eight", "2026-07-22T11:00:00.12345678+01:00") + nineDigitsZ := create("ollama-nano-nine", "2026-07-22T09:00:00.123456789Z") + invalidCalendar := create("ollama-nano-invalid", "2026-02-30T09:00:00.123456789Z") + future := create("ollama-nano-future", "2026-07-22T15:00:00.123456789Z") + + accounts, err := repo.ListDueOllamaCloudUsageAccounts(ctx, now, 10) + + require.NoError(t, err, "invalid stored values must not abort the query") + require.Equal(t, []int64{ + invalidCalendar.ID, + nineDigitsZ.ID, + eightDigitsOffset.ID, + sevenDigitsOffset.ID, + }, accountIDs(accounts)) + require.NotContains(t, accountIDs(accounts), future.ID) +} + +func TestLockAndMergeAccountProbeExtraCoalescesNullableOllamaGroupIdentity(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "ordinary-openai-without-base-url", Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-no-base-url"}, + Extra: map[string]any{service.UpstreamBillingProbeEnabledExtraKey: true}, + }) + loaded, err := newAccountRepositoryWithSQL(tx.Client(), tx, nil).GetByID(ctx, account.ID) + require.NoError(t, err) + + merged, err := lockAndMergeAccountProbeExtra(ctx, tx.Client(), loaded, nil) + + require.NoError(t, err, "a NULL Ollama eligibility expression must scan as false") + require.NotContains(t, merged, service.OllamaCloudUsageSessionExtraKey) + require.Equal(t, true, merged[service.UpstreamBillingProbeEnabledExtraKey]) +} + +func TestOllamaCloudUsageGroupWritesAreAtomicAcrossPlatformsAndURLVariants(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + create := func(name, platform, apiKey, baseURL string) *service.Account { + t.Helper() + return mustCreateAccount(t, tx.Client(), &service.Account{ + Name: name, Platform: platform, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": apiKey, "base_url": baseURL}, + Extra: map[string]any{}, + }) + } + first := create("ollama-group-openai", service.PlatformOpenAI, "shared-key", "https://ollama.com") + second := create("ollama-group-anthropic", service.PlatformAnthropic, "shared-key", "HTTPS://WWW.OLLAMA.COM:443/v1") + different := create("ollama-group-different", service.PlatformOpenAI, "different-key", "https://ollama.com") + + require.NoError(t, repo.SaveOllamaCloudUsageSession(ctx, first, "cipher:shared", false)) + for _, id := range []int64{first.ID, second.ID} { + account, err := repo.GetByID(ctx, id) + require.NoError(t, err) + require.Equal(t, "cipher:shared", account.Extra[service.OllamaCloudUsageSessionExtraKey]) + require.Equal(t, false, account.Extra[service.OllamaCloudUsageAutoRefreshExtraKey]) + } + differentLoaded, err := repo.GetByID(ctx, different.ID) + require.NoError(t, err) + require.NotContains(t, differentLoaded.Extra, service.OllamaCloudUsageSessionExtraKey) + + secondLoaded, err := repo.GetByID(ctx, second.ID) + require.NoError(t, err) + require.NoError(t, repo.SetOllamaCloudUsageAutoRefresh(ctx, secondLoaded, true)) + firstLoaded, err := repo.GetByID(ctx, first.ID) + require.NoError(t, err) + secondLoaded, err = repo.GetByID(ctx, second.ID) + require.NoError(t, err) + require.Equal(t, true, firstLoaded.Extra[service.OllamaCloudUsageAutoRefreshExtraKey]) + require.Equal(t, true, secondLoaded.Extra[service.OllamaCloudUsageAutoRefreshExtraKey]) + + now := time.Now().UTC() + snapshot := &service.OllamaCloudUsageSnapshot{ + Status: service.OllamaCloudUsageStatusOK, LastAttemptAt: now, NextRefreshAt: now.Add(time.Hour), + } + require.NoError(t, repo.UpdateOllamaCloudUsageSnapshot(ctx, firstLoaded, snapshot)) + secondLoaded, err = repo.GetByID(ctx, second.ID) + require.NoError(t, err) + require.Equal(t, service.OllamaCloudUsageStatusOK, + secondLoaded.Extra[service.OllamaCloudUsageSnapshotExtraKey].(map[string]any)["status"]) + + staleSecond := secondLoaded + require.NoError(t, repo.UpdateCredentials(ctx, second.ID, map[string]any{ + "api_key": "rotated-key", "base_url": "https://ollama.com", + })) + require.ErrorIs(t, repo.DisableOllamaCloudUsageAutoRefresh(ctx, staleSecond), service.ErrOllamaCloudUsageIdentityChanged) + firstLoaded, err = repo.GetByID(ctx, first.ID) + require.NoError(t, err) + secondLoaded, err = repo.GetByID(ctx, second.ID) + require.NoError(t, err) + require.Equal(t, "cipher:shared", firstLoaded.Extra[service.OllamaCloudUsageSessionExtraKey]) + require.Equal(t, true, firstLoaded.Extra[service.OllamaCloudUsageAutoRefreshExtraKey]) + require.NotContains(t, secondLoaded.Extra, service.OllamaCloudUsageSessionExtraKey) + require.NotContains(t, secondLoaded.Extra, service.OllamaCloudUsageAutoRefreshExtraKey) + + require.NoError(t, repo.DeleteOllamaCloudUsageSession(ctx, firstLoaded)) + firstLoaded, err = repo.GetByID(ctx, first.ID) + require.NoError(t, err) + require.NotContains(t, firstLoaded.Extra, service.OllamaCloudUsageSessionExtraKey) +} + +func TestConcurrentOllamaCloudUsageSaveAndDeleteSerializeGroupState(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + client := testEntClient(t) + repo := newAccountRepositoryWithSQL(client, integrationDB, nil) + suffix := time.Now().UnixNano() + apiKey := fmt.Sprintf("ollama-concurrent-%d", suffix) + create := func(platform string) *service.Account { + t.Helper() + return mustCreateAccount(t, client, &service.Account{ + Name: fmt.Sprintf("%s-%s", apiKey, platform), Platform: platform, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": apiKey, "base_url": "https://ollama.com"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:initial", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + }, + }) + } + first := create(service.PlatformOpenAI) + second := create(service.PlatformAnthropic) + t.Cleanup(func() { + _, _ = integrationDB.ExecContext(context.Background(), "DELETE FROM accounts WHERE id IN ($1, $2)", first.ID, second.ID) + }) + anchor, err := repo.GetByID(ctx, first.ID) + require.NoError(t, err) + + start := make(chan struct{}) + errs := make(chan error, 2) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + <-start + errs <- repo.SaveOllamaCloudUsageSession(ctx, anchor, "cipher:replacement", true) + }() + go func() { + defer wg.Done() + <-start + errs <- repo.DeleteOllamaCloudUsageSession(ctx, anchor) + }() + close(start) + wg.Wait() + close(errs) + for writeErr := range errs { + require.NoError(t, writeErr) + } + + firstLoaded, err := repo.GetByID(ctx, first.ID) + require.NoError(t, err) + secondLoaded, err := repo.GetByID(ctx, second.ID) + require.NoError(t, err) + managedState := func(account *service.Account) map[string]any { + state := make(map[string]any) + for _, key := range []string{ + service.OllamaCloudUsageSessionExtraKey, + service.OllamaCloudUsageAutoRefreshExtraKey, + service.OllamaCloudUsageSnapshotExtraKey, + } { + if value, ok := account.Extra[key]; ok { + state[key] = value + } + } + return state + } + firstState := managedState(firstLoaded) + require.Equal(t, firstState, managedState(secondLoaded), "a serialized last commit must own the whole group") + if len(firstState) > 0 { + require.Equal(t, "cipher:replacement", firstState[service.OllamaCloudUsageSessionExtraKey]) + require.Equal(t, true, firstState[service.OllamaCloudUsageAutoRefreshExtraKey]) + require.NotContains(t, firstState, service.OllamaCloudUsageSnapshotExtraKey) + } +} + +func accountIDs(accounts []service.Account) []int64 { + ids := make([]int64, len(accounts)) + for index := range accounts { + ids[index] = accounts[index].ID + } + return ids +} + +func TestOllamaCloudUsageCredentialAndBulkUpdatesPreserveManagedStateOnlyWhenSafe(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + repo := newAccountRepositoryWithSQL(tx.Client(), tx, nil) + now := time.Now().UTC() + 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": "old-key", "base_url": "https://ollama.com"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:wos-session=fixture", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + service.OllamaCloudUsageSnapshotExtraKey: map[string]any{ + "status": service.OllamaCloudUsageStatusOK, "last_attempt_at": now, "next_refresh_at": now.Add(time.Hour), + }, + }, + }) + } + + rawAccount := newAccount("ollama-raw-credentials") + require.NoError(t, repo.UpdateCredentials(ctx, rawAccount.ID, map[string]any{ + "api_key": "old-key", "base_url": "https://ollama.com/V1", + })) + rawUpdated, err := repo.GetByID(ctx, rawAccount.ID) + require.NoError(t, err) + require.NotContains(t, rawUpdated.Extra, service.OllamaCloudUsageSessionExtraKey) + require.NotContains(t, rawUpdated.Extra, service.OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, rawUpdated.Extra, service.OllamaCloudUsageSnapshotExtraKey) + + bulkAccount := newAccount("ollama-bulk-credentials") + rows, err := repo.BulkUpdate(ctx, []int64{bulkAccount.ID}, service.AccountBulkUpdate{ + Credentials: map[string]any{"base_url": "HTTPS://WWW.OLLAMA.COM:443/v1"}, + }) + require.NoError(t, err) + require.Equal(t, int64(1), rows) + bulkUnchanged, err := repo.GetByID(ctx, bulkAccount.ID) + require.NoError(t, err) + require.Contains(t, bulkUnchanged.Extra, service.OllamaCloudUsageSnapshotExtraKey) + + rows, err = repo.BulkUpdate(ctx, []int64{bulkAccount.ID}, service.AccountBulkUpdate{ + Credentials: map[string]any{"base_url": "https://ollama.com/V1"}, + }) + require.NoError(t, err) + require.Equal(t, int64(1), rows) + bulkIneligible, err := repo.GetByID(ctx, bulkAccount.ID) + require.NoError(t, err) + require.NotContains(t, bulkIneligible.Extra, service.OllamaCloudUsageSessionExtraKey) + require.NotContains(t, bulkIneligible.Extra, service.OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, bulkIneligible.Extra, service.OllamaCloudUsageSnapshotExtraKey) +} + +func TestProxyIdentityUpdateInvalidatesOllamaSnapshotAndRejectsInFlightCAS(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: "ollama-identity-proxy", Protocol: "http", Host: "old.example", Port: 8080, + Username: "old-user", Password: "old-pass", Status: service.StatusActive, + }) + now := time.Now().UTC() + account := mustCreateAccount(t, tx.Client(), &service.Account{ + Name: "ollama-proxy-account", Platform: service.PlatformAnthropic, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://ollama.com"}, + ProxyID: &proxy.ID, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:wos-session=fixture", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + service.OllamaCloudUsageSnapshotExtraKey: map[string]any{ + "status": service.OllamaCloudUsageStatusOK, "last_attempt_at": now, "next_refresh_at": now.Add(time.Hour), + }, + }, + }) + 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) + require.NotContains(t, got.Extra, service.OllamaCloudUsageSnapshotExtraKey) + require.Equal(t, "cipher:wos-session=fixture", got.Extra[service.OllamaCloudUsageSessionExtraKey]) + require.Equal(t, true, got.Extra[service.OllamaCloudUsageAutoRefreshExtraKey]) + + err = accountRepo.UpdateOllamaCloudUsageSnapshot(ctx, inFlight, &service.OllamaCloudUsageSnapshot{ + Status: service.OllamaCloudUsageStatusOK, LastAttemptAt: now, NextRefreshAt: now.Add(time.Hour), + }) + require.ErrorIs(t, err, service.ErrOllamaCloudUsageIdentityChanged) +} diff --git a/backend/internal/repository/account_repo_ollama_cloud_usage_test.go b/backend/internal/repository/account_repo_ollama_cloud_usage_test.go new file mode 100644 index 0000000000..a7f5bb92a4 --- /dev/null +++ b/backend/internal/repository/account_repo_ollama_cloud_usage_test.go @@ -0,0 +1,279 @@ +package repository + +import ( + "context" + "encoding/json" + "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 newOllamaCloudUsageRepositoryTestClient(t *testing.T) (*dbent.Client, sqlmock.Sqlmock) { + t.Helper() + 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() }) + return client, mock +} + +func ollamaCloudUsageRepositoryAccount() *service.Account { + return &service.Account{ + ID: 17, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://ollama.com"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "cipher:wos-session=secret", + service.OllamaCloudUsageAutoRefreshExtraKey: true, + }, + } +} + +func TestUpdateOllamaCloudUsageSnapshotRowsAffectedZeroIsIdentityConflict(t *testing.T) { + client, mock := newOllamaCloudUsageRepositoryTestClient(t) + mock.ExpectBegin() + expectOllamaCloudUsageGroupLock(mock, ollamaCloudUsageRepositoryAccount(), true, + `"cipher:wos-session=secret"`, `true`, `null`) + mock.ExpectExec(`(?s)`+regexp.QuoteMeta("UPDATE accounts")). + WithArgs(sqlmock.AnyArg(), "key", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectRollback() + repo := newAccountRepositoryWithSQL(client, nil, nil) + + err := repo.UpdateOllamaCloudUsageSnapshot(context.Background(), ollamaCloudUsageRepositoryAccount(), &service.OllamaCloudUsageSnapshot{ + Status: service.OllamaCloudUsageStatusOK, + LastAttemptAt: time.Now(), + NextRefreshAt: time.Now().Add(time.Hour), + }) + + require.ErrorIs(t, err, service.ErrOllamaCloudUsageIdentityChanged) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func expectOllamaCloudUsageGroupLock( + mock sqlmock.Sqlmock, + account *service.Account, + anchorMatches bool, + sessionJSON, autoJSON, snapshotJSON string, +) { + apiKey, _ := account.Credentials["api_key"].(string) + credentials, _ := json.Marshal(normalizeJSONMap(account.Credentials)) + var proxyID any + if account.ProxyID != nil { + proxyID = *account.ProxyID + } + mock.ExpectQuery(`(?s)`+regexp.QuoteMeta("SELECT")+`.*`+regexp.QuoteMeta("FOR NO KEY UPDATE")). + WithArgs(apiKey, account.ID, account.Platform, account.Type, string(credentials), proxyID). + WillReturnRows(sqlmock.NewRows([]string{"id", "anchor_matches", "session", "auto_refresh", "snapshot"}). + AddRow(account.ID, anchorMatches, sessionJSON, autoJSON, snapshotJSON)) +} + +func TestOllamaCloudUsageManagedWriteRejectsChangedProxyIdentity(t *testing.T) { + client, mock := newOllamaCloudUsageRepositoryTestClient(t) + mock.ExpectBegin() + 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)) + mock.ExpectRollback() + + account := ollamaCloudUsageRepositoryAccount() + proxyID := int64(9) + account.ProxyID = &proxyID + account.Proxy = &service.Proxy{ + ID: proxyID, Protocol: "http", Host: "old.example", Port: 3128, + Username: "user", Password: "pass", Status: service.StatusActive, + } + repo := newAccountRepositoryWithSQL(client, nil, nil) + + err := repo.SaveOllamaCloudUsageSession(context.Background(), account, "cipher:wos-session=replacement", true) + + require.ErrorIs(t, err, service.ErrOllamaCloudUsageIdentityChanged) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestSaveAndDeleteOllamaCloudUsageSessionKeepCiphertextOutOfSQL(t *testing.T) { + var capturedSQL []string + matcher := sqlmock.QueryMatcherFunc(func(expectedSQL, actualSQL string) error { + capturedSQL = append(capturedSQL, actualSQL) + return sqlmock.QueryMatcherRegexp.Match(expectedSQL, actualSQL) + }) + db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(matcher)) + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + client := dbent.NewClient(dbent.Driver(entsql.OpenDB(dialect.Postgres, db))) + t.Cleanup(func() { _ = client.Close() }) + repo := newAccountRepositoryWithSQL(client, db, nil) + account := ollamaCloudUsageRepositoryAccount() + const replacement = "cipher:wos-session=browser-cookie-secret" + + mock.ExpectBegin() + expectOllamaCloudUsageGroupLock(mock, account, true, `"cipher:wos-session=secret"`, `true`, `null`) + mock.ExpectExec(`(?s)UPDATE accounts.*ollama_cloud_usage_session.*ollama_cloud_usage_auto_refresh.*ollama_cloud_usage_snapshot`). + WithArgs(`{"ollama_cloud_usage_auto_refresh":true,"ollama_cloud_usage_session":"cipher:wos-session=browser-cookie-secret"}`, "key", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + require.NoError(t, repo.SaveOllamaCloudUsageSession(context.Background(), account, replacement, true)) + + account.Extra[service.OllamaCloudUsageSessionExtraKey] = replacement + mock.ExpectBegin() + expectOllamaCloudUsageGroupLock(mock, account, true, `"cipher:wos-session=browser-cookie-secret"`, `true`, `null`) + mock.ExpectExec(`(?s)UPDATE accounts.*ollama_cloud_usage_session.*ollama_cloud_usage_auto_refresh.*ollama_cloud_usage_snapshot`). + WithArgs(`{}`, "key", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + require.NoError(t, repo.DeleteOllamaCloudUsageSession(context.Background(), account)) + + require.NotEmpty(t, capturedSQL) + for _, query := range capturedSQL { + require.NotContains(t, query, "browser-cookie-secret") + } + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestOllamaCloudBaseURLSQLRegexMatchesServiceSemantics(t *testing.T) { + for _, baseURL := range []string{ + "https://ollama.com", + "HTTPS://WWW.OLLAMA.COM:443/v1", + "https://ollama.com/V1", + "https://ollama.com/v1/", + "https://ollama.com.evil.test/v1", + } { + t.Run(baseURL, func(t *testing.T) { + matched, err := regexp.MatchString(ollamaCloudBaseURLRegexSQL, baseURL) + require.NoError(t, err) + account := ollamaCloudUsageRepositoryAccount() + account.Credentials["base_url"] = baseURL + require.Equal(t, service.IsOllamaCloudUsageAccount(account), matched) + }) + } +} + +func TestListOllamaCloudUsageGroupAccountsUsesOneStrictBatchQuery(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + var capturedSQL string + mock.ExpectQuery("SELECT id"). + WithArgs(sqlmock.AnyArg()). + WillReturnRows(sqlmock.NewRows([]string{"id"})) + repo := newAccountRepositoryWithSQL(nil, captureQuerySQL{db: db, captured: &capturedSQL}, nil) + first := ollamaCloudUsageRepositoryAccount() + second := ollamaCloudUsageRepositoryAccount() + second.ID = 18 + second.Platform = service.PlatformAnthropic + second.Credentials = map[string]any{"api_key": "key", "base_url": "https://www.ollama.com:443/v1"} + + accounts, err := repo.ListOllamaCloudUsageGroupAccounts(context.Background(), []*service.Account{first, second}) + + require.NoError(t, err) + require.Empty(t, accounts) + query := normalizeSQLWhitespace(capturedSQL) + require.Contains(t, query, "credentials ->> 'api_key' = ANY($1)") + require.Contains(t, query, "platform IN ('openai', 'anthropic')") + require.Contains(t, query, "jsonb_typeof(credentials -> 'api_key') = 'string'") + require.Contains(t, query, ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'")) + require.NotContains(t, query, "~*") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestListDueOllamaCloudUsageAccountsFiltersOrdersAndLimits(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + now := time.Date(2026, time.July, 22, 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.ListDueOllamaCloudUsageAccounts(context.Background(), now, 20) + + require.NoError(t, err) + require.Empty(t, accounts) + normalized := normalizeSQLWhitespace(capturedSQL) + for _, clause := range []string{ + "deleted_at IS NULL", + "status = 'active'", + "platform IN ('openai', 'anthropic')", + "type = 'apikey'", + ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'"), + "jsonb_typeof(extra -> 'ollama_cloud_usage_session') = 'string'", + `extra @> '{"ollama_cloud_usage_auto_refresh": true}'::jsonb`, + "parsed_next_refresh_at::timestamptz <= $1", + "PARTITION BY api_key", + "WHERE group_rank = 1", + "LIMIT $2", + } { + require.Contains(t, normalized, clause) + } + require.NotContains(t, normalized, "~*") + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestBulkUpdateOllamaIdentityCleanupIsValueConditional(t *testing.T) { + exec := &recordingSQLExecutor{result: rowsAffectedResult(1)} + repo := newAccountRepositoryWithSQL(nil, exec, nil) + + _, err := repo.BulkUpdate(context.Background(), []int64{17}, service.AccountBulkUpdate{ + Credentials: map[string]any{"base_url": "https://www.ollama.com:443/v1"}, + }) + + require.NoError(t, err) + require.NotEmpty(t, exec.execQueries) + query := normalizeSQLWhitespace(exec.execQueries[0]) + require.Contains(t, query, "NOT ("+ollamaCloudBaseURLMatchesSQL("credentials ->> 'base_url'")) + require.Contains(t, query, ollamaCloudBaseURLMatchesSQL("$1::jsonb ->> 'base_url'")) + require.NotContains(t, query, "~*") + require.Contains(t, query, "platform IN ('openai', 'anthropic') AND type = 'apikey'") + require.Contains(t, query, "- 'ollama_cloud_usage_session' - 'ollama_cloud_usage_auto_refresh' - 'ollama_cloud_usage_snapshot'") + payload, ok := exec.execArgs[0][0].([]byte) + require.True(t, ok) + require.NotContains(t, string(payload), service.OllamaCloudUsageSnapshotExtraKey) +} + +func TestUpdateCredentialsIdentityChangeClearsAllOllamaManagedExtra(t *testing.T) { + client, mock := newOllamaCloudUsageRepositoryTestClient(t) + mock.ExpectBegin() + mock.ExpectExec(`(?s)UPDATE accounts.*credentials -> 'api_key' IS DISTINCT FROM.*ollama_cloud_usage_session.*ollama_cloud_usage_auto_refresh.*ollama_cloud_usage_snapshot`). + WithArgs(`{"api_key":"new-key","base_url":"https://ollama.com"}`, int64(17)). + 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, nil, nil) + + err := repo.UpdateCredentials(context.Background(), 17, map[string]any{ + "api_key": "new-key", "base_url": "https://ollama.com", + }) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestDisableOllamaCloudUsageAutoRefreshUsesGroupIdentityCAS(t *testing.T) { + client, mock := newOllamaCloudUsageRepositoryTestClient(t) + account := ollamaCloudUsageRepositoryAccount() + mock.ExpectBegin() + expectOllamaCloudUsageGroupLock(mock, account, true, `"cipher:wos-session=secret"`, `true`, `null`) + mock.ExpectExec(`(?s)UPDATE accounts.*ollama_cloud_usage_auto_refresh`). + WithArgs(`{"ollama_cloud_usage_auto_refresh":false,"ollama_cloud_usage_session":"cipher:wos-session=secret"}`, "key", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + repo := newAccountRepositoryWithSQL(client, nil, nil) + + err := repo.DisableOllamaCloudUsageAutoRefresh(context.Background(), account) + + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go b/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go index 2a78f38c10..497059ba5a 100644 --- a/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go +++ b/backend/internal/repository/account_repo_upstream_billing_probe_update_test.go @@ -81,8 +81,8 @@ func TestLockAndMergeAccountProbeExtraUsesCurrentDatabaseSnapshot(t *testing.T) 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)) + WillReturnRows(sqlmock.NewRows([]string{"identity_unchanged", "ollama_group_unchanged", "ollama_proxy_unchanged", "enabled", "snapshot", "ollama_session", "ollama_auto", "ollama_snapshot"}). + AddRow(tt.identityUnchanged, false, true, tt.databaseEnabled, tt.databaseSnapshot, nil, nil, nil)) account := &service.Account{ ID: 27, @@ -104,6 +104,45 @@ func TestLockAndMergeAccountProbeExtraUsesCurrentDatabaseSnapshot(t *testing.T) } } +func TestLockAndMergeAccountProbeExtraProtectsOllamaManagedFields(t *testing.T) { + for _, identityUnchanged := range []bool{true, false} { + t.Run(map[bool]string{true: "same identity keeps snapshot", false: "changed identity clears snapshot"}[identityUnchanged], 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(29), service.PlatformAnthropic, service.AccountTypeAPIKey, `{"api_key":"key","base_url":"https://ollama.com"}`, nil). + WillReturnRows(sqlmock.NewRows([]string{"identity_unchanged", "ollama_group_unchanged", "ollama_proxy_unchanged", "enabled", "snapshot", "ollama_session", "ollama_auto", "ollama_snapshot"}). + AddRow(identityUnchanged, identityUnchanged, true, nil, nil, []byte(`"local-ciphertext"`), []byte(`true`), []byte(`{"status":"ok"}`))) + + account := &service.Account{ + ID: 29, Platform: service.PlatformAnthropic, Type: service.AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "key", "base_url": "https://ollama.com"}, + Extra: map[string]any{ + service.OllamaCloudUsageSessionExtraKey: "forged-ciphertext", + service.OllamaCloudUsageAutoRefreshExtraKey: false, + service.OllamaCloudUsageSnapshotExtraKey: map[string]any{"status": "forged"}, + }, + } + got, err := lockAndMergeAccountProbeExtra(context.Background(), client, account, nil) + require.NoError(t, err) + if identityUnchanged { + require.Equal(t, "local-ciphertext", got[service.OllamaCloudUsageSessionExtraKey]) + require.Equal(t, true, got[service.OllamaCloudUsageAutoRefreshExtraKey]) + require.Equal(t, map[string]any{"status": "ok"}, got[service.OllamaCloudUsageSnapshotExtraKey]) + } else { + require.NotContains(t, got, service.OllamaCloudUsageSessionExtraKey) + require.NotContains(t, got, service.OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, got, service.OllamaCloudUsageSnapshotExtraKey) + } + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + func TestUpdateExtraExplicitProbeDisableRemovesSnapshot(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) @@ -236,8 +275,8 @@ func TestUpdateWithUpstreamBillingProbeEnabledRollsBackWhenOutboxFails(t *testin 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"}`))) + WillReturnRows(sqlmock.NewRows([]string{"identity_unchanged", "ollama_group_unchanged", "ollama_proxy_unchanged", "enabled", "snapshot", "ollama_session", "ollama_auto", "ollama_snapshot"}). + AddRow(true, false, true, []byte(`true`), []byte(`{"status":"ok"}`), nil, nil, nil)) mock.ExpectExec(`(?s)UPDATE .*accounts.*SET.*WHERE .*id.*`). WillReturnResult(sqlmock.NewResult(0, 1)) mock.ExpectQuery(`(?s)SELECT .* FROM "accounts" WHERE "id" = \$1`). diff --git a/backend/internal/repository/proxy_repo.go b/backend/internal/repository/proxy_repo.go index 6801f074e1..5eda079390 100644 --- a/backend/internal/repository/proxy_repo.go +++ b/backend/internal/repository/proxy_repo.go @@ -224,12 +224,20 @@ func lockProxyProbeIdentity(ctx context.Context, client *dbent.Client, proxyID i 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() + SET extra = COALESCE(extra, '{}'::jsonb) + - 'upstream_billing_probe' + - 'ollama_cloud_usage_snapshot', + 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 ( + (platform = 'openai' + AND extra ? 'upstream_billing_probe' + AND extra -> 'upstream_billing_probe' <> 'null'::jsonb) + OR (platform IN ('openai', 'anthropic') + AND extra ? 'ollama_cloud_usage_snapshot' + AND extra -> 'ollama_cloud_usage_snapshot' <> 'null'::jsonb) + ) AND deleted_at IS NULL RETURNING id `, proxyID) diff --git a/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go b/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go index 27a3ac1242..d4cd0ad049 100644 --- a/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go +++ b/backend/internal/repository/proxy_repo_upstream_billing_probe_test.go @@ -33,7 +33,7 @@ func TestProxyUpdateInvalidatesBoundProbeSnapshotsAndEnqueuesOutboxAtomically(t 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`). + mock.ExpectQuery(`(?s)UPDATE accounts.*- 'upstream_billing_probe'.*- 'ollama_cloud_usage_snapshot'.*type = 'apikey'.*platform = 'openai'.*extra \? 'upstream_billing_probe'.*platform IN \('openai', 'anthropic'\).*extra \? 'ollama_cloud_usage_snapshot'.*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)")). @@ -76,7 +76,7 @@ func TestProxyUpdateRollsBackWhenProbeInvalidationOutboxFails(t *testing.T) { 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`). + mock.ExpectQuery(`(?s)UPDATE accounts.*- 'upstream_billing_probe'.*- 'ollama_cloud_usage_snapshot'.*type = 'apikey'.*platform = 'openai'.*extra \? 'upstream_billing_probe'.*platform IN \('openai', 'anthropic'\).*extra \? 'ollama_cloud_usage_snapshot'.*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)")). diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index b64ec64b56..ddb78a4160 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -346,6 +346,8 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers, stepUpAu 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("/ollama-cloud-usage/settings", h.Admin.Account.GetOllamaCloudUsageSettings) + accounts.PUT("/ollama-cloud-usage/settings", h.Admin.Account.UpdateOllamaCloudUsageSettings) accounts.GET("/:id", h.Admin.Account.GetByID) accounts.POST("", h.Admin.Account.Create) accounts.POST("/:id/duplicate", h.Admin.Account.Duplicate) @@ -356,6 +358,11 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers, stepUpAu 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.GET("/:id/ollama-cloud-usage", h.Admin.Account.GetOllamaCloudUsage) + accounts.PUT("/:id/ollama-cloud-usage/session", h.Admin.Account.SaveOllamaCloudUsageSession) + accounts.DELETE("/:id/ollama-cloud-usage/session", h.Admin.Account.DeleteOllamaCloudUsageSession) + accounts.PUT("/:id/ollama-cloud-usage/auto-refresh", h.Admin.Account.SetOllamaCloudUsageAutoRefresh) + accounts.POST("/:id/ollama-cloud-usage/refresh", h.Admin.Account.RefreshOllamaCloudUsage) accounts.DELETE("/:id", h.Admin.Account.Delete) accounts.POST("/:id/test", h.Admin.Account.Test) accounts.POST("/:id/recover-state", h.Admin.Account.RecoverState) diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index c328210009..98263dca1c 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -315,7 +315,14 @@ func (s *AccountService) Update(ctx context.Context, id int64, req UpdateAccount } if req.Extra != nil { - account.Extra = *req.Extra + extra := make(map[string]any, len(*req.Extra)) + for key, value := range *req.Extra { + extra[key] = value + } + delete(extra, OllamaCloudUsageSessionExtraKey) + delete(extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(extra, OllamaCloudUsageSnapshotExtraKey) + account.Extra = extra } if req.ProxyID != nil { diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index ae3a9da367..1238452ed1 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -453,9 +453,12 @@ func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAcc } func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { - // Probe state is system-managed. New accounts always start with auto probe disabled. + // Probe/session state is system-managed. New accounts always start with automatic refresh disabled. delete(accountExtra, UpstreamBillingProbeEnabledExtraKey) delete(accountExtra, UpstreamBillingProbeExtraKey) + delete(accountExtra, OllamaCloudUsageSessionExtraKey) + delete(accountExtra, OllamaCloudUsageAutoRefreshExtraKey) + delete(accountExtra, OllamaCloudUsageSnapshotExtraKey) account := &Account{ Name: input.Name, Notes: normalizeAccountNotes(input.Notes), @@ -612,6 +615,7 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U } } previousProbeIdentity := upstreamBillingProbeIdentity(account) + previousOllamaUsageIdentity := ollamaCloudUsageIdentity(account) // 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用, // 必须在此守住,否则仅在创建时的保证可被这些路径绕过。 if account.IsCredentialShadow() { @@ -675,7 +679,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U } delete(normalizedExtra, UpstreamBillingProbeEnabledExtraKey) delete(normalizedExtra, UpstreamBillingProbeExtraKey) - // 保留配额用量字段,防止编辑账号时意外重置 + delete(normalizedExtra, OllamaCloudUsageSessionExtraKey) + delete(normalizedExtra, OllamaCloudUsageAutoRefreshExtraKey) + delete(normalizedExtra, OllamaCloudUsageSnapshotExtraKey) + // 保留配额用量和专用服务受管字段,防止普通账号编辑意外覆盖。 for _, key := range []string{ "quota_used", "quota_daily_used", @@ -685,6 +692,9 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U grokBillingExtraKey, UpstreamBillingProbeEnabledExtraKey, UpstreamBillingProbeExtraKey, + OllamaCloudUsageSessionExtraKey, + OllamaCloudUsageAutoRefreshExtraKey, + OllamaCloudUsageSnapshotExtraKey, } { if v, ok := account.Extra[key]; ok { normalizedExtra[key] = v @@ -733,6 +743,17 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U delete(account.Extra, UpstreamBillingProbeEnabledExtraKey) } } + if account.Extra != nil { + if !IsOllamaCloudUsageAccount(account) { + delete(account.Extra, OllamaCloudUsageSessionExtraKey) + delete(account.Extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } else if !reflect.DeepEqual(previousOllamaUsageIdentity, ollamaCloudUsageIdentity(account)) { + delete(account.Extra, OllamaCloudUsageSessionExtraKey) + delete(account.Extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } + } // 只在指针非 nil 时更新 Concurrency(支持设置为 0) if input.Concurrency != nil { account.Concurrency = normalizeAccountConcurrency(account.Platform, account.Type, *input.Concurrency) @@ -833,6 +854,9 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // UpdateAccountExtra 仅对 Extra JSONB 做 key 级合并,避免覆盖其它运行态键 // (如 model_rate_limits / passive_usage_* 等)。 func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error { + delete(updates, OllamaCloudUsageSessionExtraKey) + delete(updates, OllamaCloudUsageAutoRefreshExtraKey) + delete(updates, OllamaCloudUsageSnapshotExtraKey) if _, exists := updates[openAILongContextBillingEnabledKey]; exists { account, err := s.accountRepo.GetByID(ctx, id) if err != nil { @@ -851,9 +875,12 @@ 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) { - // Managed probe state may only enter through the dedicated typed field below. + // Managed probe/session state may only enter through dedicated typed endpoints. delete(input.Extra, UpstreamBillingProbeEnabledExtraKey) delete(input.Extra, UpstreamBillingProbeExtraKey) + delete(input.Extra, OllamaCloudUsageSessionExtraKey) + delete(input.Extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(input.Extra, OllamaCloudUsageSnapshotExtraKey) if len(input.AccountIDs) == 0 && input.Filters != nil { accountIDs, err := s.resolveBulkUpdateTargetIDs(ctx, input.Filters) diff --git a/backend/internal/service/crs_sync_service.go b/backend/internal/service/crs_sync_service.go index 8a86c5127a..f227ea4278 100644 --- a/backend/internal/service/crs_sync_service.go +++ b/backend/internal/service/crs_sync_service.go @@ -1160,21 +1160,39 @@ func reconcileCRSUpstreamBillingProbeExtra( targetCredentials map[string]any, extra map[string]any, ) { - delete(extra, UpstreamBillingProbeEnabledExtraKey) - delete(extra, UpstreamBillingProbeExtraKey) + for _, key := range []string{ + UpstreamBillingProbeEnabledExtraKey, + UpstreamBillingProbeExtraKey, + OllamaCloudUsageSessionExtraKey, + OllamaCloudUsageAutoRefreshExtraKey, + OllamaCloudUsageSnapshotExtraKey, + } { + delete(extra, key) + } 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 + if targetPlatform == PlatformOpenAI && targetType == AccountTypeAPIKey { + if enabled, ok := existing.Extra[UpstreamBillingProbeEnabledExtraKey]; ok { + extra[UpstreamBillingProbeEnabledExtraKey] = enabled + } + if reflect.DeepEqual(upstreamBillingProbeIdentity(existing), upstreamBillingProbeIdentity(target)) { + if snapshot, ok := existing.Extra[UpstreamBillingProbeExtraKey]; ok { + extra[UpstreamBillingProbeExtraKey] = snapshot + } + } + } + if IsOllamaCloudUsageAccount(existing) && IsOllamaCloudUsageAccount(target) && + reflect.DeepEqual(ollamaCloudUsageIdentity(existing), ollamaCloudUsageIdentity(target)) { + if session, ok := existing.Extra[OllamaCloudUsageSessionExtraKey]; ok { + extra[OllamaCloudUsageSessionExtraKey] = session + } + if enabled, ok := existing.Extra[OllamaCloudUsageAutoRefreshExtraKey]; ok { + extra[OllamaCloudUsageAutoRefreshExtraKey] = enabled + } + if snapshot, ok := existing.Extra[OllamaCloudUsageSnapshotExtraKey]; ok { + extra[OllamaCloudUsageSnapshotExtraKey] = snapshot } } } diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 279ebdba04..755790f0ef 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -385,6 +385,9 @@ const ( // for probing remote Sub2API API-key billing metadata. SettingKeyUpstreamBillingProbeSettings = "upstream_billing_probe_settings" + // SettingKeyOllamaCloudUsageSettings stores the opt-in global runner switch and interval. + SettingKeyOllamaCloudUsageSettings = "ollama_cloud_usage_settings" + // ========================= // Overload Cooldown (529) // ========================= diff --git a/backend/internal/service/ollama_cloud_usage.go b/backend/internal/service/ollama_cloud_usage.go new file mode 100644 index 0000000000..2b50bab62c --- /dev/null +++ b/backend/internal/service/ollama_cloud_usage.go @@ -0,0 +1,1033 @@ +package service + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "maps" + "math/rand/v2" + "net/http" + "net/url" + "strconv" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/google/uuid" + "golang.org/x/net/http/httpguts" + "golang.org/x/sync/errgroup" + "golang.org/x/sync/singleflight" +) + +const ( + OllamaCloudUsageSessionExtraKey = "ollama_cloud_usage_session" + OllamaCloudUsageAutoRefreshExtraKey = "ollama_cloud_usage_auto_refresh" + OllamaCloudUsageSnapshotExtraKey = "ollama_cloud_usage_snapshot" + + ollamaCloudUsageSettingsURL = "https://ollama.com/settings" + ollamaCloudUsageDefaultIntervalMinutes = 60 + ollamaCloudUsageMinIntervalMinutes = 15 + ollamaCloudUsageMaxIntervalMinutes = 24 * 60 + ollamaCloudUsageCycleInterval = time.Minute + ollamaCloudUsageManualRefreshInterval = 30 * time.Second + ollamaCloudUsageRequestTimeout = 15 * time.Second + ollamaCloudUsageMaxBodyBytes = 512 * 1024 + ollamaCloudUsageMaxSessionBytes = 16 * 1024 + ollamaCloudUsageMaxPerCycle = 20 + ollamaCloudUsageConcurrency = 4 + ollamaCloudUsageMaxDelay = 24 * time.Hour + ollamaCloudUsageLeaderLockKey = "ollama:cloud:usage:leader" + ollamaCloudUsageLeaderLockTTL = 2 * time.Minute +) + +var ( + ErrOllamaCloudUsageUnavailable = infraerrors.ServiceUnavailable( + "OLLAMA_CLOUD_USAGE_UNAVAILABLE", "Ollama Cloud usage is unavailable", + ) + ErrOllamaCloudUsageAccountInvalid = infraerrors.BadRequest( + "OLLAMA_CLOUD_USAGE_ACCOUNT_INVALID", "account must be an OpenAI or Anthropic API key account using https://ollama.com", + ) + ErrOllamaCloudUsageSessionRequired = infraerrors.BadRequest( + "OLLAMA_CLOUD_USAGE_SESSION_REQUIRED", "an Ollama web session must be configured first", + ) + ErrOllamaCloudUsageEncryptionKey = infraerrors.BadRequest( + "OLLAMA_CLOUD_USAGE_ENCRYPTION_KEY_NOT_CONFIGURED", "cannot store an Ollama web session without a fixed TOTP_ENCRYPTION_KEY", + ) + ErrOllamaCloudUsageIdentityChanged = infraerrors.Conflict( + "OLLAMA_CLOUD_USAGE_IDENTITY_CHANGED", "account identity or Ollama web session changed during refresh; retry", + ) + ErrOllamaCloudUsageRefreshRateLimited = infraerrors.TooManyRequests( + "OLLAMA_CLOUD_USAGE_REFRESH_RATE_LIMITED", "Ollama Cloud usage can be refreshed manually once every 30 seconds", + ) + errOllamaCloudUsageUnauthorizedHTML = errors.New("settings HTML is a sign-in page") +) + +const ( + OllamaCloudUsageStatusOK = "ok" + OllamaCloudUsageStatusUnauthorized = "unauthorized" + OllamaCloudUsageStatusFailed = "failed" +) + +// OllamaCloudUsageSettings controls the opt-in periodic refresh runner. +type OllamaCloudUsageSettings struct { + Enabled bool `json:"enabled"` + IntervalMinutes int `json:"interval_minutes"` +} + +// OllamaCloudUsageWindow is a narrow, sanitized view of one official usage window. +type OllamaCloudUsageWindow struct { + UsedPercent float64 `json:"used_percent"` + ResetAt *time.Time `json:"reset_at,omitempty"` + ResetText string `json:"reset_text,omitempty"` +} + +// OllamaCloudUsageModelWindow identifies the official window for a model count. +type OllamaCloudUsageModelWindow string + +const ( + OllamaCloudUsageModelWindowFiveHour OllamaCloudUsageModelWindow = "five_hour" + OllamaCloudUsageModelWindowSevenDay OllamaCloudUsageModelWindow = "seven_day" +) + +// OllamaCloudUsageModel is the window-scoped model/request pair exposed by Ollama's usage DOM. +type OllamaCloudUsageModel struct { + Model string `json:"model"` + Window OllamaCloudUsageModelWindow `json:"window"` + Requests int64 `json:"requests"` +} + +// OllamaCloudUsageData intentionally excludes raw HTML and browser-session data. +type OllamaCloudUsageData struct { + Plan string `json:"plan,omitempty"` + FiveHour *OllamaCloudUsageWindow `json:"five_hour,omitempty"` + SevenDay *OllamaCloudUsageWindow `json:"seven_day,omitempty"` + Balance string `json:"balance,omitempty"` + Models []OllamaCloudUsageModel `json:"models,omitempty"` +} + +// OllamaCloudUsageSnapshot is the only usage observation persisted in account extra. +type OllamaCloudUsageSnapshot struct { + Status string `json:"status"` + Data *OllamaCloudUsageData `json:"data,omitempty"` + FetchedAt *time.Time `json:"fetched_at,omitempty"` + LastAttemptAt time.Time `json:"last_attempt_at"` + NextRefreshAt time.Time `json:"next_refresh_at"` + FailureCount int `json:"failure_count,omitempty"` + HTTPStatus int `json:"http_status,omitempty"` + LastError string `json:"last_error,omitempty"` +} + +// OllamaCloudUsageState is the dedicated DTO exposed to administrators. +type OllamaCloudUsageState struct { + AccountID int64 `json:"account_id"` + Eligible bool `json:"eligible"` + Configured bool `json:"configured"` + AutoRefreshEnabled bool `json:"auto_refresh_enabled"` + EncryptionKeyConfigured bool `json:"encryption_key_configured"` + Snapshot *OllamaCloudUsageSnapshot `json:"snapshot,omitempty"` +} + +type ollamaCloudUsageRepository interface { + ListOllamaCloudUsageGroupAccounts(context.Context, []*Account) ([]Account, error) + SaveOllamaCloudUsageSession(context.Context, *Account, string, bool) error + DeleteOllamaCloudUsageSession(context.Context, *Account) error + SetOllamaCloudUsageAutoRefresh(context.Context, *Account, bool) error + UpdateOllamaCloudUsageSnapshot(context.Context, *Account, *OllamaCloudUsageSnapshot) error + DisableOllamaCloudUsageAutoRefresh(context.Context, *Account) error + ListDueOllamaCloudUsageAccounts(context.Context, time.Time, int) ([]Account, error) +} + +// GetOllamaCloudUsageSettings returns fail-safe defaults when the setting is absent. +func (s *SettingService) GetOllamaCloudUsageSettings(ctx context.Context) (*OllamaCloudUsageSettings, error) { + defaults := defaultOllamaCloudUsageSettings() + if s == nil || s.settingRepo == nil { + return defaults, nil + } + raw, err := s.settingRepo.GetValue(ctx, SettingKeyOllamaCloudUsageSettings) + if err != nil { + if errors.Is(err, ErrSettingNotFound) { + return defaults, nil + } + return nil, fmt.Errorf("get Ollama Cloud usage settings: %w", err) + } + if strings.TrimSpace(raw) == "" { + return defaults, nil + } + settings := *defaults + if err := json.Unmarshal([]byte(raw), &settings); err != nil { + return nil, fmt.Errorf("parse Ollama Cloud usage settings: %w", err) + } + if settings.IntervalMinutes == 0 { + settings.IntervalMinutes = defaults.IntervalMinutes + } + normalizeOllamaCloudUsageSettings(&settings) + return &settings, nil +} + +func (s *SettingService) SetOllamaCloudUsageSettings(ctx context.Context, settings *OllamaCloudUsageSettings) error { + if s == nil || s.settingRepo == nil { + return ErrOllamaCloudUsageUnavailable + } + if settings == nil { + return infraerrors.BadRequest("INVALID_OLLAMA_CLOUD_USAGE_SETTINGS", "settings cannot be nil") + } + if settings.IntervalMinutes < ollamaCloudUsageMinIntervalMinutes || settings.IntervalMinutes > ollamaCloudUsageMaxIntervalMinutes { + return infraerrors.BadRequest( + "INVALID_OLLAMA_CLOUD_USAGE_INTERVAL", + fmt.Sprintf("interval_minutes must be between %d and %d", ollamaCloudUsageMinIntervalMinutes, ollamaCloudUsageMaxIntervalMinutes), + ) + } + normalizeOllamaCloudUsageSettings(settings) + data, err := json.Marshal(settings) + if err != nil { + return fmt.Errorf("marshal Ollama Cloud usage settings: %w", err) + } + return s.settingRepo.Set(ctx, SettingKeyOllamaCloudUsageSettings, string(data)) +} + +func defaultOllamaCloudUsageSettings() *OllamaCloudUsageSettings { + return &OllamaCloudUsageSettings{Enabled: false, IntervalMinutes: ollamaCloudUsageDefaultIntervalMinutes} +} + +func normalizeOllamaCloudUsageSettings(settings *OllamaCloudUsageSettings) { + if settings.IntervalMinutes < ollamaCloudUsageMinIntervalMinutes { + settings.IntervalMinutes = ollamaCloudUsageMinIntervalMinutes + } + if settings.IntervalMinutes > ollamaCloudUsageMaxIntervalMinutes { + settings.IntervalMinutes = ollamaCloudUsageMaxIntervalMinutes + } +} + +// OllamaCloudUsageService refreshes the official settings HTML without affecting routing state. +type OllamaCloudUsageService struct { + accountRepo AccountRepository + httpUpstream HTTPUpstream + settingService *SettingService + encryptor SecretEncryptor + encryptionKeyConfigured bool + + parentCtx context.Context + parentCancel context.CancelFunc + wg sync.WaitGroup + mu sync.Mutex + started bool + stopped bool + cycleMu sync.Mutex + refreshGroup singleflight.Group + refreshSlots chan struct{} + now func() time.Time + lockCache LeaderLockCache + db *sql.DB + instanceID string +} + +func NewOllamaCloudUsageService( + accountRepo AccountRepository, + httpUpstream HTTPUpstream, + settingService *SettingService, + encryptor SecretEncryptor, + encryptionKeyConfigured bool, +) *OllamaCloudUsageService { + ctx, cancel := context.WithCancel(context.Background()) + return &OllamaCloudUsageService{ + accountRepo: accountRepo, + httpUpstream: httpUpstream, + settingService: settingService, + encryptor: encryptor, + encryptionKeyConfigured: encryptionKeyConfigured, + parentCtx: ctx, + parentCancel: cancel, + refreshSlots: make(chan struct{}, ollamaCloudUsageConcurrency), + now: time.Now, + instanceID: uuid.NewString(), + } +} + +func ProvideOllamaCloudUsageService( + accountRepo AccountRepository, + httpUpstream HTTPUpstream, + settingService *SettingService, + encryptor SecretEncryptor, + cfg *config.Config, + lockCache LeaderLockCache, + db *sql.DB, +) *OllamaCloudUsageService { + keyConfigured := cfg != nil && cfg.Totp.EncryptionKeyConfigured + svc := NewOllamaCloudUsageService(accountRepo, httpUpstream, settingService, encryptor, keyConfigured) + svc.lockCache = lockCache + svc.db = db + svc.Start() + return svc +} + +func (s *OllamaCloudUsageService) 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 *OllamaCloudUsageService) 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 *OllamaCloudUsageService) runLoop() { + defer s.wg.Done() + _ = s.RunDue(s.parentCtx) + ticker := time.NewTicker(ollamaCloudUsageCycleInterval) + defer ticker.Stop() + for { + select { + case <-s.parentCtx.Done(): + return + case <-ticker.C: + if err := s.RunDue(s.parentCtx); err != nil { + logger.LegacyPrintf("service.ollama_cloud_usage", "run_due_failed: err=%v", err) + } + } + } +} + +func (s *OllamaCloudUsageService) GetSettings(ctx context.Context) (*OllamaCloudUsageSettings, error) { + if s == nil || s.settingService == nil { + return defaultOllamaCloudUsageSettings(), nil + } + return s.settingService.GetOllamaCloudUsageSettings(ctx) +} + +func (s *OllamaCloudUsageService) UpdateSettings(ctx context.Context, settings *OllamaCloudUsageSettings) error { + if s == nil || s.settingService == nil { + return ErrOllamaCloudUsageUnavailable + } + return s.settingService.SetOllamaCloudUsageSettings(ctx, settings) +} + +func (s *OllamaCloudUsageService) GetState(ctx context.Context, accountID int64) (*OllamaCloudUsageState, error) { + if s == nil || s.accountRepo == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, err + } + if err := s.ResolveAccounts(ctx, []*Account{account}); err != nil { + return nil, err + } + state := OllamaCloudUsageStateFromAccount(account) + s.EnrichState(state) + return state, nil +} + +// ResolveAccounts overlays group-owned managed state onto the supplied account +// objects. The repository resolves all matching siblings in one bounded query, +// so account-list responses do not issue one query per row. +func (s *OllamaCloudUsageService) ResolveAccounts(ctx context.Context, accounts []*Account) error { + if s == nil || s.accountRepo == nil || len(accounts) == 0 { + return nil + } + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return nil + } + eligible := make([]*Account, 0, len(accounts)) + for _, account := range accounts { + if _, ok := ollamaCloudUsageGroupFingerprint(account); ok { + eligible = append(eligible, account) + } + } + if len(eligible) == 0 { + return nil + } + siblings, err := writer.ListOllamaCloudUsageGroupAccounts(ctx, eligible) + if err != nil { + return fmt.Errorf("resolve Ollama Cloud usage groups: %w", err) + } + sources := make(map[string]*Account) + for index := range siblings { + candidate := &siblings[index] + fingerprint, valid := ollamaCloudUsageGroupFingerprint(candidate) + if !valid || !ollamaCloudUsageConfigured(candidate) { + continue + } + current := sources[fingerprint] + if current == nil || candidate.UpdatedAt.After(current.UpdatedAt) || + (candidate.UpdatedAt.Equal(current.UpdatedAt) && candidate.ID < current.ID) { + sources[fingerprint] = candidate + } + } + resolvedSources := make(map[string]*Account, len(sources)) + for fingerprint, source := range sources { + clone := *source + clone.Extra = make(map[string]any, len(source.Extra)) + maps.Copy(clone.Extra, source.Extra) + resolvedSources[fingerprint] = &clone + } + for index := range siblings { + candidate := &siblings[index] + fingerprint, valid := ollamaCloudUsageGroupFingerprint(candidate) + source := resolvedSources[fingerprint] + if !valid || source == nil || !sameOllamaCloudUsageSession(source, candidate) { + continue + } + candidateSnapshot := decodeOllamaCloudUsageSnapshot(candidate.Extra) + currentSnapshot := decodeOllamaCloudUsageSnapshot(source.Extra) + if candidateSnapshot != nil && (currentSnapshot == nil || candidateSnapshot.LastAttemptAt.After(currentSnapshot.LastAttemptAt)) { + source.Extra[OllamaCloudUsageSnapshotExtraKey] = candidate.Extra[OllamaCloudUsageSnapshotExtraKey] + } + } + for _, account := range eligible { + fingerprint, _ := ollamaCloudUsageGroupFingerprint(account) + applyOllamaCloudUsageManagedExtra(account, resolvedSources[fingerprint]) + } + return nil +} + +func sameOllamaCloudUsageSession(left, right *Account) bool { + if left == nil || right == nil || left.Extra == nil || right.Extra == nil { + return false + } + leftSession, leftOK := left.Extra[OllamaCloudUsageSessionExtraKey].(string) + rightSession, rightOK := right.Extra[OllamaCloudUsageSessionExtraKey].(string) + return leftOK && rightOK && leftSession != "" && leftSession == rightSession +} + +func applyOllamaCloudUsageManagedExtra(target, source *Account) { + if target == nil { + return + } + if target.Extra == nil { + target.Extra = make(map[string]any) + } + for _, key := range []string{ + OllamaCloudUsageSessionExtraKey, + OllamaCloudUsageAutoRefreshExtraKey, + OllamaCloudUsageSnapshotExtraKey, + } { + delete(target.Extra, key) + if source != nil && source.Extra != nil { + if value, ok := source.Extra[key]; ok { + target.Extra[key] = value + } + } + } +} + +func (s *OllamaCloudUsageService) SaveSession(ctx context.Context, accountID int64, session string) (*OllamaCloudUsageState, error) { + if s == nil || s.accountRepo == nil || s.encryptor == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + if !s.encryptionKeyConfigured { + return nil, ErrOllamaCloudUsageEncryptionKey + } + normalized, err := normalizeOllamaCloudUsageCookie(session) + if err != nil { + return nil, infraerrors.BadRequest("INVALID_OLLAMA_CLOUD_USAGE_SESSION", err.Error()) + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, err + } + if !IsOllamaCloudUsageAccount(account) { + return nil, ErrOllamaCloudUsageAccountInvalid + } + if err := s.ResolveAccounts(ctx, []*Account{account}); err != nil { + return nil, err + } + ciphertext, err := s.encryptor.Encrypt(normalized) + if err != nil { + return nil, fmt.Errorf("encrypt Ollama web session: %w", err) + } + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return nil, ErrOllamaCloudUsageUnavailable + } + preserveAutoRefresh := ollamaCloudUsageConfigured(account) && ollamaCloudUsageAutoRefreshEnabled(account) + if err := writer.SaveOllamaCloudUsageSession(ctx, account, ciphertext, preserveAutoRefresh); err != nil { + return nil, err + } + return s.GetState(ctx, accountID) +} + +func (s *OllamaCloudUsageService) DeleteSession(ctx context.Context, accountID int64) (*OllamaCloudUsageState, error) { + if s == nil || s.accountRepo == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, err + } + if !IsOllamaCloudUsageAccount(account) { + return nil, ErrOllamaCloudUsageAccountInvalid + } + if err := s.ResolveAccounts(ctx, []*Account{account}); err != nil { + return nil, err + } + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return nil, ErrOllamaCloudUsageUnavailable + } + if err := writer.DeleteOllamaCloudUsageSession(ctx, account); err != nil { + return nil, err + } + return s.GetState(ctx, accountID) +} + +func (s *OllamaCloudUsageService) SetAutoRefresh(ctx context.Context, accountID int64, enabled bool) (*OllamaCloudUsageState, error) { + if s == nil || s.accountRepo == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, err + } + if !IsOllamaCloudUsageAccount(account) { + return nil, ErrOllamaCloudUsageAccountInvalid + } + if err := s.ResolveAccounts(ctx, []*Account{account}); err != nil { + return nil, err + } + if enabled && !ollamaCloudUsageConfigured(account) { + return nil, ErrOllamaCloudUsageSessionRequired + } + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return nil, ErrOllamaCloudUsageUnavailable + } + if err := writer.SetOllamaCloudUsageAutoRefresh(ctx, account, enabled); err != nil { + return nil, err + } + return s.GetState(ctx, accountID) +} + +func (s *OllamaCloudUsageService) Refresh(ctx context.Context, accountID int64) (*OllamaCloudUsageState, error) { + settings, err := s.GetSettings(ctx) + if err != nil { + return nil, err + } + if _, err := s.refreshAccount(ctx, accountID, settings.IntervalMinutes, false); err != nil { + return nil, err + } + return s.GetState(ctx, accountID) +} + +func (s *OllamaCloudUsageService) 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 + } + release, acquired := tryAcquireSingletonLeaderLock(ctx, s.lockCache, s.db, ollamaCloudUsageLeaderLockKey, s.instanceID, ollamaCloudUsageLeaderLockTTL) + if !acquired { + return nil + } + defer release() + + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return ErrOllamaCloudUsageUnavailable + } + now := s.currentTime() + accounts, err := writer.ListDueOllamaCloudUsageAccounts(ctx, now, ollamaCloudUsageMaxPerCycle) + if err != nil { + return fmt.Errorf("list due Ollama Cloud usage accounts: %w", err) + } + var group errgroup.Group + seenGroups := make(map[string]struct{}, len(accounts)) + for index := range accounts { + account := accounts[index] + fingerprint, valid := ollamaCloudUsageGroupFingerprint(&account) + if !valid || !account.IsActive() || !ollamaCloudUsageConfigured(&account) || !ollamaCloudUsageAutoRefreshEnabled(&account) { + continue + } + if _, duplicate := seenGroups[fingerprint]; duplicate { + continue + } + seenGroups[fingerprint] = struct{}{} + if snapshot := decodeOllamaCloudUsageSnapshot(account.Extra); snapshot != nil && now.Before(snapshot.NextRefreshAt) { + continue + } + accountID := account.ID + expected := account + group.Go(func() error { + if _, refreshErr := s.refreshAccount(ctx, accountID, settings.IntervalMinutes, true); refreshErr != nil { + if errors.Is(refreshErr, ErrOllamaCloudUsageIdentityChanged) { + if disableErr := writer.DisableOllamaCloudUsageAutoRefresh(ctx, &expected); disableErr != nil { + logger.LegacyPrintf("service.ollama_cloud_usage", "disable_auto_refresh_failed: account_id=%d err=%v", accountID, disableErr) + } + return nil + } + logger.LegacyPrintf("service.ollama_cloud_usage", "refresh_due_failed: account_id=%d err=%v", accountID, refreshErr) + } + return nil + }) + } + return group.Wait() +} + +func (s *OllamaCloudUsageService) refreshAccount(ctx context.Context, accountID int64, intervalMinutes int, requireEnabled bool) (*OllamaCloudUsageSnapshot, error) { + if s == nil || s.accountRepo == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + anchor, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, err + } + key, valid := ollamaCloudUsageGroupFingerprint(anchor) + if !valid { + return nil, ErrOllamaCloudUsageAccountInvalid + } + value, err, _ := s.refreshGroup.Do(key, func() (any, error) { + select { + case s.refreshSlots <- struct{}{}: + defer func() { <-s.refreshSlots }() + case <-ctx.Done(): + return nil, ctx.Err() + } + account, loadErr := s.accountRepo.GetByID(ctx, accountID) + if loadErr != nil { + return nil, loadErr + } + currentKey, currentValid := ollamaCloudUsageGroupFingerprint(account) + if !currentValid { + return nil, ErrOllamaCloudUsageAccountInvalid + } + if currentKey != key { + return nil, ErrOllamaCloudUsageIdentityChanged + } + if err := s.ResolveAccounts(ctx, []*Account{account}); err != nil { + return nil, err + } + if !ollamaCloudUsageConfigured(account) { + return nil, ErrOllamaCloudUsageSessionRequired + } + if !requireEnabled { + if snapshot := decodeOllamaCloudUsageSnapshot(account.Extra); snapshot != nil && !snapshot.LastAttemptAt.IsZero() { + retryAt := snapshot.LastAttemptAt.Add(ollamaCloudUsageManualRefreshInterval) + if now := s.currentTime(); now.Before(retryAt) { + remaining := retryAt.Sub(now) + seconds := int((remaining + time.Second - 1) / time.Second) + return nil, ErrOllamaCloudUsageRefreshRateLimited.WithMetadata(map[string]string{ + "retry_after_seconds": strconv.Itoa(seconds), + }) + } + } + } + if requireEnabled { + if !account.IsActive() || !ollamaCloudUsageAutoRefreshEnabled(account) { + return nil, nil + } + if snapshot := decodeOllamaCloudUsageSnapshot(account.Extra); snapshot != nil && s.currentTime().Before(snapshot.NextRefreshAt) { + return nil, nil + } + } + return s.refreshLoadedAccount(ctx, account, intervalMinutes) + }) + if err != nil || value == nil { + return nil, err + } + snapshot, ok := value.(*OllamaCloudUsageSnapshot) + if !ok { + return nil, fmt.Errorf("invalid Ollama Cloud usage refresh result") + } + return snapshot, nil +} + +func (s *OllamaCloudUsageService) refreshLoadedAccount(ctx context.Context, account *Account, intervalMinutes int) (*OllamaCloudUsageSnapshot, error) { + now := s.currentTime().UTC() + ciphertext, _ := account.Extra[OllamaCloudUsageSessionExtraKey].(string) + if ciphertext == "" { + return nil, ErrOllamaCloudUsageSessionRequired + } + if !s.encryptionKeyConfigured || s.encryptor == nil { + return nil, ErrOllamaCloudUsageEncryptionKey + } + cookie, err := s.encryptor.Decrypt(ciphertext) + if err != nil { + return nil, infraerrors.ServiceUnavailable("OLLAMA_CLOUD_USAGE_SESSION_DECRYPT_FAILED", "stored Ollama web session cannot be decrypted") + } + cookie, err = normalizeOllamaCloudUsageCookie(cookie) + if err != nil { + return nil, infraerrors.ServiceUnavailable("OLLAMA_CLOUD_USAGE_SESSION_INVALID", "stored Ollama web session is invalid") + } + if s.httpUpstream == nil { + return nil, ErrOllamaCloudUsageUnavailable + } + proxyURL := "" + if account.ProxyID != nil { + if account.Proxy == nil || account.Proxy.ID != *account.ProxyID { + return nil, ErrOllamaCloudUsageIdentityChanged + } + proxyURL = account.Proxy.URL() + } + requestCtx, cancel := context.WithTimeout(WithHTTPUpstreamRedirectsDisabled(ctx), ollamaCloudUsageRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(requestCtx, http.MethodGet, ollamaCloudUsageSettingsURL, nil) + if err != nil || !isExactOllamaCloudSettingsURL(req.URL) { + return nil, ErrOllamaCloudUsageUnavailable + } + req.Header.Set("Accept", "text/html,application/xhtml+xml") + req.Header.Set("Cookie", cookie) + req.Header.Set("User-Agent", "sub2api-ollama-usage/1") + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return s.persistFailure(ctx, account, intervalMinutes, now, 0, "request_failed", 0, false) + } + if resp == nil || resp.Body == nil { + return s.persistFailure(ctx, account, intervalMinutes, now, 0, "empty_response", 0, false) + } + defer func() { _ = resp.Body.Close() }() + if resp.Request != nil && !isExactOllamaCloudSettingsURL(resp.Request.URL) { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "response_host_mismatch", 0, false) + } + if resp.StatusCode >= 300 && resp.StatusCode < 400 { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "redirect_blocked", retryAfter(resp.Header, now), false) + } + if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "unauthorized", retryAfter(resp.Header, now), true) + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "http_error", retryAfter(resp.Header, now), false) + } + body, readErr := io.ReadAll(io.LimitReader(resp.Body, ollamaCloudUsageMaxBodyBytes+1)) + if readErr != nil { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "response_read_failed", 0, false) + } + if len(body) > ollamaCloudUsageMaxBodyBytes { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "response_too_large", 0, false) + } + data, parseErr := parseOllamaCloudUsageHTML(body) + if errors.Is(parseErr, errOllamaCloudUsageUnauthorizedHTML) { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "unauthorized", 0, true) + } + if parseErr != nil { + return s.persistFailure(ctx, account, intervalMinutes, now, resp.StatusCode, "invalid_html", 0, false) + } + snapshot := &OllamaCloudUsageSnapshot{ + Status: OllamaCloudUsageStatusOK, + Data: data, + FetchedAt: &now, + LastAttemptAt: now, + NextRefreshAt: now.Add(nextOllamaCloudUsageDelay(intervalMinutes, 0, 0)), + HTTPStatus: resp.StatusCode, + } + if err := s.updateSnapshot(ctx, account, snapshot); err != nil { + return nil, err + } + return snapshot, nil +} + +func (s *OllamaCloudUsageService) persistFailure( + ctx context.Context, + account *Account, + intervalMinutes int, + now time.Time, + httpStatus int, + reason string, + retryAfterDuration time.Duration, + unauthorized bool, +) (*OllamaCloudUsageSnapshot, error) { + previous := decodeOllamaCloudUsageSnapshot(account.Extra) + failureCount := 1 + if previous != nil { + failureCount = previous.FailureCount + 1 + } + status := OllamaCloudUsageStatusFailed + if unauthorized { + status = OllamaCloudUsageStatusUnauthorized + } + snapshot := &OllamaCloudUsageSnapshot{ + Status: status, + LastAttemptAt: now, + NextRefreshAt: now.Add(nextOllamaCloudUsageDelay(intervalMinutes, failureCount, retryAfterDuration)), + FailureCount: failureCount, + HTTPStatus: httpStatus, + LastError: reason, + } + if previous != nil { + snapshot.Data = previous.Data + snapshot.FetchedAt = previous.FetchedAt + } + if err := s.updateSnapshot(ctx, account, snapshot); err != nil { + return nil, err + } + return snapshot, nil +} + +func (s *OllamaCloudUsageService) updateSnapshot(ctx context.Context, account *Account, snapshot *OllamaCloudUsageSnapshot) error { + writer, ok := s.accountRepo.(ollamaCloudUsageRepository) + if !ok { + return ErrOllamaCloudUsageUnavailable + } + return writer.UpdateOllamaCloudUsageSnapshot(ctx, account, snapshot) +} + +// EnrichState adds service-owned runtime configuration to an account-derived state. +func (s *OllamaCloudUsageService) EnrichState(state *OllamaCloudUsageState) { + if state == nil { + return + } + state.EncryptionKeyConfigured = s != nil && s.encryptionKeyConfigured +} + +func OllamaCloudUsageStateFromAccount(account *Account) *OllamaCloudUsageState { + state := &OllamaCloudUsageState{} + if account == nil { + return state + } + state.AccountID = account.ID + state.Eligible = IsOllamaCloudUsageAccount(account) + if !state.Eligible { + return state + } + state.Configured = ollamaCloudUsageConfigured(account) + state.AutoRefreshEnabled = state.Configured && ollamaCloudUsageAutoRefreshEnabled(account) + state.Snapshot = decodeOllamaCloudUsageSnapshot(account.Extra) + return state +} + +func IsOllamaCloudUsageAccount(account *Account) bool { + if account == nil || account.Type != AccountTypeAPIKey || (account.Platform != PlatformOpenAI && account.Platform != PlatformAnthropic) { + return false + } + baseURL, _ := account.Credentials["base_url"].(string) + return isOllamaCloudBaseURL(baseURL) +} + +func isOllamaCloudBaseURL(raw string) bool { + raw = strings.TrimSpace(raw) + if raw == "" || strings.ContainsAny(raw, "?#") { + return false + } + parsed, err := url.Parse(raw) + if err != nil || parsed.Opaque != "" || !strings.EqualFold(parsed.Scheme, "https") || parsed.User != nil || parsed.ForceQuery || parsed.RawQuery != "" || parsed.Fragment != "" || parsed.RawFragment != "" { + return false + } + hostname := strings.ToLower(parsed.Hostname()) + if hostname != "ollama.com" && hostname != "www.ollama.com" { + return false + } + authority := strings.ToLower(parsed.Host) + if authority != hostname && authority != hostname+":443" { + return false + } + if parsed.RawPath != "" { + return false + } + return parsed.Path == "" || parsed.Path == "/v1" +} + +func ollamaCloudUsageIdentity(account *Account) map[string]any { + if !IsOllamaCloudUsageAccount(account) { + return nil + } + apiKey, ok := account.Credentials["api_key"].(string) + if !ok || apiKey == "" { + return nil + } + return map[string]any{"host": "ollama.com", "api_key": apiKey} +} + +func ollamaCloudUsageGroupFingerprint(account *Account) (string, bool) { + identity := ollamaCloudUsageIdentity(account) + if identity == nil { + return "", false + } + apiKey, _ := identity["api_key"].(string) + sum := sha256.Sum256([]byte("ollama.com\x00" + apiKey)) + return hex.EncodeToString(sum[:]), true +} + +func isExactOllamaCloudSettingsURL(parsed *url.URL) bool { + return parsed != nil && parsed.Scheme == "https" && parsed.Host == "ollama.com" && parsed.Path == "/settings" && + parsed.User == nil && parsed.RawQuery == "" && parsed.Fragment == "" && parsed.RawPath == "" +} + +func normalizeOllamaCloudUsageCookie(raw string) (string, error) { + if len(raw) > ollamaCloudUsageMaxSessionBytes { + return "", errors.New("session is too large") + } + raw = strings.TrimSpace(raw) + if strings.ContainsAny(raw, "\r\n") { + return "", errors.New("session contains invalid header characters") + } + if raw == "" { + return "", errors.New("session cannot be empty") + } + if !httpguts.ValidHeaderFieldValue(raw) { + return "", errors.New("session contains invalid header characters") + } + blockedAttributes := map[string]struct{}{ + "domain": {}, "path": {}, "expires": {}, "max-age": {}, "samesite": {}, "secure": {}, "httponly": {}, "partitioned": {}, + } + parts := strings.Split(raw, ";") + normalized := make([]string, 0, len(parts)) + seen := make(map[string]struct{}, len(parts)) + for _, part := range parts { + part = strings.TrimSpace(part) + name, value, ok := strings.Cut(part, "=") + name = strings.TrimSpace(name) + value = strings.TrimSpace(value) + if !ok || name == "" || value == "" || !httpguts.ValidHeaderFieldName(name) || strings.HasPrefix(name, "$") { + return "", errors.New("session must be a Cookie header containing name=value pairs") + } + lowerName := strings.ToLower(name) + if _, blocked := blockedAttributes[lowerName]; blocked { + return "", errors.New("paste a Cookie header, not a Set-Cookie value with attributes") + } + if _, duplicate := seen[lowerName]; duplicate { + return "", errors.New("session contains duplicate cookie names") + } + if strings.ContainsAny(value, ";\r\n") { + return "", errors.New("session contains an invalid cookie value") + } + seen[lowerName] = struct{}{} + if isAllowedOllamaCloudSessionCookie(name) { + normalized = append(normalized, name+"="+value) + } + } + if len(normalized) == 0 { + return "", errors.New("session does not contain an allowed Ollama session cookie") + } + return strings.Join(normalized, "; "), nil +} + +func isAllowedOllamaCloudSessionCookie(name string) bool { + switch name { + case "wos-session", "__Secure-session", "session", "ollama_session", "__Host-ollama_session": + return true + } + for _, base := range []string{ + "next-auth.session-token", + "__Secure-next-auth.session-token", + "authjs.session-token", + "__Secure-authjs.session-token", + } { + if name == base { + return true + } + if suffix, ok := strings.CutPrefix(name, base+"."); ok && suffix != "" { + validShard := true + for _, char := range suffix { + if char < '0' || char > '9' { + validShard = false + break + } + } + if validShard { + return true + } + } + } + return false +} + +func ollamaCloudUsageConfigured(account *Account) bool { + if account == nil || account.Extra == nil { + return false + } + value, ok := account.Extra[OllamaCloudUsageSessionExtraKey].(string) + return ok && strings.TrimSpace(value) != "" +} + +func ollamaCloudUsageAutoRefreshEnabled(account *Account) bool { + if account == nil || account.Extra == nil { + return false + } + enabled, ok := account.Extra[OllamaCloudUsageAutoRefreshExtraKey].(bool) + return ok && enabled +} + +func decodeOllamaCloudUsageSnapshot(extra map[string]any) *OllamaCloudUsageSnapshot { + if extra == nil { + return nil + } + value, ok := extra[OllamaCloudUsageSnapshotExtraKey] + if !ok || value == nil { + return nil + } + raw, err := json.Marshal(value) + if err != nil { + return nil + } + var snapshot OllamaCloudUsageSnapshot + if err := json.Unmarshal(raw, &snapshot); err != nil { + return nil + } + if snapshot.Status != OllamaCloudUsageStatusOK && snapshot.Status != OllamaCloudUsageStatusUnauthorized && snapshot.Status != OllamaCloudUsageStatusFailed { + return nil + } + return &snapshot +} + +func nextOllamaCloudUsageDelay(intervalMinutes, failureCount int, retryAfterDuration time.Duration) time.Duration { + minimumDelay := retryAfterDuration + base := time.Duration(intervalMinutes) * time.Minute + if base < ollamaCloudUsageMinIntervalMinutes*time.Minute { + base = ollamaCloudUsageMinIntervalMinutes * time.Minute + } + if failureCount > 0 { + shift := min(failureCount-1, 6) + base *= time.Duration(1 << shift) + } + if base > ollamaCloudUsageMaxDelay { + base = ollamaCloudUsageMaxDelay + } + if retryAfterDuration > base { + base = retryAfterDuration + } + jitterRange := base / 10 + if jitterRange > 5*time.Minute { + jitterRange = 5 * time.Minute + } + if jitterRange > 0 { + base += time.Duration(rand.Int64N(int64(jitterRange)*2+1)) - jitterRange + } + if base < minimumDelay { + return minimumDelay + } + if base < time.Minute { + return time.Minute + } + return base +} + +func (s *OllamaCloudUsageService) currentTime() time.Time { + if s != nil && s.now != nil { + return s.now() + } + return time.Now() +} diff --git a/backend/internal/service/ollama_cloud_usage_parser.go b/backend/internal/service/ollama_cloud_usage_parser.go new file mode 100644 index 0000000000..beb949e2e6 --- /dev/null +++ b/backend/internal/service/ollama_cloud_usage_parser.go @@ -0,0 +1,456 @@ +package service + +import ( + "fmt" + "regexp" + "sort" + "strconv" + "strings" + "time" + + "golang.org/x/net/html" +) + +var ( + ollamaUsagePercentPattern = regexp.MustCompile(`(?i)([0-9]+(?:\.[0-9]+)?)\s*%`) + ollamaUsageWidthPattern = regexp.MustCompile(`(?i)(?:^|;)\s*width\s*:\s*([0-9]+(?:\.[0-9]+)?)%`) + ollamaBalancePattern = regexp.MustCompile(`(?i)(?:balance|credits?)(?:\s+[[:alpha:]]+){0,4}\s*[:\n]?\s*((?:USD\s*)?\$?\s*-?[0-9][0-9,]*(?:\.[0-9]{1,4})?)`) + ollamaResetPattern = regexp.MustCompile(`(?i)\breset(?:s|ting)?\s*(?:at|in|on)?\s*[:\-]?\s*([^\n|]+)`) + ollamaModelFallbackPattern = regexp.MustCompile(`(?i)^(.+?)\s+([0-9][0-9,]*)\s+requests?$`) + + ollamaFiveHourUsageAliases = []string{"session usage", "5 hour usage", "5-hour usage", "5h usage", "5 hour limit", "5-hour limit"} + ollamaSevenDayUsageAliases = []string{"weekly usage", "7 day usage", "7-day usage", "7d usage", "weekly limit", "7 day limit"} +) + +func parseOllamaCloudUsageHTML(body []byte) (*OllamaCloudUsageData, error) { + doc, err := html.Parse(strings.NewReader(string(body))) + if err != nil { + return nil, fmt.Errorf("parse settings HTML: %w", err) + } + pageText := normalizedNodeText(doc) + lowerPage := strings.ToLower(pageText) + if containsAny(lowerPage, "sign in to ollama", "log in to ollama", "continue to sign in") { + return nil, errOllamaCloudUsageUnauthorizedHTML + } + + data := &OllamaCloudUsageData{} + data.Plan = valueBesideLabel(doc, []string{"cloud usage"}, 80) + if data.Plan == "" { + data.Plan = valueBesideLabel(doc, []string{"plan", "subscription"}, 80) + } + data.FiveHour = parseOllamaUsageWindow(doc, ollamaFiveHourUsageAliases) + data.SevenDay = parseOllamaUsageWindow(doc, ollamaSevenDayUsageAliases) + data.Balance = valueBesideLabel(doc, []string{"balance remaining"}, 80) + if data.Balance == "" { + if match := ollamaBalancePattern.FindStringSubmatch(pageText); len(match) == 2 { + data.Balance = strings.Join(strings.Fields(match[1]), "") + } + } + data.Models = parseOllamaModels(doc) + + if data.Plan == "" && data.FiveHour == nil && data.SevenDay == nil && data.Balance == "" && len(data.Models) == 0 { + return nil, fmt.Errorf("settings HTML does not contain recognizable usage fields") + } + return data, nil +} + +func parseOllamaUsageWindow(root *html.Node, aliases []string) *OllamaCloudUsageWindow { + label := findLabelElement(root, aliases) + if label == nil { + return nil + } + var candidate *OllamaCloudUsageWindow + block := label + for depth := 0; block != nil && depth < 6; depth, block = depth+1, block.Parent { + text := normalizedNodeText(block) + if len(text) > 600 { + break + } + percent, ok := ollamaUsagePercentFromText(text) + if !ok { + percent, ok = ollamaUsagePercentFromTrack(block) + } + if !ok { + continue + } + if strings.Contains(strings.ToLower(text), "remaining") && !strings.Contains(strings.ToLower(text), "used") { + percent = 100 - percent + } + window := &OllamaCloudUsageWindow{UsedPercent: percent} + window.ResetAt = timeElementValue(block) + if reset := ollamaResetPattern.FindStringSubmatch(text); len(reset) == 2 { + window.ResetText = strings.TrimSpace(reset[1]) + if window.ResetAt == nil { + window.ResetAt = parseOllamaResetTime(window.ResetText) + } + } + if candidate == nil { + candidate = window + } + if window.ResetAt != nil || window.ResetText != "" { + return window + } + } + return candidate +} + +func valueBesideLabel(root *html.Node, aliases []string, maxLen int) string { + label := findLabelElement(root, aliases) + if label == nil { + return "" + } + if sibling := nextElementSibling(label); sibling != nil { + value := strings.Trim(normalizedNodeText(sibling), ":-| ") + if value != "" && len(value) <= maxLen && !strings.EqualFold(value, "manage") { + return value + } + } + for depth, block := 0, label.Parent; block != nil && depth < 4; depth, block = depth+1, block.Parent { + text := normalizedNodeText(block) + if text == "" || len(text) > maxLen { + continue + } + value := text + for _, alias := range aliases { + if index := strings.Index(strings.ToLower(value), alias); index >= 0 { + value = strings.TrimSpace(value[:index] + " " + value[index+len(alias):]) + break + } + } + value = strings.Trim(value, ":-| ") + if value != "" && !strings.EqualFold(value, "manage") { + return value + } + } + return "" +} + +func findLabelElement(root *html.Node, aliases []string) *html.Node { + var best *html.Node + bestLen := int(^uint(0) >> 1) + walkHTML(root, func(node *html.Node) { + if node.Type != html.ElementNode || !isOllamaParserContainer(node.Data) { + return + } + text := strings.ToLower(normalizedNodeText(node)) + if text == "" { + return + } + for _, alias := range aliases { + if (text == alias || strings.HasPrefix(text, alias+" ") || strings.HasPrefix(text, alias+":")) && len(text) < bestLen { + best, bestLen = node, len(text) + } + } + }) + return best +} + +func parseOllamaModels(root *html.Node) []OllamaCloudUsageModel { + models := make([]OllamaCloudUsageModel, 0) + seen := make(map[string]struct{}) + appendModel := func(node *html.Node, modelValue, requestsValue string) { + model := strings.TrimSpace(modelValue) + requests, ok := parseOllamaRequestCount(requestsValue) + window := ollamaModelWindow(node) + if !ok || model == "" || len(model) > 128 || window == "" { + return + } + key := model + "\x00" + string(window) + if _, duplicate := seen[key]; duplicate { + return + } + seen[key] = struct{}{} + models = append(models, OllamaCloudUsageModel{Model: model, Window: window, Requests: requests}) + } + + walkHTML(root, func(node *html.Node) { + if node.Type != html.ElementNode { + return + } + if _, segment := htmlAttribute(node, "data-usage-segment"); !segment { + return + } + model, modelOK := htmlAttributeInSubtree(node, "data-model") + requests, requestsOK := htmlAttributeInSubtree(node, "data-requests") + if modelOK && requestsOK { + appendModel(node, model, requests) + } + }) + + // Older settings variants may expose the same narrow attributes without the + // segment marker. Do not infer request counts from percentages or limits. + if len(models) == 0 { + walkHTML(root, func(node *html.Node) { + if node.Type != html.ElementNode { + return + } + model, modelOK := htmlAttribute(node, "data-model") + requests, requestsOK := htmlAttribute(node, "data-requests") + if modelOK && requestsOK { + appendModel(node, model, requests) + } + }) + } + + if len(models) == 0 { + heading := findLabelElement(root, []string{"models", "available models"}) + for depth, block := 0, parentNode(heading); block != nil && depth < 4; depth, block = depth+1, block.Parent { + walkHTML(block, func(node *html.Node) { + if node.Type != html.ElementNode || (node.Data != "li" && node.Data != "code") { + return + } + match := ollamaModelFallbackPattern.FindStringSubmatch(strings.TrimSpace(normalizedNodeText(node))) + if len(match) == 3 { + appendModel(node, match[1], match[2]) + } + }) + if len(models) > 0 { + break + } + } + } + + sort.Slice(models, func(i, j int) bool { + if models[i].Window != models[j].Window { + return models[i].Window < models[j].Window + } + return models[i].Model < models[j].Model + }) + if len(models) > 100 { + models = models[:100] + } + return models +} + +func ollamaModelWindow(node *html.Node) OllamaCloudUsageModelWindow { + for block := parentNode(node); block != nil; block = block.Parent { + if value, ok := htmlAttribute(block, "data-usage-window"); ok { + normalized := strings.ToLower(strings.NewReplacer("-", " ", "_", " ").Replace(value)) + if containsAny(normalized, "five hour", "5 hour", "5h", "session") { + return OllamaCloudUsageModelWindowFiveHour + } + if containsAny(normalized, "seven day", "7 day", "7d", "weekly") { + return OllamaCloudUsageModelWindowSevenDay + } + } + text := strings.ToLower(normalizedNodeText(block)) + fiveHour := containsAny(text, ollamaFiveHourUsageAliases...) + sevenDay := containsAny(text, ollamaSevenDayUsageAliases...) + if fiveHour && !sevenDay { + return OllamaCloudUsageModelWindowFiveHour + } + if sevenDay && !fiveHour { + return OllamaCloudUsageModelWindowSevenDay + } + } + return "" +} + +func parseOllamaRequestCount(value string) (int64, bool) { + value = strings.ReplaceAll(strings.TrimSpace(value), ",", "") + requests, err := strconv.ParseInt(value, 10, 64) + return requests, err == nil && requests >= 0 +} + +func parentNode(node *html.Node) *html.Node { + if node == nil { + return nil + } + return node.Parent +} + +func nextElementSibling(node *html.Node) *html.Node { + if node == nil { + return nil + } + for sibling := node.NextSibling; sibling != nil; sibling = sibling.NextSibling { + if sibling.Type == html.ElementNode { + return sibling + } + } + return nil +} + +func timeElementValue(root *html.Node) *time.Time { + var parsed *time.Time + walkHTML(root, func(node *html.Node) { + if parsed != nil || node.Type != html.ElementNode { + return + } + if node.Data == "time" { + if value, ok := htmlAttribute(node, "datetime"); ok { + parsed = parseOllamaResetTime(value) + } + if parsed != nil { + return + } + } + if node.Data == "local-time" || htmlClassToken(node, "local-time") { + if value, ok := htmlAttribute(node, "data-time"); ok { + parsed = parseOllamaResetTime(value) + } + } + }) + return parsed +} + +func ollamaUsagePercentFromText(value string) (float64, bool) { + match := ollamaUsagePercentPattern.FindStringSubmatch(value) + if len(match) != 2 { + return 0, false + } + percent, err := strconv.ParseFloat(match[1], 64) + return percent, err == nil && percent >= 0 && percent <= 100 +} + +func ollamaUsagePercentFromTrack(root *html.Node) (float64, bool) { + var percent float64 + var found bool + walkHTML(root, func(node *html.Node) { + if found || node.Type != html.ElementNode { + return + } + if _, ok := htmlAttribute(node, "data-usage-track"); !ok { + return + } + var segmentTotal float64 + var segmentFound bool + walkHTML(node, func(segment *html.Node) { + if segment.Type != html.ElementNode { + return + } + if _, ok := htmlAttribute(segment, "data-usage-segment"); !ok { + return + } + if value, ok := cssWidthPercent(segment); ok { + segmentTotal += value + segmentFound = true + } + }) + if segmentFound && segmentTotal >= 0 && segmentTotal <= 100 { + percent, found = segmentTotal, true + return + } + walkHTML(node, func(child *html.Node) { + if found || child.Type != html.ElementNode { + return + } + if value, ok := cssWidthPercent(child); ok { + percent, found = value, true + } + }) + }) + return percent, found +} + +func cssWidthPercent(node *html.Node) (float64, bool) { + style, ok := htmlAttribute(node, "style") + if !ok { + return 0, false + } + match := ollamaUsageWidthPattern.FindStringSubmatch(style) + if len(match) != 2 { + return 0, false + } + percent, err := strconv.ParseFloat(match[1], 64) + return percent, err == nil && percent >= 0 && percent <= 100 +} + +func htmlAttribute(node *html.Node, key string) (string, bool) { + if node == nil { + return "", false + } + for _, attr := range node.Attr { + if strings.EqualFold(attr.Key, key) { + return attr.Val, true + } + } + return "", false +} + +func htmlAttributeInSubtree(root *html.Node, key string) (string, bool) { + var value string + var found bool + walkHTML(root, func(node *html.Node) { + if found || node.Type != html.ElementNode { + return + } + value, found = htmlAttribute(node, key) + }) + return value, found +} + +func htmlClassToken(node *html.Node, token string) bool { + value, ok := htmlAttribute(node, "class") + if !ok { + return false + } + for _, className := range strings.Fields(value) { + if className == token { + return true + } + } + return false +} + +func parseOllamaResetTime(value string) *time.Time { + value = strings.TrimSpace(value) + for _, layout := range []string{time.RFC3339Nano, time.RFC3339, "Jan 2, 2006 3:04 PM MST", "January 2, 2006 3:04 PM MST"} { + if parsed, err := time.Parse(layout, value); err == nil && !parsed.IsZero() { + parsed = parsed.UTC() + return &parsed + } + } + return nil +} + +func normalizedNodeText(node *html.Node) string { + if node == nil { + return "" + } + var parts []string + var collect func(*html.Node) + collect = func(current *html.Node) { + if current.Type == html.ElementNode && (current.Data == "script" || current.Data == "style" || current.Data == "noscript") { + return + } + if current.Type == html.TextNode { + if value := strings.TrimSpace(current.Data); value != "" { + parts = append(parts, value) + } + } + for child := current.FirstChild; child != nil; child = child.NextSibling { + collect(child) + } + } + collect(node) + return strings.Join(strings.Fields(strings.Join(parts, "\n")), " ") +} + +func walkHTML(node *html.Node, visit func(*html.Node)) { + if node == nil { + return + } + visit(node) + for child := node.FirstChild; child != nil; child = child.NextSibling { + walkHTML(child, visit) + } +} + +func isOllamaParserContainer(tag string) bool { + switch tag { + case "div", "section", "article", "li", "p", "span", "dt", "dd", "h1", "h2", "h3", "h4": + return true + default: + return false + } +} + +func containsAny(value string, candidates ...string) bool { + for _, candidate := range candidates { + if strings.Contains(value, candidate) { + return true + } + } + return false +} diff --git a/backend/internal/service/ollama_cloud_usage_test.go b/backend/internal/service/ollama_cloud_usage_test.go new file mode 100644 index 0000000000..851b30cb48 --- /dev/null +++ b/backend/internal/service/ollama_cloud_usage_test.go @@ -0,0 +1,955 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/stretchr/testify/require" +) + +type ollamaUsageTestEncryptor struct{} + +func (ollamaUsageTestEncryptor) Encrypt(value string) (string, error) { return "cipher:" + value, nil } +func (ollamaUsageTestEncryptor) Decrypt(value string) (string, error) { + if !strings.HasPrefix(value, "cipher:") { + return "", errors.New("authentication failed") + } + return strings.TrimPrefix(value, "cipher:"), nil +} + +type ollamaUsageTestRepo struct { + *upstreamBillingProbeAccountRepo + due []Account + beforeSnapshot func() + disableAutoAttempts atomic.Int64 + disableAutoCalls atomic.Int64 + groupResolveCalls atomic.Int64 +} + +func (r *ollamaUsageTestRepo) ListOllamaCloudUsageGroupAccounts(_ context.Context, anchors []*Account) ([]Account, error) { + r.groupResolveCalls.Add(1) + r.mu.Lock() + defer r.mu.Unlock() + wanted := make(map[string]struct{}, len(anchors)) + for _, anchor := range anchors { + if fingerprint, ok := ollamaCloudUsageGroupFingerprint(anchor); ok { + wanted[fingerprint] = struct{}{} + } + } + result := make([]Account, 0, len(r.accounts)) + for _, account := range r.accounts { + fingerprint, ok := ollamaCloudUsageGroupFingerprint(account) + if _, match := wanted[fingerprint]; !ok || !match { + continue + } + clone := *account + clone.Extra = make(map[string]any, len(account.Extra)) + for key, value := range account.Extra { + clone.Extra[key] = value + } + result = append(result, clone) + } + return result, nil +} + +func (r *ollamaUsageTestRepo) SaveOllamaCloudUsageSession(_ context.Context, expected *Account, ciphertext string, autoRefresh bool) error { + r.mu.Lock() + defer r.mu.Unlock() + members, err := r.ollamaGroupMembersLocked(expected) + if err != nil { + return err + } + for _, account := range members { + account.Extra[OllamaCloudUsageSessionExtraKey] = ciphertext + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = autoRefresh + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } + return nil +} + +func (r *ollamaUsageTestRepo) DeleteOllamaCloudUsageSession(_ context.Context, expected *Account) error { + r.mu.Lock() + defer r.mu.Unlock() + members, err := r.ollamaGroupMembersLocked(expected) + if err != nil { + return err + } + for _, account := range members { + delete(account.Extra, OllamaCloudUsageSessionExtraKey) + delete(account.Extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } + return nil +} + +func (r *ollamaUsageTestRepo) SetOllamaCloudUsageAutoRefresh(_ context.Context, expected *Account, enabled bool) error { + r.mu.Lock() + defer r.mu.Unlock() + members, err := r.ollamaGroupMembersLocked(expected) + if err != nil || !r.ollamaExpectedSessionExistsLocked(members, expected) { + return ErrOllamaCloudUsageIdentityChanged + } + for _, account := range members { + applyOllamaUsageTestManagedExtra(account, expected) + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = enabled + } + return nil +} + +func (r *ollamaUsageTestRepo) UpdateOllamaCloudUsageSnapshot(_ context.Context, expected *Account, snapshot *OllamaCloudUsageSnapshot) error { + if r.beforeSnapshot != nil { + r.beforeSnapshot() + } + r.mu.Lock() + defer r.mu.Unlock() + members, err := r.ollamaGroupMembersLocked(expected) + if err != nil || !r.ollamaExpectedSessionExistsLocked(members, expected) { + return ErrOllamaCloudUsageIdentityChanged + } + for _, account := range members { + applyOllamaUsageTestManagedExtra(account, expected) + account.Extra[OllamaCloudUsageSnapshotExtraKey] = snapshot + } + return nil +} + +func (r *ollamaUsageTestRepo) DisableOllamaCloudUsageAutoRefresh(_ context.Context, expected *Account) error { + r.disableAutoAttempts.Add(1) + r.mu.Lock() + defer r.mu.Unlock() + members, err := r.ollamaGroupMembersLocked(expected) + if err != nil || !r.ollamaExpectedSessionExistsLocked(members, expected) { + return ErrOllamaCloudUsageIdentityChanged + } + for _, account := range members { + applyOllamaUsageTestManagedExtra(account, expected) + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = false + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } + r.disableAutoCalls.Add(1) + return nil +} + +func (r *ollamaUsageTestRepo) ollamaGroupMembersLocked(expected *Account) ([]*Account, error) { + anchor := r.accounts[expected.ID] + if !sameOllamaUsageTestIdentity(anchor, expected) { + return nil, ErrOllamaCloudUsageIdentityChanged + } + fingerprint, ok := ollamaCloudUsageGroupFingerprint(expected) + if !ok { + return nil, ErrOllamaCloudUsageAccountInvalid + } + members := make([]*Account, 0, len(r.accounts)) + for _, account := range r.accounts { + candidate, valid := ollamaCloudUsageGroupFingerprint(account) + if valid && candidate == fingerprint { + if account.Extra == nil { + account.Extra = make(map[string]any) + } + members = append(members, account) + } + } + return members, nil +} + +func (r *ollamaUsageTestRepo) ollamaExpectedSessionExistsLocked(members []*Account, expected *Account) bool { + for _, member := range members { + if member.Extra[OllamaCloudUsageSessionExtraKey] == expected.Extra[OllamaCloudUsageSessionExtraKey] { + return true + } + } + return false +} + +func applyOllamaUsageTestManagedExtra(account, source *Account) { + for _, key := range []string{OllamaCloudUsageSessionExtraKey, OllamaCloudUsageAutoRefreshExtraKey, OllamaCloudUsageSnapshotExtraKey} { + delete(account.Extra, key) + if value, ok := source.Extra[key]; ok { + account.Extra[key] = value + } + } +} + +func (r *ollamaUsageTestRepo) ListDueOllamaCloudUsageAccounts(_ context.Context, _ time.Time, limit int) ([]Account, error) { + if len(r.due) > 0 { + return append([]Account(nil), r.due[:min(limit, len(r.due))]...), nil + } + r.mu.Lock() + defer r.mu.Unlock() + out := make([]Account, 0, len(r.accounts)) + for _, account := range r.accounts { + out = append(out, *account) + if len(out) == limit { + break + } + } + return out, nil +} + +type ollamaRefreshPreflightIdentityChangeRepo struct { + *ollamaUsageTestRepo + getCalls atomic.Int64 +} + +func (r *ollamaRefreshPreflightIdentityChangeRepo) GetByID(ctx context.Context, id int64) (*Account, error) { + if r.getCalls.Add(1) == 2 { + r.mu.Lock() + r.accounts[id].Credentials["api_key"] = "rotated-before-refresh" + r.mu.Unlock() + } + return r.upstreamBillingProbeAccountRepo.GetByID(ctx, id) +} + +type ollamaManagedExtraUpdateRepo struct { + AccountRepository + account *Account + updated *Account +} + +func (r *ollamaManagedExtraUpdateRepo) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *ollamaManagedExtraUpdateRepo) Update(_ context.Context, account *Account) error { + r.updated = account + return nil +} + +func sameOllamaUsageTestIdentity(left, right *Account) bool { + return left != nil && right != nil && left.Platform == right.Platform && left.Type == right.Type && + reflect.DeepEqual(left.Credentials, right.Credentials) && reflect.DeepEqual(left.ProxyID, right.ProxyID) +} + +type ollamaUsageHTTPStub struct { + status int + body []byte + header http.Header + calls atomic.Int64 + active atomic.Int64 + maxActive atomic.Int64 + beforeResponse func(*http.Request) + lastRequest *http.Request + lastProxyURL string + mu sync.Mutex +} + +func (s *ollamaUsageHTTPStub) Do(req *http.Request, proxyURL string, _ int64, _ int) (*http.Response, error) { + s.calls.Add(1) + active := s.active.Add(1) + defer s.active.Add(-1) + for { + peak := s.maxActive.Load() + if active <= peak || s.maxActive.CompareAndSwap(peak, active) { + break + } + } + s.mu.Lock() + s.lastRequest = req + s.lastProxyURL = proxyURL + s.mu.Unlock() + if s.beforeResponse != nil { + s.beforeResponse(req) + } + status := s.status + if status == 0 { + status = http.StatusOK + } + header := s.header + if header == nil { + header = http.Header{"Content-Type": []string{"text/html; charset=utf-8"}} + } + return &http.Response{StatusCode: status, Header: header, Body: io.NopCloser(strings.NewReader(string(s.body))), Request: req}, nil +} + +func (s *ollamaUsageHTTPStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, concurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return s.Do(req, proxyURL, accountID, concurrency) +} + +func ollamaUsageAccount(id int64) *Account { + return &Account{ + ID: id, Name: fmt.Sprintf("ollama-%d", id), Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ollama.com", "api_key": fmt.Sprintf("key-%d", id)}, + Extra: map[string]any{}, Status: StatusActive, Schedulable: true, Concurrency: 1, + } +} + +func newOllamaUsageTestService(t *testing.T, repo *ollamaUsageTestRepo, upstream HTTPUpstream, settingsRepo SettingRepository, fixedKey bool) *OllamaCloudUsageService { + t.Helper() + svc := NewOllamaCloudUsageService(repo, upstream, NewSettingService(settingsRepo, nil), ollamaUsageTestEncryptor{}, fixedKey) + t.Cleanup(svc.Stop) + return svc +} + +func ollamaUsageFixture(t *testing.T) []byte { + t.Helper() + body, err := os.ReadFile("testdata/ollama_settings_usage.html") + require.NoError(t, err) + return body +} + +func TestOllamaCloudUsageSettingsDefaultOffAndValidation(t *testing.T) { + repo := &upstreamBillingProbeSettingRepo{} + settingsService := NewSettingService(repo, nil) + settings, err := settingsService.GetOllamaCloudUsageSettings(context.Background()) + require.NoError(t, err) + require.False(t, settings.Enabled) + require.Equal(t, 60, settings.IntervalMinutes) + + err = settingsService.SetOllamaCloudUsageSettings(context.Background(), &OllamaCloudUsageSettings{Enabled: true, IntervalMinutes: 14}) + require.Error(t, err) + err = settingsService.SetOllamaCloudUsageSettings(context.Background(), &OllamaCloudUsageSettings{Enabled: true, IntervalMinutes: 90}) + require.NoError(t, err) + settings, err = settingsService.GetOllamaCloudUsageSettings(context.Background()) + require.NoError(t, err) + require.True(t, settings.Enabled) + require.Equal(t, 90, settings.IntervalMinutes) +} + +func TestIsOllamaCloudUsageAccountStrictOfficialHost(t *testing.T) { + tests := []struct { + baseURL string + platform string + want bool + }{ + {"https://ollama.com", PlatformOpenAI, true}, + {"HTTPS://OLLAMA.COM", PlatformAnthropic, true}, + {"https://www.OLLAMA.com:443/v1", PlatformOpenAI, true}, + {"https://ollama.com:443", PlatformOpenAI, true}, + {"https://ollama.com/", PlatformAnthropic, false}, + {"https://ollama.com/v1/", PlatformOpenAI, false}, + {"http://ollama.com", PlatformOpenAI, false}, + {"https://ollama.com.evil.test", PlatformOpenAI, false}, + {"https://ollama.com:444", PlatformOpenAI, false}, + {"https://user@ollama.com", PlatformOpenAI, false}, + {"https://ollama.com/v2", PlatformOpenAI, false}, + {"https://ollama.com?next=https://evil.test", PlatformOpenAI, false}, + {"https://ollama.com#usage", PlatformOpenAI, false}, + } + for _, test := range tests { + t.Run(test.baseURL+test.platform, func(t *testing.T) { + account := ollamaUsageAccount(1) + account.Platform = test.platform + account.Credentials["base_url"] = test.baseURL + require.Equal(t, test.want, IsOllamaCloudUsageAccount(account)) + }) + } +} + +func TestNormalizeOllamaCloudUsageCookieAllowlist(t *testing.T) { + normalized, err := normalizeOllamaCloudUsageCookie(" tracking=discard ; wos-session=secret ; __Secure-authjs.session-token.0=part-a ; device=discard ") + require.NoError(t, err) + require.Equal(t, "wos-session=secret; __Secure-authjs.session-token.0=part-a", normalized) + + normalized, err = normalizeOllamaCloudUsageCookie(" \t\r\nwos-session=secret; tracking=discard\r\n\t ") + require.NoError(t, err) + require.Equal(t, "wos-session=secret", normalized) + + _, err = normalizeOllamaCloudUsageCookie("wos-session=secret\r\nHost: evil.test") + require.ErrorContains(t, err, "invalid header") + + for _, allowed := range []string{ + "wos-session", "__Secure-session", "session", "ollama_session", "__Host-ollama_session", + "next-auth.session-token", "next-auth.session-token.0", "__Secure-next-auth.session-token.12", + "authjs.session-token", "__Secure-authjs.session-token.1", + } { + normalized, err := normalizeOllamaCloudUsageCookie(allowed + "=value") + require.NoError(t, err, allowed) + require.Equal(t, allowed+"=value", normalized) + } + + for _, invalid := range []string{ + "", "Domain=ollama.com; wos-session=x", "wos-session=x; Path=/", + "wos-session=x; wos-session=y", "Secure", "tracking=only", "__session=arbitrary", + "authjs.session-token.bad=not-a-shard", "Authjs.session-token=wrong-case", + } { + _, err := normalizeOllamaCloudUsageCookie(invalid) + require.Error(t, err, invalid) + } + _, err = normalizeOllamaCloudUsageCookie("wos-session=" + strings.Repeat("x", ollamaCloudUsageMaxSessionBytes)) + require.ErrorContains(t, err, "too large") +} + +func TestParseOllamaCloudUsageHTMLFixture(t *testing.T) { + data, err := parseOllamaCloudUsageHTML(ollamaUsageFixture(t)) + require.NoError(t, err) + require.Equal(t, "max", data.Plan) + require.NotNil(t, data.FiveHour) + require.Equal(t, 5.6, data.FiveHour.UsedPercent) + require.NotNil(t, data.FiveHour.ResetAt) + require.Equal(t, time.Date(2026, time.July, 23, 3, 0, 0, 0, time.UTC), *data.FiveHour.ResetAt) + require.NotNil(t, data.SevenDay) + require.Equal(t, 14.2, data.SevenDay.UsedPercent) + require.NotNil(t, data.SevenDay.ResetAt) + require.Equal(t, time.Date(2026, time.July, 29, 0, 0, 0, 0, time.UTC), *data.SevenDay.ResetAt) + require.Equal(t, "$0", data.Balance) + require.Equal(t, []OllamaCloudUsageModel{ + {Model: "gpt-oss:120b-cloud", Window: OllamaCloudUsageModelWindowFiveHour, Requests: 2}, + {Model: "qwen3-coder:480b-cloud", Window: OllamaCloudUsageModelWindowFiveHour, Requests: 3}, + {Model: "gpt-oss:120b-cloud", Window: OllamaCloudUsageModelWindowSevenDay, Requests: 12}, + {Model: "qwen3-coder:480b-cloud", Window: OllamaCloudUsageModelWindowSevenDay, Requests: 13}, + }, data.Models) + + _, err = parseOllamaCloudUsageHTML([]byte(`
Sign in to Ollama
`)) + require.ErrorIs(t, err, errOllamaCloudUsageUnauthorizedHTML) + _, err = parseOllamaCloudUsageHTML([]byte(`

5 hour usage 42% used

Sign in to Ollama
`)) + require.ErrorIs(t, err, errOllamaCloudUsageUnauthorizedHTML) + _, err = parseOllamaCloudUsageHTML([]byte(`
unrelated settings
`)) + require.Error(t, err) +} + +func TestParseOllamaCloudUsageHTMLMissingOptionalFieldsAndCSSWidthFallback(t *testing.T) { + data, err := parseOllamaCloudUsageHTML([]byte(` +
+

5 hour usage

+
+
+
+
+
`)) + require.NoError(t, err) + require.Equal(t, 23.5, data.FiveHour.UsedPercent) + require.Nil(t, data.FiveHour.ResetAt) + require.Empty(t, data.Plan) + require.Nil(t, data.SevenDay) + require.Empty(t, data.Balance) + require.Equal(t, []OllamaCloudUsageModel{{ + Model: "model-a", Window: OllamaCloudUsageModelWindowFiveHour, Requests: 1234, + }}, data.Models) +} + +func TestParseOllamaCloudUsageHTMLResetElementVariants(t *testing.T) { + const want = "2026-07-23T03:00:00Z" + for name, element := range map[string]string{ + "time datetime": ``, + "custom element": `2 hours.`, + "class token": `2 hours.`, + } { + t.Run(name, func(t *testing.T) { + data, err := parseOllamaCloudUsageHTML([]byte( + `
Session usage1% used
Resets in ` + element + `
`, + )) + require.NoError(t, err) + require.NotNil(t, data.FiveHour) + require.NotNil(t, data.FiveHour.ResetAt) + require.Equal(t, want, data.FiveHour.ResetAt.Format(time.RFC3339)) + }) + } +} + +func TestParseOllamaCloudUsageHTMLPlanAndBalanceFallbacks(t *testing.T) { + data, err := parseOllamaCloudUsageHTML([]byte(` +
+

Cloud usagemax

+
PlanPro
+

Credits currently available: USD $9.50

+
`)) + require.NoError(t, err) + require.Equal(t, "max", data.Plan) + require.Equal(t, "USD$9.50", data.Balance) + + data, err = parseOllamaCloudUsageHTML([]byte(`
SubscriptionPro
`)) + require.NoError(t, err) + require.Equal(t, "Pro", data.Plan) +} + +func TestOllamaCloudUsageManagedExtraCannotBeImported(t *testing.T) { + remoteExtra := map[string]any{ + OllamaCloudUsageSessionExtraKey: "remote-ciphertext", + OllamaCloudUsageAutoRefreshExtraKey: true, + OllamaCloudUsageSnapshotExtraKey: map[string]any{"status": "forged"}, + } + created, err := buildAccountForCreate(&CreateAccountInput{ + Name: "ollama", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"base_url": "https://ollama.com", "api_key": "key"}, + Concurrency: 1, + }, mergeMap(nil, remoteExtra)) + require.NoError(t, err) + require.NotContains(t, created.Extra, OllamaCloudUsageSessionExtraKey) + require.NotContains(t, created.Extra, OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, created.Extra, OllamaCloudUsageSnapshotExtraKey) + + existing := ollamaUsageAccount(6) + existing.Extra = map[string]any{ + OllamaCloudUsageSessionExtraKey: "local-ciphertext", + OllamaCloudUsageAutoRefreshExtraKey: false, + OllamaCloudUsageSnapshotExtraKey: map[string]any{"status": OllamaCloudUsageStatusOK}, + } + targetExtra := mergeMap(existing.Extra, remoteExtra) + reconcileCRSUpstreamBillingProbeExtra(existing, existing.Platform, existing.Type, mergeMap(existing.Credentials, nil), targetExtra) + require.Equal(t, "local-ciphertext", targetExtra[OllamaCloudUsageSessionExtraKey]) + require.Equal(t, false, targetExtra[OllamaCloudUsageAutoRefreshExtraKey]) + require.Equal(t, map[string]any{"status": OllamaCloudUsageStatusOK}, targetExtra[OllamaCloudUsageSnapshotExtraKey]) + + changedCredentials := mergeMap(existing.Credentials, map[string]any{"api_key": "rotated"}) + targetExtra = mergeMap(existing.Extra, remoteExtra) + reconcileCRSUpstreamBillingProbeExtra(existing, existing.Platform, existing.Type, changedCredentials, targetExtra) + require.NotContains(t, targetExtra, OllamaCloudUsageSessionExtraKey) + require.NotContains(t, targetExtra, OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, targetExtra, OllamaCloudUsageSnapshotExtraKey) +} + +func TestAccountServiceUpdateStripsOllamaManagedExtra(t *testing.T) { + account := ollamaUsageAccount(61) + account.Extra = map[string]any{ + OllamaCloudUsageSessionExtraKey: "local-ciphertext", + OllamaCloudUsageAutoRefreshExtraKey: true, + OllamaCloudUsageSnapshotExtraKey: map[string]any{"status": OllamaCloudUsageStatusOK}, + } + repo := &ollamaManagedExtraUpdateRepo{account: account} + svc := NewAccountService(repo, nil) + requestedExtra := map[string]any{ + "note": "preserved", + OllamaCloudUsageSessionExtraKey: "forged-ciphertext", + OllamaCloudUsageAutoRefreshExtraKey: nil, + OllamaCloudUsageSnapshotExtraKey: nil, + } + + _, err := svc.Update(context.Background(), account.ID, UpdateAccountRequest{Extra: &requestedExtra}) + require.NoError(t, err) + require.Equal(t, "preserved", repo.updated.Extra["note"]) + require.NotContains(t, repo.updated.Extra, OllamaCloudUsageSessionExtraKey) + require.NotContains(t, repo.updated.Extra, OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, repo.updated.Extra, OllamaCloudUsageSnapshotExtraKey) + // The request map is not mutated while managed fields are stripped. + require.Contains(t, requestedExtra, OllamaCloudUsageSessionExtraKey) +} + +func TestOllamaCloudUsageSessionEncryptionFailClosedAndWriteOnlyState(t *testing.T) { + account := ollamaUsageAccount(7) + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{7: account}}} + settings := &upstreamBillingProbeSettingRepo{} + + ephemeral := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{}, settings, false) + _, err := ephemeral.SaveSession(context.Background(), 7, "wos-session=plaintext-secret") + require.ErrorIs(t, err, ErrOllamaCloudUsageEncryptionKey) + require.NotContains(t, account.Extra, OllamaCloudUsageSessionExtraKey) + + svc := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{}, settings, true) + _, err = svc.SaveSession(context.Background(), 7, "tracking=arbitrary-only") + require.Error(t, err) + require.NotContains(t, account.Extra, OllamaCloudUsageSessionExtraKey) + + state, err := svc.SaveSession(context.Background(), 7, "tracking=must-not-persist; wos-session=plaintext-secret") + require.NoError(t, err) + require.True(t, state.Configured) + stored, ok := account.Extra[OllamaCloudUsageSessionExtraKey].(string) + require.True(t, ok) + require.Equal(t, "cipher:wos-session=plaintext-secret", stored) + require.NotContains(t, stored, "tracking") + raw, err := json.Marshal(state) + require.NoError(t, err) + require.NotContains(t, string(raw), "plaintext-secret") + require.NotContains(t, string(raw), "cipher:") + + account.Extra[OllamaCloudUsageSessionExtraKey] = "plaintext-secret" + _, err = svc.Refresh(context.Background(), 7) + require.ErrorContains(t, err, "cannot be decrypted") +} + +func TestOllamaCloudUsageGroupSharesAcrossPlatformsURLVariantsAndDynamicSiblings(t *testing.T) { + source := ollamaUsageAccount(71) + source.Credentials["api_key"] = "shared-key" + source.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=shared" + source.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + source.Extra[OllamaCloudUsageSnapshotExtraKey] = &OllamaCloudUsageSnapshot{ + Status: OllamaCloudUsageStatusOK, + Data: &OllamaCloudUsageData{Plan: "pro"}, + } + source.UpdatedAt = time.Now().Add(-time.Minute) + sibling := ollamaUsageAccount(72) + sibling.Platform = PlatformAnthropic + sibling.Credentials = map[string]any{"base_url": "HTTPS://WWW.OLLAMA.COM:443/v1", "api_key": "shared-key"} + sibling.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=shared" + sibling.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + sibling.UpdatedAt = time.Now() + different := ollamaUsageAccount(73) + different.Credentials["api_key"] = "different-key" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + source.ID: source, sibling.ID: sibling, different.ID: different, + }}} + svc := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{}, &upstreamBillingProbeSettingRepo{}, true) + + state, err := svc.GetState(context.Background(), sibling.ID) + require.NoError(t, err) + require.True(t, state.Configured) + require.True(t, state.AutoRefreshEnabled) + require.Equal(t, "pro", state.Snapshot.Data.Plan) + + differentState, err := svc.GetState(context.Background(), different.ID) + require.NoError(t, err) + require.False(t, differentState.Configured) + + newSibling := ollamaUsageAccount(74) + newSibling.Platform = PlatformAnthropic + newSibling.Credentials = map[string]any{"base_url": "https://ollama.com:443", "api_key": "shared-key"} + repo.mu.Lock() + repo.accounts[newSibling.ID] = newSibling + repo.mu.Unlock() + newState, err := svc.GetState(context.Background(), newSibling.ID) + require.NoError(t, err) + require.True(t, newState.Configured) + require.Equal(t, state.Snapshot, newState.Snapshot) + + before := repo.groupResolveCalls.Load() + require.NoError(t, svc.ResolveAccounts(context.Background(), []*Account{source, sibling, different, newSibling})) + require.Equal(t, before+1, repo.groupResolveCalls.Load(), "one list batch must issue one group lookup") +} + +func TestOllamaCloudUsageSaveAutoRefreshAndDeleteAreGroupScoped(t *testing.T) { + first := ollamaUsageAccount(81) + first.Credentials["api_key"] = "shared-key" + second := ollamaUsageAccount(82) + second.Platform = PlatformAnthropic + second.Credentials = map[string]any{"base_url": "https://www.ollama.com/v1", "api_key": "shared-key"} + different := ollamaUsageAccount(83) + different.Credentials["api_key"] = "different-key" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + first.ID: first, second.ID: second, different.ID: different, + }}} + svc := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{}, &upstreamBillingProbeSettingRepo{}, true) + + state, err := svc.SaveSession(context.Background(), second.ID, "wos-session=shared-browser") + require.NoError(t, err) + require.True(t, state.Configured) + require.Equal(t, "cipher:wos-session=shared-browser", first.Extra[OllamaCloudUsageSessionExtraKey]) + require.Equal(t, first.Extra[OllamaCloudUsageSessionExtraKey], second.Extra[OllamaCloudUsageSessionExtraKey]) + require.NotContains(t, different.Extra, OllamaCloudUsageSessionExtraKey) + + state, err = svc.SetAutoRefresh(context.Background(), first.ID, true) + require.NoError(t, err) + require.True(t, state.AutoRefreshEnabled) + require.Equal(t, true, first.Extra[OllamaCloudUsageAutoRefreshExtraKey]) + require.Equal(t, true, second.Extra[OllamaCloudUsageAutoRefreshExtraKey]) + + state, err = svc.DeleteSession(context.Background(), second.ID) + require.NoError(t, err) + require.False(t, state.Configured) + for _, member := range []*Account{first, second} { + require.NotContains(t, member.Extra, OllamaCloudUsageSessionExtraKey) + require.NotContains(t, member.Extra, OllamaCloudUsageAutoRefreshExtraKey) + require.NotContains(t, member.Extra, OllamaCloudUsageSnapshotExtraKey) + } +} + +func TestOllamaCloudUsageRefreshSingleflightAndRunnerDeduplicateSharedGroup(t *testing.T) { + first := ollamaUsageAccount(91) + first.Credentials["api_key"] = "shared-key" + first.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=shared" + first.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + second := ollamaUsageAccount(92) + second.Platform = PlatformAnthropic + second.Credentials = map[string]any{"base_url": "https://www.ollama.com:443/v1", "api_key": "shared-key"} + second.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=shared" + second.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + repo := &ollamaUsageTestRepo{ + upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{first.ID: first, second.ID: second}}, + due: []Account{*first, *second}, + } + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyOllamaCloudUsageSettings: `{"enabled":true,"interval_minutes":60}`, + }} + started := make(chan struct{}) + release := make(chan struct{}) + var once sync.Once + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t), beforeResponse: func(*http.Request) { + once.Do(func() { close(started) }) + <-release + }} + svc := newOllamaUsageTestService(t, repo, upstream, settingsRepo, true) + + errs := make(chan error, 2) + go func() { _, err := svc.Refresh(context.Background(), first.ID); errs <- err }() + <-started + go func() { _, err := svc.Refresh(context.Background(), second.ID); errs <- err }() + close(release) + require.NoError(t, <-errs) + require.NoError(t, <-errs) + require.Equal(t, int64(1), upstream.calls.Load()) + require.NotNil(t, decodeOllamaCloudUsageSnapshot(first.Extra)) + require.Equal(t, decodeOllamaCloudUsageSnapshot(first.Extra), decodeOllamaCloudUsageSnapshot(second.Extra)) + + delete(first.Extra, OllamaCloudUsageSnapshotExtraKey) + delete(second.Extra, OllamaCloudUsageSnapshotExtraKey) + upstream.beforeResponse = nil + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(2), upstream.calls.Load(), "RunDue must issue one request for the shared group") +} + +func TestOllamaCloudUsageRefreshRejectsGroupChangeBeforeUpstreamRequest(t *testing.T) { + account := ollamaUsageAccount(94) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + base := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{account.ID: account}}} + repo := &ollamaRefreshPreflightIdentityChangeRepo{ollamaUsageTestRepo: base} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := NewOllamaCloudUsageService(repo, upstream, NewSettingService(&upstreamBillingProbeSettingRepo{}, nil), ollamaUsageTestEncryptor{}, true) + t.Cleanup(svc.Stop) + + _, err := svc.Refresh(context.Background(), account.ID) + + require.ErrorIs(t, err, ErrOllamaCloudUsageIdentityChanged) + require.Zero(t, upstream.calls.Load()) + require.NotContains(t, account.Extra, OllamaCloudUsageSnapshotExtraKey) +} + +func TestOllamaCloudUsageRefreshUsesFixedURLCookieAndNoRedirects(t *testing.T) { + account := ollamaUsageAccount(8) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=browser-secret; tracking=must-not-send" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{8: account}}} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := newOllamaUsageTestService(t, repo, upstream, &upstreamBillingProbeSettingRepo{}, true) + fixedNow := time.Date(2026, time.July, 22, 15, 0, 0, 0, time.UTC) + svc.now = func() time.Time { return fixedNow } + + state, err := svc.Refresh(context.Background(), 8) + require.NoError(t, err) + require.Equal(t, OllamaCloudUsageStatusOK, state.Snapshot.Status) + require.Equal(t, "https://ollama.com/settings", upstream.lastRequest.URL.String()) + require.Equal(t, "ollama.com", upstream.lastRequest.Host) + require.Equal(t, "wos-session=browser-secret", upstream.lastRequest.Header.Get("Cookie")) + require.NotContains(t, upstream.lastRequest.Header.Get("Cookie"), "tracking") + require.Empty(t, upstream.lastRequest.Header.Get("Authorization")) + require.True(t, HTTPUpstreamRedirectsDisabled(upstream.lastRequest.Context())) +} + +func TestOllamaCloudUsageManualRefreshUsesShortIndependentInterval(t *testing.T) { + account := ollamaUsageAccount(12) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=initial" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{12: account}}} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := newOllamaUsageTestService(t, repo, upstream, &upstreamBillingProbeSettingRepo{}, true) + fixedNow := time.Date(2026, time.July, 22, 15, 0, 0, 0, time.UTC) + svc.now = func() time.Time { return fixedNow } + + _, err := svc.Refresh(context.Background(), 12) + require.NoError(t, err) + _, err = svc.Refresh(context.Background(), 12) + require.ErrorIs(t, err, ErrOllamaCloudUsageRefreshRateLimited) + require.Equal(t, int64(1), upstream.calls.Load()) + + // Saving a repaired session clears the prior snapshot, so the global 60-minute + // next_refresh_at does not block immediate administrator verification. + _, err = svc.SaveSession(context.Background(), 12, "wos-session=repaired") + require.NoError(t, err) + _, err = svc.Refresh(context.Background(), 12) + require.NoError(t, err) + require.Equal(t, int64(2), upstream.calls.Load()) +} + +func TestOllamaCloudUsageRefreshUsesHydratedProxyIdentity(t *testing.T) { + account := ollamaUsageAccount(13) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + proxyID := int64(4) + account.ProxyID = &proxyID + account.Proxy = &Proxy{ + ID: proxyID, Protocol: "http", Host: "127.0.0.1", Port: 3128, + Username: "proxy-user", Password: "proxy-pass", Status: StatusActive, + } + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{13: account}}} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := newOllamaUsageTestService(t, repo, upstream, &upstreamBillingProbeSettingRepo{}, true) + + _, err := svc.Refresh(context.Background(), 13) + require.NoError(t, err) + require.Equal(t, account.Proxy.URL(), upstream.lastProxyURL) +} + +func TestOllamaCloudUsageRedirectAndBodyLimitArePersistedSafely(t *testing.T) { + for _, test := range []struct { + name string + status int + body []byte + reason string + }{ + {"redirect", http.StatusFound, nil, "redirect_blocked"}, + {"body limit", http.StatusOK, make([]byte, ollamaCloudUsageMaxBodyBytes+1), "response_too_large"}, + } { + t.Run(test.name, func(t *testing.T) { + account := ollamaUsageAccount(9) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{9: account}}} + svc := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{status: test.status, body: test.body}, &upstreamBillingProbeSettingRepo{}, true) + state, err := svc.Refresh(context.Background(), 9) + require.NoError(t, err) + require.Equal(t, OllamaCloudUsageStatusFailed, state.Snapshot.Status) + require.Equal(t, test.reason, state.Snapshot.LastError) + }) + } +} + +func TestOllamaCloudUsageRefreshRejectsIdentityChange(t *testing.T) { + account := ollamaUsageAccount(10) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{10: account}}} + repo.beforeSnapshot = func() { account.Credentials["api_key"] = "rotated" } + svc := newOllamaUsageTestService(t, repo, &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)}, &upstreamBillingProbeSettingRepo{}, true) + _, err := svc.Refresh(context.Background(), 10) + require.ErrorIs(t, err, ErrOllamaCloudUsageIdentityChanged) + require.NotContains(t, account.Extra, OllamaCloudUsageSnapshotExtraKey) +} + +func TestOllamaCloudUsageRunnerHonorsLeaderLockAndBackoff(t *testing.T) { + account := ollamaUsageAccount(11) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{11: account}}} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyOllamaCloudUsageSettings: `{"enabled":true,"interval_minutes":60}`, + }} + cache := &fakeLeaderLockCache{} + _, acquired := tryAcquireSingletonLeaderLock(context.Background(), cache, nil, ollamaCloudUsageLeaderLockKey, "peer", time.Minute) + require.True(t, acquired) + svc := newOllamaUsageTestService(t, repo, upstream, settingsRepo, true) + svc.lockCache = cache + require.NoError(t, svc.RunDue(context.Background())) + require.Zero(t, upstream.calls.Load()) + require.NoError(t, cache.ReleaseLeaderLock(context.Background(), ollamaCloudUsageLeaderLockKey, "peer")) + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), upstream.calls.Load()) + + firstFailure := nextOllamaCloudUsageDelay(60, 1, 0) + thirdFailure := nextOllamaCloudUsageDelay(60, 3, 0) + require.Greater(t, thirdFailure, firstFailure) + require.GreaterOrEqual(t, nextOllamaCloudUsageDelay(60, 1, 3*time.Hour), 3*time.Hour) + require.LessOrEqual(t, nextOllamaCloudUsageDelay(60, 20, 0), ollamaCloudUsageMaxDelay+5*time.Minute) +} + +func TestOllamaCloudUsageRunnerDisablesAutoRefreshAfterUnpersistableIdentityError(t *testing.T) { + account := ollamaUsageAccount(14) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + missingProxyID := int64(99) + account.ProxyID = &missingProxyID + account.Proxy = nil + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{14: account}}} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyOllamaCloudUsageSettings: `{"enabled":true,"interval_minutes":60}`, + }} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := newOllamaUsageTestService(t, repo, upstream, settingsRepo, true) + + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), repo.disableAutoCalls.Load()) + require.Equal(t, false, account.Extra[OllamaCloudUsageAutoRefreshExtraKey]) + require.Zero(t, upstream.calls.Load()) + + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), repo.disableAutoCalls.Load()) + require.Zero(t, upstream.calls.Load()) +} + +func TestOllamaCloudUsageRunnerIdentityChangePreservesOldGroupAndDoesNotLoop(t *testing.T) { + anchor := ollamaUsageAccount(15) + anchor.Credentials["api_key"] = "shared-before-rotation" + anchor.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + anchor.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + sibling := ollamaUsageAccount(16) + sibling.Platform = PlatformAnthropic + sibling.Credentials = map[string]any{"api_key": "shared-before-rotation", "base_url": "https://www.ollama.com:443/v1"} + sibling.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + sibling.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + dueAnchor := *anchor + dueAnchor.Credentials = mergeMap(nil, anchor.Credentials) + dueAnchor.Extra = mergeMap(nil, anchor.Extra) + repo := &ollamaUsageTestRepo{ + upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{ + anchor.ID: anchor, sibling.ID: sibling, + }}, + due: []Account{dueAnchor}, + } + var rotateOnce sync.Once + repo.beforeSnapshot = func() { + rotateOnce.Do(func() { + repo.mu.Lock() + defer repo.mu.Unlock() + anchor.Credentials["api_key"] = "rotated-account-key" + delete(anchor.Extra, OllamaCloudUsageSessionExtraKey) + delete(anchor.Extra, OllamaCloudUsageAutoRefreshExtraKey) + delete(anchor.Extra, OllamaCloudUsageSnapshotExtraKey) + }) + } + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{ + SettingKeyOllamaCloudUsageSettings: `{"enabled":true,"interval_minutes":60}`, + }} + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t)} + svc := newOllamaUsageTestService(t, repo, upstream, settingsRepo, true) + + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), repo.disableAutoAttempts.Load()) + require.Zero(t, repo.disableAutoCalls.Load(), "the stale anchor CAS must not disable the old sibling group") + require.Equal(t, true, sibling.Extra[OllamaCloudUsageAutoRefreshExtraKey]) + require.NotContains(t, anchor.Extra, OllamaCloudUsageAutoRefreshExtraKey) + + repo.due = []Account{*anchor, *sibling} + require.NoError(t, svc.RunDue(context.Background())) + require.Equal(t, int64(1), repo.disableAutoAttempts.Load(), "the changed account must not be retried") + require.Equal(t, true, sibling.Extra[OllamaCloudUsageAutoRefreshExtraKey]) + require.NotNil(t, decodeOllamaCloudUsageSnapshot(sibling.Extra), "the still-valid sibling must refresh normally") + require.Equal(t, int64(2), upstream.calls.Load()) +} + +func TestOllamaCloudUsageSingleflightConcurrencyAndRunnerSwitches(t *testing.T) { + accounts := make(map[int64]*Account) + for id := int64(1); id <= 7; id++ { + account := ollamaUsageAccount(id) + account.Extra[OllamaCloudUsageSessionExtraKey] = "cipher:wos-session=secret" + account.Extra[OllamaCloudUsageAutoRefreshExtraKey] = true + accounts[id] = account + } + repo := &ollamaUsageTestRepo{upstreamBillingProbeAccountRepo: &upstreamBillingProbeAccountRepo{accounts: accounts}} + unblock := make(chan struct{}) + entered := make(chan struct{}, 10) + upstream := &ollamaUsageHTTPStub{body: ollamaUsageFixture(t), beforeResponse: func(*http.Request) { + entered <- struct{}{} + <-unblock + }} + settingsRepo := &upstreamBillingProbeSettingRepo{values: map[string]string{}} + svc := newOllamaUsageTestService(t, repo, upstream, settingsRepo, true) + + // Global automatic refresh is fail-safe off by default. + require.NoError(t, svc.RunDue(context.Background())) + require.Zero(t, upstream.calls.Load()) + + settingsRepo.values[SettingKeyOllamaCloudUsageSettings] = `{"enabled":true,"interval_minutes":60}` + var singleflight sync.WaitGroup + singleflight.Add(2) + for range 2 { + go func() { + defer singleflight.Done() + _, _ = svc.Refresh(context.Background(), 1) + }() + } + <-entered + close(unblock) + singleflight.Wait() + require.Equal(t, int64(1), upstream.calls.Load()) + + // Clear snapshots so all accounts are due, then verify the shared four-slot bound. + for _, account := range accounts { + delete(account.Extra, OllamaCloudUsageSnapshotExtraKey) + } + unblock2 := make(chan struct{}) + upstream.beforeResponse = func(*http.Request) { <-unblock2 } + done := make(chan struct{}) + go func() { + _ = svc.RunDue(context.Background()) + close(done) + }() + require.Eventually(t, func() bool { return upstream.active.Load() == ollamaCloudUsageConcurrency }, time.Second, 10*time.Millisecond) + close(unblock2) + <-done + require.LessOrEqual(t, upstream.maxActive.Load(), int64(ollamaCloudUsageConcurrency)) + require.Equal(t, int64(8), upstream.calls.Load()) +} diff --git a/backend/internal/service/testdata/ollama_settings_usage.html b/backend/internal/service/testdata/ollama_settings_usage.html new file mode 100644 index 0000000000..b76f75b56c --- /dev/null +++ b/backend/internal/service/testdata/ollama_settings_usage.html @@ -0,0 +1,57 @@ + + + + Settings - Ollama + +
+
+

Cloud usagemax

+
Balance remaining
$0
+
+ +
+
+
Session usage5.6% used
+
+
+
+
+
Resets in 2 hours.
+
+ +
+
Weekly usage14.2% used
+
+
+
+
+
Resets in 6 days.
+
+
+
+ + diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 166c1ff832..e1210208cb 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -726,6 +726,7 @@ var ProviderSet = wire.NewSet( ProvideAccountUsageService, ProvideAccountTestService, ProvideUpstreamBillingProbeService, + ProvideOllamaCloudUsageService, ProvideSettingService, NewDataManagementService, ProvideBackupService, diff --git a/frontend/src/api/__tests__/admin.accounts.ollamaCloudUsage.spec.ts b/frontend/src/api/__tests__/admin.accounts.ollamaCloudUsage.spec.ts new file mode 100644 index 0000000000..1777819cb1 --- /dev/null +++ b/frontend/src/api/__tests__/admin.accounts.ollamaCloudUsage.spec.ts @@ -0,0 +1,68 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { get, post, put, del } = vi.hoisted(() => ({ + get: vi.fn(), + post: vi.fn(), + put: vi.fn(), + del: vi.fn() +})) + +vi.mock('@/api/client', () => ({ + apiClient: { get, post, put, delete: del } +})) + +import { + deleteOllamaCloudUsageSession, + getOllamaCloudUsage, + getOllamaCloudUsageSettings, + refreshOllamaCloudUsage, + saveOllamaCloudUsageSession, + setOllamaCloudUsageAutoRefresh, + updateOllamaCloudUsageSettings +} from '@/api/admin/accounts' + +const state = { + account_id: 7, + eligible: true, + configured: true, + auto_refresh_enabled: false, + encryption_key_configured: true +} + +describe('admin Ollama Cloud usage API', () => { + beforeEach(() => { + get.mockReset() + post.mockReset() + put.mockReset() + del.mockReset() + }) + + it('uses dedicated global settings endpoints', async () => { + const settings = { enabled: false, interval_minutes: 60 } + get.mockResolvedValueOnce({ data: settings }) + put.mockResolvedValueOnce({ data: settings }) + + await expect(getOllamaCloudUsageSettings()).resolves.toEqual(settings) + await expect(updateOllamaCloudUsageSettings(settings)).resolves.toEqual(settings) + expect(get).toHaveBeenCalledWith('/admin/accounts/ollama-cloud-usage/settings') + expect(put).toHaveBeenCalledWith('/admin/accounts/ollama-cloud-usage/settings', settings) + }) + + it('keeps session configuration write-only and separate from account updates', async () => { + get.mockResolvedValueOnce({ data: state }) + put.mockResolvedValueOnce({ data: state }).mockResolvedValueOnce({ data: state }) + del.mockResolvedValueOnce({ data: { ...state, configured: false } }) + post.mockResolvedValueOnce({ data: state }) + + await expect(getOllamaCloudUsage(7)).resolves.toEqual(state) + await expect(saveOllamaCloudUsageSession(7, 'wos-session=secret')).resolves.toEqual(state) + await expect(setOllamaCloudUsageAutoRefresh(7, true)).resolves.toEqual(state) + await expect(refreshOllamaCloudUsage(7)).resolves.toEqual(state) + await expect(deleteOllamaCloudUsageSession(7)).resolves.toMatchObject({ configured: false }) + + expect(put).toHaveBeenNthCalledWith(1, '/admin/accounts/7/ollama-cloud-usage/session', { session: 'wos-session=secret' }) + expect(put).toHaveBeenNthCalledWith(2, '/admin/accounts/7/ollama-cloud-usage/auto-refresh', { enabled: true }) + expect(post).toHaveBeenCalledWith('/admin/accounts/7/ollama-cloud-usage/refresh') + expect(del).toHaveBeenCalledWith('/admin/accounts/7/ollama-cloud-usage/session') + }) +}) diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index 5bc9752155..40b80c1861 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -22,7 +22,9 @@ import type { CheckMixedChannelRequest, CheckMixedChannelResponse, UpstreamBillingProbeResult, - UpstreamBillingProbeSettings + UpstreamBillingProbeSettings, + OllamaCloudUsageSettings, + OllamaCloudUsageState } from '@/types' /** @@ -882,6 +884,50 @@ export async function probeUpstreamBillingBatch(accountIds: number[]): Promise { + const { data } = await apiClient.get('/admin/accounts/ollama-cloud-usage/settings') + return data +} + +export async function updateOllamaCloudUsageSettings( + settings: OllamaCloudUsageSettings +): Promise { + const { data } = await apiClient.put( + '/admin/accounts/ollama-cloud-usage/settings', + settings + ) + return data +} + +export async function getOllamaCloudUsage(id: number): Promise { + const { data } = await apiClient.get(`/admin/accounts/${id}/ollama-cloud-usage`) + return data +} + +export async function saveOllamaCloudUsageSession(id: number, session: string): Promise { + const { data } = await apiClient.put(`/admin/accounts/${id}/ollama-cloud-usage/session`, { + session + }) + return data +} + +export async function deleteOllamaCloudUsageSession(id: number): Promise { + const { data } = await apiClient.delete(`/admin/accounts/${id}/ollama-cloud-usage/session`) + return data +} + +export async function setOllamaCloudUsageAutoRefresh(id: number, enabled: boolean): Promise { + const { data } = await apiClient.put(`/admin/accounts/${id}/ollama-cloud-usage/auto-refresh`, { + enabled + }) + return data +} + +export async function refreshOllamaCloudUsage(id: number): Promise { + const { data } = await apiClient.post(`/admin/accounts/${id}/ollama-cloud-usage/refresh`) + return data +} + export const accountsAPI = { list, listWithEtag, @@ -933,7 +979,14 @@ export const accountsAPI = { updateUpstreamBillingProbeSettings, setUpstreamBillingProbeEnabled, probeUpstreamBilling, - probeUpstreamBillingBatch + probeUpstreamBillingBatch, + getOllamaCloudUsageSettings, + updateOllamaCloudUsageSettings, + getOllamaCloudUsage, + saveOllamaCloudUsageSession, + deleteOllamaCloudUsageSession, + setOllamaCloudUsageAutoRefresh, + refreshOllamaCloudUsage } export default accountsAPI diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 9a3ac13f02..b83e0308cd 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -552,6 +552,10 @@
+
-
-
+
-
@@ -627,6 +634,7 @@ import UsageProgressBar from './UsageProgressBar.vue' import AccountQuotaInfo from './AccountQuotaInfo.vue' import OpenAIQuotaResetCell from './OpenAIQuotaResetCell.vue' import GrokQuotaProbeCell from './GrokQuotaProbeCell.vue' +import OllamaCloudUsageCell from './OllamaCloudUsageCell.vue' // Module-level cache shared across all AccountUsageCell instances const _usageCache = new Map() diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index df4838fd31..b15259ef3e 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1635,6 +1635,12 @@ /> + +
props.account?.parent_account_id != null) +const handleOllamaCloudUsageUpdated = (state: OllamaCloudUsageState) => { + if (props.account) emit('updated', { ...props.account, ollama_cloud_usage: state }) +} + // Platform-specific hint for Base URL const baseUrlHint = computed(() => { if (!props.account) return t('admin.accounts.baseUrlHint') diff --git a/frontend/src/components/account/OllamaCloudUsageCell.vue b/frontend/src/components/account/OllamaCloudUsageCell.vue new file mode 100644 index 0000000000..201e271fe5 --- /dev/null +++ b/frontend/src/components/account/OllamaCloudUsageCell.vue @@ -0,0 +1,35 @@ + + + diff --git a/frontend/src/components/account/OllamaCloudUsageSettings.vue b/frontend/src/components/account/OllamaCloudUsageSettings.vue new file mode 100644 index 0000000000..7caaa173a0 --- /dev/null +++ b/frontend/src/components/account/OllamaCloudUsageSettings.vue @@ -0,0 +1,270 @@ +