diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index eac5392b0c..e164af0bec 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.139 +0.1.141 diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 844348c024..0d3bf88705 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -136,6 +136,7 @@ var DefaultBedrockModelMapping = map[string]string{ "claude-opus-4-1": "us.anthropic.claude-opus-4-1-20250805-v1:0", "claude-opus-4-20250514": "us.anthropic.claude-opus-4-20250514-v1:0", // Claude Sonnet + "claude-sonnet-5": "us.anthropic.claude-sonnet-5-v1", "claude-sonnet-4-6-thinking": "us.anthropic.claude-sonnet-4-6", "claude-sonnet-4-6": "us.anthropic.claude-sonnet-4-6", "claude-sonnet-4-5": "us.anthropic.claude-sonnet-4-5-20250929-v1:0", diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 4f4807130e..13d4c3fdbc 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -259,6 +259,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks, EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection, RewriteMessageCacheControl: settings.RewriteMessageCacheControl, + EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization, AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion, OpenAICodexUserAgent: settings.OpenAICodexUserAgent, MinCodexVersion: settings.MinCodexVersion, @@ -598,6 +599,7 @@ type UpdateSettingsRequest struct { ClaudeOAuthSystemPromptBlocks *string `json:"claude_oauth_system_prompt_blocks"` EnableAnthropicCacheTTL1hInjection *bool `json:"enable_anthropic_cache_ttl_1h_injection"` RewriteMessageCacheControl *bool `json:"rewrite_message_cache_control"` + EnableClientDatelineNormalization *bool `json:"enable_client_dateline_normalization"` AntigravityUserAgentVersion *string `json:"antigravity_user_agent_version"` OpenAICodexUserAgent *string `json:"openai_codex_user_agent"` @@ -1731,6 +1733,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.RewriteMessageCacheControl }(), + EnableClientDatelineNormalization: func() bool { + if req.EnableClientDatelineNormalization != nil { + return *req.EnableClientDatelineNormalization + } + return previousSettings.EnableClientDatelineNormalization + }(), AntigravityUserAgentVersion: func() string { if req.AntigravityUserAgentVersion != nil { return *req.AntigravityUserAgentVersion @@ -2143,6 +2151,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks, EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection, RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl, + EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization, AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion, OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent, MinCodexVersion: updatedSettings.MinCodexVersion, @@ -2644,6 +2653,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.RewriteMessageCacheControl != after.RewriteMessageCacheControl { changed = append(changed, "rewrite_message_cache_control") } + if before.EnableClientDatelineNormalization != after.EnableClientDatelineNormalization { + changed = append(changed, "enable_client_dateline_normalization") + } if before.AntigravityUserAgentVersion != after.AntigravityUserAgentVersion { changed = append(changed, "antigravity_user_agent_version") } diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 11d06331af..c7334f1bea 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -577,7 +577,7 @@ func AccountSummaryFromService(a *service.Account) *AccountSummary { } func usageLogFromServiceUser(l *service.UsageLog) UsageLog { - // 普通用户 DTO:严禁包含管理员字段(例如 account_rate_multiplier、ip_address、account)。 + // 普通用户 DTO:严禁包含管理员字段(例如 account_rate_multiplier、account、upstream_model)。 requestType := l.EffectiveRequestType() stream, openAIWSMode := service.ApplyLegacyRequestFields(requestType, l.Stream, l.OpenAIWSMode) requestedModel := l.RequestedModel @@ -594,7 +594,6 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog { ServiceTier: l.ServiceTier, ReasoningEffort: l.ReasoningEffort, InboundEndpoint: l.InboundEndpoint, - UpstreamEndpoint: l.UpstreamEndpoint, GroupID: l.GroupID, SubscriptionID: l.SubscriptionID, InputTokens: l.InputTokens, @@ -626,6 +625,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog { ImageSizeBreakdown: l.ImageSizeBreakdown, MediaType: l.MediaType, UserAgent: l.UserAgent, + IPAddress: l.IPAddress, CacheTTLOverridden: l.CacheTTLOverridden, BillingMode: l.BillingMode, CreatedAt: l.CreatedAt, @@ -637,7 +637,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog { } // UsageLogFromService converts a service UsageLog to DTO for regular users. -// It excludes Account details and IP address - users should not see these. +// It excludes admin-only account/upstream internals while keeping user billing and request metadata. func UsageLogFromService(l *service.UsageLog) *UsageLog { if l == nil { return nil @@ -652,8 +652,10 @@ func UsageLogFromServiceAdmin(l *service.UsageLog) *AdminUsageLog { if l == nil { return nil } + usageLog := usageLogFromServiceUser(l) + usageLog.UpstreamEndpoint = l.UpstreamEndpoint return &AdminUsageLog{ - UsageLog: usageLogFromServiceUser(l), + UsageLog: usageLog, UpstreamModel: l.UpstreamModel, ChannelID: l.ChannelID, ModelMappingChain: l.ModelMappingChain, diff --git a/backend/internal/handler/dto/mappers_usage_test.go b/backend/internal/handler/dto/mappers_usage_test.go index eca838b9af..6cbcbada56 100644 --- a/backend/internal/handler/dto/mappers_usage_test.go +++ b/backend/internal/handler/dto/mappers_usage_test.go @@ -95,8 +95,7 @@ func TestUsageLogFromService_IncludesServiceTierForUserAndAdmin(t *testing.T) { require.Equal(t, serviceTier, *userDTO.ServiceTier) require.NotNil(t, userDTO.InboundEndpoint) require.Equal(t, inboundEndpoint, *userDTO.InboundEndpoint) - require.NotNil(t, userDTO.UpstreamEndpoint) - require.Equal(t, upstreamEndpoint, *userDTO.UpstreamEndpoint) + require.Nil(t, userDTO.UpstreamEndpoint) require.NotNil(t, adminDTO.ServiceTier) require.Equal(t, serviceTier, *adminDTO.ServiceTier) require.NotNil(t, adminDTO.InboundEndpoint) @@ -133,6 +132,45 @@ func TestUsageLogFromService_UsesRequestedModelAndKeepsUpstreamAdminOnly(t *test require.Contains(t, string(adminJSON), `"upstream_model":"claude-sonnet-4-20250514"`) } +func TestUsageLogFromService_KeepsUserBillingAndIPWithoutAdminCostFields(t *testing.T) { + t.Parallel() + + ipAddress := "203.0.113.10" + accountRateMultiplier := 1.5 + accountStatsCost := 0.21 + log := &service.UsageLog{ + RequestID: "req_user_visible_billing", + Model: "gpt-5.4", + InputCost: 0.01, + OutputCost: 0.02, + CacheCreationCost: 0.03, + CacheReadCost: 0.04, + TotalCost: 0.10, + ActualCost: 0.08, + RateMultiplier: 0.8, + IPAddress: &ipAddress, + AccountRateMultiplier: &accountRateMultiplier, + AccountStatsCost: &accountStatsCost, + } + + userDTO := UsageLogFromService(log) + require.Equal(t, 0.01, userDTO.InputCost) + require.Equal(t, 0.02, userDTO.OutputCost) + require.Equal(t, 0.03, userDTO.CacheCreationCost) + require.Equal(t, 0.04, userDTO.CacheReadCost) + require.Equal(t, 0.10, userDTO.TotalCost) + require.Equal(t, 0.08, userDTO.ActualCost) + require.Equal(t, 0.8, userDTO.RateMultiplier) + require.NotNil(t, userDTO.IPAddress) + require.Equal(t, ipAddress, *userDTO.IPAddress) + + userJSON, err := json.Marshal(userDTO) + require.NoError(t, err) + require.NotContains(t, string(userJSON), "account_rate_multiplier") + require.NotContains(t, string(userJSON), "account_stats_cost") + require.NotContains(t, string(userJSON), "account_cost") +} + func TestUsageLogFromService_FallsBackToLegacyModelWhenRequestedModelMissing(t *testing.T) { t.Parallel() diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index e55e06b59d..e117d19ba3 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -186,6 +186,7 @@ type SystemSettings struct { ClaudeOAuthSystemPromptBlocks string `json:"claude_oauth_system_prompt_blocks"` EnableAnthropicCacheTTL1hInjection bool `json:"enable_anthropic_cache_ttl_1h_injection"` RewriteMessageCacheControl bool `json:"rewrite_message_cache_control"` + EnableClientDatelineNormalization bool `json:"enable_client_dateline_normalization"` AntigravityUserAgentVersion string `json:"antigravity_user_agent_version"` OpenAICodexUserAgent string `json:"openai_codex_user_agent"` diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 1832b6e490..d6bc0b80a8 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -485,6 +485,8 @@ type UsageLog struct { // User-Agent UserAgent *string `json:"user_agent"` + // IPAddress is visible to the owner of the usage record. + IPAddress *string `json:"ip_address,omitempty"` // Cache TTL Override 标记 CacheTTLOverridden bool `json:"cache_ttl_overridden"` @@ -520,7 +522,7 @@ type AdminUsageLog struct { // AccountStatsCost 自定义定价规则计算的账号统计费用(nil 表示使用默认公式) AccountStatsCost *float64 `json:"account_stats_cost,omitempty"` - // IPAddress 用户请求 IP(仅管理员可见) + // IPAddress 用户请求 IP IPAddress *string `json:"ip_address,omitempty"` // Account 最小账号信息(避免泄露敏感字段) diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go index 23bb62dd18..9d0f1d8fac 100644 --- a/backend/internal/handler/usage_handler.go +++ b/backend/internal/handler/usage_handler.go @@ -17,6 +17,33 @@ import ( "github.com/gin-gonic/gin" ) +type userUsageFilters struct { + Filters usagestats.UsageLogFilters + StartTime time.Time + EndTime time.Time +} + +type userModelStat struct { + Model string `json:"model"` + Requests int64 `json:"requests"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + CacheCreationTokens int64 `json:"cache_creation_tokens"` + CacheReadTokens int64 `json:"cache_read_tokens"` + TotalTokens int64 `json:"total_tokens"` + Cost float64 `json:"cost"` + ActualCost float64 `json:"actual_cost"` +} + +type userGroupStat struct { + GroupID int64 `json:"group_id"` + GroupName string `json:"group_name"` + Requests int64 `json:"requests"` + TotalTokens int64 `json:"total_tokens"` + Cost float64 `json:"cost"` + ActualCost float64 `json:"actual_cost"` +} + // UsageHandler handles usage-related requests type UsageHandler struct { usageService *service.UsageService @@ -40,41 +67,45 @@ func NewUsageHandler( } } -// List handles listing usage records with pagination -// GET /api/v1/usage -func (h *UsageHandler) List(c *gin.Context) { +func (h *UsageHandler) parseUserUsageFilters(c *gin.Context, requireRange bool) (*userUsageFilters, bool) { subject, ok := middleware2.GetAuthSubjectFromContext(c) if !ok { response.Unauthorized(c, "User not authenticated") - return + return nil, false } - page, pageSize := response.ParsePagination(c) - var apiKeyID int64 - if apiKeyIDStr := c.Query("api_key_id"); apiKeyIDStr != "" { + if apiKeyIDStr := strings.TrimSpace(c.Query("api_key_id")); apiKeyIDStr != "" { id, err := strconv.ParseInt(apiKeyIDStr, 10, 64) if err != nil { response.BadRequest(c, "Invalid api_key_id") - return + return nil, false + } + if h.apiKeyService == nil { + response.InternalError(c, "API key service not available") + return nil, false } - - // [Security Fix] Verify API Key ownership to prevent horizontal privilege escalation apiKey, err := h.apiKeyService.GetByID(c.Request.Context(), id) if err != nil { response.ErrorFrom(c, err) - return + return nil, false } if apiKey.UserID != subject.UserID { response.Forbidden(c, "Not authorized to access this API key's usage records") - return + return nil, false } - apiKeyID = id } - // Parse additional filters - model := c.Query("model") + var groupID int64 + if groupIDStr := strings.TrimSpace(c.Query("group_id")); groupIDStr != "" { + id, err := strconv.ParseInt(groupIDStr, 10, 64) + if err != nil { + response.BadRequest(c, "Invalid group_id") + return nil, false + } + groupID = id + } var requestType *int16 var stream *bool @@ -82,51 +113,119 @@ func (h *UsageHandler) List(c *gin.Context) { parsed, err := service.ParseUsageRequestType(requestTypeStr) if err != nil { response.BadRequest(c, err.Error()) - return + return nil, false } value := int16(parsed) requestType = &value - } else if streamStr := c.Query("stream"); streamStr != "" { + } else if streamStr := strings.TrimSpace(c.Query("stream")); streamStr != "" { val, err := strconv.ParseBool(streamStr) if err != nil { response.BadRequest(c, "Invalid stream value, use true or false") - return + return nil, false } stream = &val } var billingType *int8 - if billingTypeStr := c.Query("billing_type"); billingTypeStr != "" { + if billingTypeStr := strings.TrimSpace(c.Query("billing_type")); billingTypeStr != "" { val, err := strconv.ParseInt(billingTypeStr, 10, 8) if err != nil { response.BadRequest(c, "Invalid billing_type") - return + return nil, false } bt := int8(val) billingType = &bt } - // Parse date range - var startTime, endTime *time.Time - userTZ := c.Query("timezone") // Get user's timezone from request - if startDateStr := c.Query("start_date"); startDateStr != "" { + billingMode := strings.TrimSpace(c.Query("billing_mode")) + if billingMode != "" && !service.BillingMode(billingMode).IsValid() { + response.BadRequest(c, "Invalid billing_mode") + return nil, false + } + + userTZ := c.Query("timezone") + now := timezone.NowInUserLocation(userTZ) + var startTime, endTime time.Time + var startPtr, endPtr *time.Time + startDateStr := strings.TrimSpace(c.Query("start_date")) + endDateStr := strings.TrimSpace(c.Query("end_date")) + + if startDateStr != "" { t, err := timezone.ParseInUserLocation("2006-01-02", startDateStr, userTZ) if err != nil { response.BadRequest(c, "Invalid start_date format, use YYYY-MM-DD") - return + return nil, false } - startTime = &t + startTime = t + startPtr = &startTime } - - if endDateStr := c.Query("end_date"); endDateStr != "" { + if endDateStr != "" { t, err := timezone.ParseInUserLocation("2006-01-02", endDateStr, userTZ) if err != nil { response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") - return + return nil, false } - // Use half-open range [start, end), move to next calendar day start (DST-safe). - t = t.AddDate(0, 0, 1) - endTime = &t + endTime = t.AddDate(0, 0, 1) + endPtr = &endTime + } + + if requireRange { + if startPtr == nil { + switch c.DefaultQuery("period", "") { + case "today": + startTime = timezone.StartOfDayInUserLocation(now, userTZ) + case "week": + startTime = now.AddDate(0, 0, -7) + case "month": + startTime = now.AddDate(0, -1, 0) + default: + startTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, -7), userTZ) + } + startPtr = &startTime + } + if endPtr == nil { + if strings.TrimSpace(c.Query("period")) != "" { + endTime = now + } else { + endTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, 1), userTZ) + } + endPtr = &endTime + } + } + + return &userUsageFilters{ + Filters: usagestats.UsageLogFilters{ + UserID: subject.UserID, + APIKeyID: apiKeyID, + GroupID: groupID, + Model: strings.TrimSpace(c.Query("model")), + ModelFilterSource: usagestats.ModelSourceRequested, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + BillingMode: billingMode, + StartTime: startPtr, + EndTime: endPtr, + }, + StartTime: derefTime(startPtr), + EndTime: derefTime(endPtr), + }, true +} + +func derefTime(value *time.Time) time.Time { + if value == nil { + return time.Time{} + } + return *value +} + +// List handles listing usage records with pagination +// GET /api/v1/usage +func (h *UsageHandler) List(c *gin.Context) { + page, pageSize := response.ParsePagination(c) + parsed, ok := h.parseUserUsageFilters(c, false) + if !ok { + return } params := pagination.PaginationParams{ @@ -135,18 +234,8 @@ func (h *UsageHandler) List(c *gin.Context) { SortBy: c.DefaultQuery("sort_by", "created_at"), SortOrder: c.DefaultQuery("sort_order", "desc"), } - filters := usagestats.UsageLogFilters{ - UserID: subject.UserID, // Always filter by current user for security - APIKeyID: apiKeyID, - Model: model, - RequestType: requestType, - Stream: stream, - BillingType: billingType, - StartTime: startTime, - EndTime: endTime, - } - records, result, err := h.usageService.ListWithFilters(c.Request.Context(), params, filters) + records, result, err := h.usageService.ListWithFilters(c.Request.Context(), params, parsed.Filters) if err != nil { response.ErrorFrom(c, err) return @@ -303,122 +392,23 @@ func (h *UsageHandler) GetByID(c *gin.Context) { // Stats handles getting usage statistics // GET /api/v1/usage/stats func (h *UsageHandler) Stats(c *gin.Context) { - subject, ok := middleware2.GetAuthSubjectFromContext(c) + parsed, ok := h.parseUserUsageFilters(c, true) if !ok { - response.Unauthorized(c, "User not authenticated") return } - var apiKeyID int64 - if apiKeyIDStr := c.Query("api_key_id"); apiKeyIDStr != "" { - id, err := strconv.ParseInt(apiKeyIDStr, 10, 64) - if err != nil { - response.BadRequest(c, "Invalid api_key_id") - return - } - - // [Security Fix] Verify API Key ownership to prevent horizontal privilege escalation - apiKey, err := h.apiKeyService.GetByID(c.Request.Context(), id) - if err != nil { - response.NotFound(c, "API key not found") - return - } - if apiKey.UserID != subject.UserID { - response.Forbidden(c, "Not authorized to access this API key's statistics") - return - } - - apiKeyID = id - } - - // 获取时间范围参数 - userTZ := c.Query("timezone") // Get user's timezone from request - now := timezone.NowInUserLocation(userTZ) - var startTime, endTime time.Time - - // 优先使用 start_date 和 end_date 参数 - startDateStr := c.Query("start_date") - endDateStr := c.Query("end_date") - - if startDateStr != "" && endDateStr != "" { - // 使用自定义日期范围 - var err error - startTime, err = timezone.ParseInUserLocation("2006-01-02", startDateStr, userTZ) - if err != nil { - response.BadRequest(c, "Invalid start_date format, use YYYY-MM-DD") - return - } - endTime, err = timezone.ParseInUserLocation("2006-01-02", endDateStr, userTZ) - if err != nil { - response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD") - return - } - // 与 SQL 条件 created_at < end 对齐,使用次日 00:00 作为上边界(DST-safe)。 - endTime = endTime.AddDate(0, 0, 1) - } else { - // 使用 period 参数 - period := c.DefaultQuery("period", "today") - switch period { - case "today": - startTime = timezone.StartOfDayInUserLocation(now, userTZ) - case "week": - startTime = now.AddDate(0, 0, -7) - case "month": - startTime = now.AddDate(0, -1, 0) - default: - startTime = timezone.StartOfDayInUserLocation(now, userTZ) - } - endTime = now - } - - var stats *service.UsageStats - var err error - if apiKeyID > 0 { - stats, err = h.usageService.GetStatsByAPIKey(c.Request.Context(), apiKeyID, startTime, endTime) - } else { - stats, err = h.usageService.GetStatsByUser(c.Request.Context(), subject.UserID, startTime, endTime) - } + stats, err := h.usageService.GetStatsWithFilters(c.Request.Context(), parsed.Filters) if err != nil { response.ErrorFrom(c, err) return } + stats.TotalAccountCost = nil + stats.UpstreamEndpoints = nil + stats.EndpointPaths = nil response.Success(c, stats) } -// parseUserTimeRange parses start_date, end_date query parameters for user dashboard -// Uses user's timezone if provided, otherwise falls back to server timezone -func parseUserTimeRange(c *gin.Context) (time.Time, time.Time) { - userTZ := c.Query("timezone") // Get user's timezone from request - now := timezone.NowInUserLocation(userTZ) - startDate := c.Query("start_date") - endDate := c.Query("end_date") - - var startTime, endTime time.Time - - if startDate != "" { - if t, err := timezone.ParseInUserLocation("2006-01-02", startDate, userTZ); err == nil { - startTime = t - } else { - startTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, -7), userTZ) - } - } else { - startTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, -7), userTZ) - } - - if endDate != "" { - if t, err := timezone.ParseInUserLocation("2006-01-02", endDate, userTZ); err == nil { - endTime = t.Add(24 * time.Hour) // Include the end date - } else { - endTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, 1), userTZ) - } - } else { - endTime = timezone.StartOfDayInUserLocation(now.AddDate(0, 0, 1), userTZ) - } - - return startTime, endTime -} - const ( defaultAPIKeyDailyUsageDays = 30 maxAPIKeyDailyUsageDays = 90 @@ -463,16 +453,13 @@ func (h *UsageHandler) DashboardStats(c *gin.Context) { // DashboardTrend handles getting user usage trend data // GET /api/v1/usage/dashboard/trend func (h *UsageHandler) DashboardTrend(c *gin.Context) { - subject, ok := middleware2.GetAuthSubjectFromContext(c) + parsed, ok := h.parseUserUsageFilters(c, true) if !ok { - response.Unauthorized(c, "User not authenticated") return } - - startTime, endTime := parseUserTimeRange(c) granularity := c.DefaultQuery("granularity", "day") - trend, err := h.usageService.GetUserUsageTrendByUserID(c.Request.Context(), subject.UserID, startTime, endTime, granularity) + trend, err := h.usageService.GetUsageTrendWithFilters(c.Request.Context(), parsed.StartTime, parsed.EndTime, granularity, parsed.Filters) if err != nil { response.ErrorFrom(c, err) return @@ -480,8 +467,8 @@ func (h *UsageHandler) DashboardTrend(c *gin.Context) { response.Success(c, gin.H{ "trend": trend, - "start_date": startTime.Format("2006-01-02"), - "end_date": endTime.Add(-24 * time.Hour).Format("2006-01-02"), + "start_date": parsed.StartTime.Format("2006-01-02"), + "end_date": parsed.EndTime.Add(-24 * time.Hour).Format("2006-01-02"), "granularity": granularity, }) } @@ -489,27 +476,136 @@ func (h *UsageHandler) DashboardTrend(c *gin.Context) { // DashboardModels handles getting user model usage statistics // GET /api/v1/usage/dashboard/models func (h *UsageHandler) DashboardModels(c *gin.Context) { - subject, ok := middleware2.GetAuthSubjectFromContext(c) + parsed, ok := h.parseUserUsageFilters(c, true) if !ok { - response.Unauthorized(c, "User not authenticated") return } - startTime, endTime := parseUserTimeRange(c) + modelSource := strings.TrimSpace(c.Query("model_source")) + if modelSource != "" && modelSource != usagestats.ModelSourceRequested { + response.BadRequest(c, "Invalid model_source, user usage only supports requested") + return + } - stats, err := h.usageService.GetUserModelStats(c.Request.Context(), subject.UserID, startTime, endTime) + stats, err := h.usageService.GetModelStatsWithFiltersBySource(c.Request.Context(), parsed.StartTime, parsed.EndTime, parsed.Filters, usagestats.ModelSourceRequested) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, gin.H{ - "models": stats, - "start_date": startTime.Format("2006-01-02"), - "end_date": endTime.Add(-24 * time.Hour).Format("2006-01-02"), + "models": userModelStatsFromUsageStats(stats), + "start_date": parsed.StartTime.Format("2006-01-02"), + "end_date": parsed.EndTime.Add(-24 * time.Hour).Format("2006-01-02"), }) } +// DashboardSnapshotV2 returns usage-page chart data scoped to the current user. +// GET /api/v1/usage/dashboard/snapshot-v2 +func (h *UsageHandler) DashboardSnapshotV2(c *gin.Context) { + parsed, ok := h.parseUserUsageFilters(c, true) + if !ok { + return + } + + granularity := strings.TrimSpace(c.DefaultQuery("granularity", "day")) + if granularity != "hour" { + granularity = "day" + } + includeTrend, ok := parseBoolQueryWithDefault(c, "include_trend", true) + if !ok { + return + } + includeModels, ok := parseBoolQueryWithDefault(c, "include_model_stats", true) + if !ok { + return + } + includeGroups, ok := parseBoolQueryWithDefault(c, "include_group_stats", false) + if !ok { + return + } + + resp := gin.H{ + "generated_at": time.Now().UTC().Format(time.RFC3339), + "start_date": parsed.StartTime.Format("2006-01-02"), + "end_date": parsed.EndTime.Add(-24 * time.Hour).Format("2006-01-02"), + "granularity": granularity, + } + + if includeTrend { + trend, err := h.usageService.GetUsageTrendWithFilters(c.Request.Context(), parsed.StartTime, parsed.EndTime, granularity, parsed.Filters) + if err != nil { + response.ErrorFrom(c, err) + return + } + resp["trend"] = trend + } + if includeModels { + models, err := h.usageService.GetModelStatsWithFiltersBySource(c.Request.Context(), parsed.StartTime, parsed.EndTime, parsed.Filters, usagestats.ModelSourceRequested) + if err != nil { + response.ErrorFrom(c, err) + return + } + resp["models"] = userModelStatsFromUsageStats(models) + } + if includeGroups { + groups, err := h.usageService.GetGroupStatsWithFilters(c.Request.Context(), parsed.StartTime, parsed.EndTime, parsed.Filters) + if err != nil { + response.ErrorFrom(c, err) + return + } + resp["groups"] = userGroupStatsFromUsageStats(groups) + } + + response.Success(c, resp) +} + +func userModelStatsFromUsageStats(stats []usagestats.ModelStat) []userModelStat { + out := make([]userModelStat, 0, len(stats)) + for _, stat := range stats { + out = append(out, userModelStat{ + Model: stat.Model, + Requests: stat.Requests, + InputTokens: stat.InputTokens, + OutputTokens: stat.OutputTokens, + CacheCreationTokens: stat.CacheCreationTokens, + CacheReadTokens: stat.CacheReadTokens, + TotalTokens: stat.TotalTokens, + Cost: stat.Cost, + ActualCost: stat.ActualCost, + }) + } + return out +} + +func userGroupStatsFromUsageStats(stats []usagestats.GroupStat) []userGroupStat { + out := make([]userGroupStat, 0, len(stats)) + for _, stat := range stats { + out = append(out, userGroupStat{ + GroupID: stat.GroupID, + GroupName: stat.GroupName, + Requests: stat.Requests, + TotalTokens: stat.TotalTokens, + Cost: stat.Cost, + ActualCost: stat.ActualCost, + }) + } + return out +} + +func parseBoolQueryWithDefault(c *gin.Context, key string, fallback bool) (bool, bool) { + raw := c.Query(key) + if strings.TrimSpace(raw) == "" { + return fallback, true + } + parsed, err := strconv.ParseBool(raw) + if err != nil { + response.BadRequest(c, "Invalid "+key+" value, use true or false") + return false, false + } + return parsed, true +} + // BatchAPIKeysUsageRequest represents the request for batch API keys usage type BatchAPIKeysUsageRequest struct { APIKeyIDs []int64 `json:"api_key_ids" binding:"required"` diff --git a/backend/internal/handler/usage_handler_request_type_test.go b/backend/internal/handler/usage_handler_request_type_test.go index ed08c5a81a..1dcb1b83a4 100644 --- a/backend/internal/handler/usage_handler_request_type_test.go +++ b/backend/internal/handler/usage_handler_request_type_test.go @@ -5,6 +5,7 @@ import ( "net/http" "net/http/httptest" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" @@ -16,21 +17,67 @@ import ( type userUsageRepoCapture struct { service.UsageLogRepository - listParams pagination.PaginationParams - listFilters usagestats.UsageLogFilters + listParams pagination.PaginationParams + listFilters usagestats.UsageLogFilters + statsFilters usagestats.UsageLogFilters + trendFilters usagestats.UsageLogFilters + groupFilters usagestats.UsageLogFilters + listRows []service.UsageLog + stats *usagestats.UsageStats + modelStats []usagestats.ModelStat + groupStats []usagestats.GroupStat } func (s *userUsageRepoCapture) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters usagestats.UsageLogFilters) ([]service.UsageLog, *pagination.PaginationResult, error) { s.listParams = params s.listFilters = filters - return []service.UsageLog{}, &pagination.PaginationResult{ - Total: 0, + return s.listRows, &pagination.PaginationResult{ + Total: int64(len(s.listRows)), Page: params.Page, PageSize: params.PageSize, Pages: 0, }, nil } +func (s *userUsageRepoCapture) GetStatsWithFilters(ctx context.Context, filters usagestats.UsageLogFilters) (*usagestats.UsageStats, error) { + s.statsFilters = filters + if s.stats != nil { + return s.stats, nil + } + return &usagestats.UsageStats{}, nil +} + +func (s *userUsageRepoCapture) GetUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) ([]usagestats.TrendDataPoint, error) { + s.trendFilters = usagestats.UsageLogFilters{ + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + Model: model, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + } + return []usagestats.TrendDataPoint{}, nil +} + +func (s *userUsageRepoCapture) GetModelStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) ([]usagestats.ModelStat, error) { + return s.modelStats, nil +} + +func (s *userUsageRepoCapture) GetGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) ([]usagestats.GroupStat, error) { + s.groupFilters = usagestats.UsageLogFilters{ + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + } + return s.groupStats, nil +} + func newUserUsageRequestTypeTestRouter(repo *userUsageRepoCapture) *gin.Engine { gin.SetMode(gin.TestMode) usageSvc := service.NewUsageService(repo, nil, nil, nil) @@ -41,6 +88,9 @@ func newUserUsageRequestTypeTestRouter(repo *userUsageRepoCapture) *gin.Engine { c.Next() }) router.GET("/usage", handler.List) + router.GET("/usage/stats", handler.Stats) + router.GET("/usage/dashboard/models", handler.DashboardModels) + router.GET("/usage/dashboard/snapshot-v2", handler.DashboardSnapshotV2) return router } @@ -80,3 +130,197 @@ func TestUserUsageListInvalidStream(t *testing.T) { require.Equal(t, http.StatusBadRequest, rec.Code) } + +func TestUserUsageListAdvancedFilters(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage?group_id=7&model=gpt-5&billing_type=1&billing_mode=image&start_date=2026-03-01&end_date=2026-03-02", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, int64(42), repo.listFilters.UserID) + require.Equal(t, int64(7), repo.listFilters.GroupID) + require.Equal(t, "gpt-5", repo.listFilters.Model) + require.Equal(t, usagestats.ModelSourceRequested, repo.listFilters.ModelFilterSource) + require.NotNil(t, repo.listFilters.BillingType) + require.Equal(t, int8(1), *repo.listFilters.BillingType) + require.Equal(t, "image", repo.listFilters.BillingMode) + require.NotNil(t, repo.listFilters.StartTime) + require.NotNil(t, repo.listFilters.EndTime) +} + +func TestUserUsageListInvalidBillingMode(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage?billing_mode=bad", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) +} + +func TestUserUsageListKeepsUserBillingAndIPWithoutAdminCostFields(t *testing.T) { + ipAddress := "203.0.113.10" + upstreamModel := "upstream-private-model" + billingTier := "internal-tier" + channelID := int64(99) + accountRateMultiplier := 1.7 + accountStatsCost := 0.12 + repo := &userUsageRepoCapture{ + listRows: []service.UsageLog{{ + ID: 1, + UserID: 42, + APIKeyID: 7, + AccountID: 5, + RequestID: "req_user_billing", + Model: "gpt-5", + InputCost: 0.01, + OutputCost: 0.02, + CacheCreationCost: 0.03, + CacheReadCost: 0.04, + TotalCost: 0.10, + ActualCost: 0.08, + RateMultiplier: 0.8, + IPAddress: &ipAddress, + UpstreamModel: &upstreamModel, + BillingTier: &billingTier, + ChannelID: &channelID, + AccountRateMultiplier: &accountRateMultiplier, + AccountStatsCost: &accountStatsCost, + }}, + } + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + body := rec.Body.String() + require.Contains(t, body, `"input_cost":0.01`) + require.Contains(t, body, `"output_cost":0.02`) + require.Contains(t, body, `"cache_creation_cost":0.03`) + require.Contains(t, body, `"cache_read_cost":0.04`) + require.Contains(t, body, `"total_cost":0.1`) + require.Contains(t, body, `"actual_cost":0.08`) + require.Contains(t, body, `"rate_multiplier":0.8`) + require.Contains(t, body, `"ip_address":"203.0.113.10"`) + require.NotContains(t, body, "upstream_endpoint") + require.NotContains(t, body, "account_rate_multiplier") + require.NotContains(t, body, "account_stats_cost") + require.NotContains(t, body, "upstream_model") + require.NotContains(t, body, "billing_tier") + require.NotContains(t, body, "channel_id") + require.NotContains(t, body, `"account":`) +} + +func TestUserUsageStatsUsesScopedFilters(t *testing.T) { + accountCost := 0.12 + repo := &userUsageRepoCapture{ + stats: &usagestats.UsageStats{ + TotalCost: 0.10, + TotalActualCost: 0.08, + TotalAccountCost: &accountCost, + UpstreamEndpoints: []usagestats.EndpointStat{{ + Endpoint: "/v1/responses", + }}, + EndpointPaths: []usagestats.EndpointStat{{ + Endpoint: "/v1/chat/completions -> /v1/responses", + }}, + }, + } + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage/stats?group_id=9&request_type=sync&billing_mode=token&start_date=2026-03-01&end_date=2026-03-02", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, int64(42), repo.statsFilters.UserID) + require.Equal(t, int64(9), repo.statsFilters.GroupID) + require.Equal(t, usagestats.ModelSourceRequested, repo.statsFilters.ModelFilterSource) + require.NotNil(t, repo.statsFilters.RequestType) + require.Equal(t, int16(service.RequestTypeSync), *repo.statsFilters.RequestType) + require.Equal(t, "token", repo.statsFilters.BillingMode) + require.Contains(t, rec.Body.String(), `"total_cost":0.1`) + require.Contains(t, rec.Body.String(), `"total_actual_cost":0.08`) + require.NotContains(t, rec.Body.String(), "total_account_cost") + require.NotContains(t, rec.Body.String(), "upstream_endpoints") + require.NotContains(t, rec.Body.String(), "endpoint_paths") +} + +func TestUserUsageDashboardModelsOmitsAccountCost(t *testing.T) { + repo := &userUsageRepoCapture{ + modelStats: []usagestats.ModelStat{{ + Model: "gpt-5", + Requests: 2, + TotalTokens: 30, + Cost: 0.10, + ActualCost: 0.08, + AccountCost: 0.07, + }}, + } + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage/dashboard/models?start_date=2026-03-01&end_date=2026-03-02", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + body := rec.Body.String() + require.Contains(t, body, `"cost":0.1`) + require.Contains(t, body, `"actual_cost":0.08`) + require.NotContains(t, body, "account_cost") +} + +func TestUserUsageDashboardModelsRejectsAdminModelSources(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage/dashboard/models?model_source=upstream&start_date=2026-03-01&end_date=2026-03-02", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) +} + +func TestUserUsageSnapshotUsesScopedFilters(t *testing.T) { + repo := &userUsageRepoCapture{ + modelStats: []usagestats.ModelStat{{Model: "gpt-5", AccountCost: 0.07}}, + groupStats: []usagestats.GroupStat{{GroupID: 1, GroupName: "default", AccountCost: 0.06}}, + } + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage/dashboard/snapshot-v2?include_trend=true&include_model_stats=true&include_group_stats=true&group_id=11&request_type=stream&start_date=2026-03-01&end_date=2026-03-02", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, int64(42), repo.trendFilters.UserID) + require.Equal(t, int64(11), repo.trendFilters.GroupID) + require.NotNil(t, repo.trendFilters.RequestType) + require.Equal(t, int16(service.RequestTypeStream), *repo.trendFilters.RequestType) + require.Equal(t, int64(42), repo.groupFilters.UserID) + require.Equal(t, int64(11), repo.groupFilters.GroupID) + require.NotContains(t, rec.Body.String(), "account_cost") +} + +func TestUserUsageSnapshotRejectsInvalidIncludeFlags(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + for _, query := range []string{ + "include_trend=bad", + "include_model_stats=bad", + "include_group_stats=bad", + } { + req := httptest.NewRequest(http.MethodGet, "/usage/dashboard/snapshot-v2?start_date=2026-03-01&end_date=2026-03-02&"+query, nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code, query) + } +} diff --git a/backend/internal/pkg/anthropicfp/dateline.go b/backend/internal/pkg/anthropicfp/dateline.go new file mode 100644 index 0000000000..ab07765582 --- /dev/null +++ b/backend/internal/pkg/anthropicfp/dateline.go @@ -0,0 +1,288 @@ +// Package anthropicfp provides pure helpers for suppressing client-side +// fingerprints that would otherwise be visible to upstream Anthropic when a +// forwarding gateway sits between the client and api.anthropic.com. +// +// Currently exposes NormalizeDateline: it rewrites the "Today's date is +// YYYY-MM-DD." sentence inside a request body back to a canonical ASCII form, +// erasing three bits of steganographic signal (four apostrophe code points and +// a date-separator variant) that some clients embed in that sentence when +// they detect a non-official base URL. +package anthropicfp + +import ( + "fmt" + "regexp" + "sort" + "strings" + + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +// datelineRegexes matches the fingerprinted sentence with any of the four +// apostrophe code points seen in the wild and either separator. Two regexes +// are used because Go's RE2-based regexp package does not support +// backreferences: matching `-` and `/` in two passes keeps the two separators +// inside YYYY?MM?DD forced to agree, so mixed-separator strings like +// "Today's date is 2026-07/01." never match. This is what filters out +// user-authored prose like "Today is foo." or "His date is 2026-06-30." from +// being touched. +var ( + datelineRegexHyphen = regexp.MustCompile(`Today(['’ʼʹ])s date is (\d{4})-(\d{2})-(\d{2})\.`) + datelineRegexSlash = regexp.MustCompile(`Today(['’ʼʹ])s date is (\d{4})/(\d{2})/(\d{2})\.`) +) + +// systemReminderRegex matches a block. The dateline lives in +// this block once the conversation has advanced past the first turn (system +// prompt caching hides the top-level system block for subsequent turns), so +// the messages[].content[] scan is confined to what lives inside these tags. +var systemReminderRegex = regexp.MustCompile(`(?s).*?`) + +// DatelineHit records what a single rewrite normalized, for observability. +type DatelineHit struct { + // ApostropheVariant is one of "ascii" (U+0027), "u2019", "u02bc", "u02b9". + ApostropheVariant string + // DateSeparator is either "-" or "/" as seen before normalization. + DateSeparator string +} + +// canonicalize returns the canonical form of a matched dateline sentence. +// The output always uses ASCII apostrophe and hyphen separators. +func canonicalize(year, month, day string) string { + return fmt.Sprintf("Today's date is %s-%s-%s.", year, month, day) +} + +func apostropheVariant(r rune) string { + switch r { + case '’': + return "u2019" + case 'ʼ': + return "u02bc" + case 'ʹ': + return "u02b9" + default: + return "ascii" + } +} + +type datelineMatch struct { + start, end int + apoRune rune + sep string + year, month, day string +} + +func collectMatches(text string, re *regexp.Regexp, sep string) []datelineMatch { + locs := re.FindAllStringSubmatchIndex(text, -1) + if len(locs) == 0 { + return nil + } + out := make([]datelineMatch, 0, len(locs)) + for _, m := range locs { + var apoRune rune + for _, r := range text[m[2]:m[3]] { + apoRune = r + break + } + out = append(out, datelineMatch{ + start: m[0], + end: m[1], + apoRune: apoRune, + sep: sep, + year: text[m[4]:m[5]], + month: text[m[6]:m[7]], + day: text[m[8]:m[9]], + }) + } + return out +} + +// NormalizeText replaces every fingerprinted dateline sentence in text with +// its canonical form. It returns the possibly-rewritten text and the list of +// hits observed. When no match is found the original string is returned +// verbatim (byte-identical), and the returned hit slice is nil. +func NormalizeText(text string) (string, []DatelineHit) { + if !strings.Contains(text, "date is ") { + return text, nil + } + matches := collectMatches(text, datelineRegexHyphen, "-") + matches = append(matches, collectMatches(text, datelineRegexSlash, "/")...) + if len(matches) == 0 { + return text, nil + } + sort.Slice(matches, func(i, j int) bool { return matches[i].start < matches[j].start }) + + var b strings.Builder + b.Grow(len(text)) + prev := 0 + hits := make([]DatelineHit, 0, len(matches)) + changed := false + for _, m := range matches { + full := text[m.start:m.end] + canonical := canonicalize(m.year, m.month, m.day) + if canonical == full { + // Already canonical: no rewrite, no hit. + continue + } + _, _ = b.WriteString(text[prev:m.start]) + _, _ = b.WriteString(canonical) + prev = m.end + changed = true + hits = append(hits, DatelineHit{ + ApostropheVariant: apostropheVariant(m.apoRune), + DateSeparator: m.sep, + }) + } + if !changed { + return text, nil + } + _, _ = b.WriteString(text[prev:]) + return b.String(), hits +} + +// normalizeSystemReminderScopedText scans only the blocks +// inside text and normalizes datelines inside them. Text outside the blocks is +// preserved byte-for-byte, so user prose, tool_result content, code blocks, +// or shell commands that happen to contain an apostrophe or a slash date are +// never touched. +func normalizeSystemReminderScopedText(text string) (string, []DatelineHit) { + if !strings.Contains(text, "") { + return text, nil + } + locs := systemReminderRegex.FindAllStringIndex(text, -1) + if len(locs) == 0 { + return text, nil + } + var b strings.Builder + b.Grow(len(text)) + prev := 0 + var hits []DatelineHit + changed := false + for _, loc := range locs { + _, _ = b.WriteString(text[prev:loc[0]]) + block := text[loc[0]:loc[1]] + normalized, blockHits := NormalizeText(block) + if normalized != block { + changed = true + } + _, _ = b.WriteString(normalized) + hits = append(hits, blockHits...) + prev = loc[1] + } + if !changed { + return text, nil + } + _, _ = b.WriteString(text[prev:]) + return b.String(), hits +} + +// NormalizeDateline scans an Anthropic /v1/messages request body and rewrites +// every fingerprinted dateline sentence back to its canonical ASCII form. +// +// Scope (mirroring where genuine clients place the sentence): +// - `system` string, or `.text` field of each text-typed block in `system`. +// - Text bodies inside `messages[i].content` — but ONLY the substrings that +// appear inside `...` tags. Free user +// prose, tool_use.input, tool_result.content, and other block types are +// never scanned, guaranteeing that legitimate text like a code block, a +// shell command, or a chat message that mentions today's date is never +// accidentally rewritten. +// +// The function is a pure transform: it never modifies the input slice, and if +// no rewrite is needed it returns the original slice (identity), a nil hit +// slice, and changed=false. +func NormalizeDateline(body []byte) ([]byte, []DatelineHit, bool) { + if len(body) == 0 { + return body, nil, false + } + out := body + var hits []DatelineHit + changed := false + + sys := gjson.GetBytes(out, "system") + if sys.Exists() { + switch { + case sys.Type == gjson.String: + normalized, sysHits := NormalizeText(sys.String()) + if normalized != sys.String() { + if next, err := sjson.SetBytes(out, "system", normalized); err == nil { + out = next + changed = true + hits = append(hits, sysHits...) + } + } + case sys.IsArray(): + idx := 0 + sys.ForEach(func(_, item gjson.Result) bool { + if item.Get("type").String() == "text" { + t := item.Get("text") + if t.Exists() && t.Type == gjson.String { + normalized, textHits := NormalizeText(t.String()) + if normalized != t.String() { + path := fmt.Sprintf("system.%d.text", idx) + if next, err := sjson.SetBytes(out, path, normalized); err == nil { + out = next + changed = true + hits = append(hits, textHits...) + } + } + } + } + idx++ + return true + }) + } + } + + messages := gjson.GetBytes(out, "messages") + if messages.IsArray() { + msgIdx := -1 + messages.ForEach(func(_, msg gjson.Result) bool { + msgIdx++ + content := msg.Get("content") + if !content.Exists() { + return true + } + switch { + case content.Type == gjson.String: + normalized, contentHits := normalizeSystemReminderScopedText(content.String()) + if normalized != content.String() { + path := fmt.Sprintf("messages.%d.content", msgIdx) + if next, err := sjson.SetBytes(out, path, normalized); err == nil { + out = next + changed = true + hits = append(hits, contentHits...) + } + } + case content.IsArray(): + contentIdx := -1 + content.ForEach(func(_, block gjson.Result) bool { + contentIdx++ + if block.Get("type").String() != "text" { + return true + } + t := block.Get("text") + if !t.Exists() || t.Type != gjson.String { + return true + } + normalized, textHits := normalizeSystemReminderScopedText(t.String()) + if normalized != t.String() { + path := fmt.Sprintf("messages.%d.content.%d.text", msgIdx, contentIdx) + if next, err := sjson.SetBytes(out, path, normalized); err == nil { + out = next + changed = true + hits = append(hits, textHits...) + } + } + return true + }) + } + return true + }) + } + + if !changed { + return body, nil, false + } + return out, hits, true +} diff --git a/backend/internal/pkg/anthropicfp/dateline_test.go b/backend/internal/pkg/anthropicfp/dateline_test.go new file mode 100644 index 0000000000..0c1abcd5c4 --- /dev/null +++ b/backend/internal/pkg/anthropicfp/dateline_test.go @@ -0,0 +1,248 @@ +package anthropicfp + +import ( + "bytes" + "strings" + "testing" +) + +func TestNormalizeText_ASCIIHyphenIsIdentity(t *testing.T) { + in := "Today's date is 2026-07-01." + out, hits := NormalizeText(in) + if out != in { + t.Fatalf("canonical form should be returned identity, got %q", out) + } + if len(hits) != 0 { + t.Fatalf("no hits expected on canonical input, got %d", len(hits)) + } +} + +func TestNormalizeText_SlashSeparatorASCIIApostrophe(t *testing.T) { + in := "Today's date is 2026/07/01." + out, hits := NormalizeText(in) + want := "Today's date is 2026-07-01." + if out != want { + t.Fatalf("want %q, got %q", want, out) + } + if len(hits) != 1 || hits[0].ApostropheVariant != "ascii" || hits[0].DateSeparator != "/" { + t.Fatalf("unexpected hit: %+v", hits) + } +} + +func TestNormalizeText_U2019Apostrophe(t *testing.T) { + in := "Today’s date is 2026-07-01." + out, hits := NormalizeText(in) + want := "Today's date is 2026-07-01." + if out != want { + t.Fatalf("want %q, got %q", want, out) + } + if len(hits) != 1 || hits[0].ApostropheVariant != "u2019" || hits[0].DateSeparator != "-" { + t.Fatalf("unexpected hit: %+v", hits) + } +} + +func TestNormalizeText_U02BCApostropheWithSlash(t *testing.T) { + in := "Todayʼs date is 2026/07/01." + out, hits := NormalizeText(in) + want := "Today's date is 2026-07-01." + if out != want { + t.Fatalf("want %q, got %q", want, out) + } + if len(hits) != 1 || hits[0].ApostropheVariant != "u02bc" || hits[0].DateSeparator != "/" { + t.Fatalf("unexpected hit: %+v", hits) + } +} + +func TestNormalizeText_U02B9Apostrophe(t *testing.T) { + in := "Todayʹs date is 2026/07/01." + out, hits := NormalizeText(in) + want := "Today's date is 2026-07-01." + if out != want { + t.Fatalf("want %q, got %q", want, out) + } + if len(hits) != 1 || hits[0].ApostropheVariant != "u02b9" || hits[0].DateSeparator != "/" { + t.Fatalf("unexpected hit: %+v", hits) + } +} + +func TestNormalizeText_MixedSeparatorsNoMatch(t *testing.T) { + // backreference \3 forces the two separators to agree + in := "Today's date is 2026-07/01." + out, hits := NormalizeText(in) + if out != in { + t.Fatalf("mixed-separator input should not be matched, got %q", out) + } + if len(hits) != 0 { + t.Fatalf("expected no hits, got %d", len(hits)) + } +} + +func TestNormalizeText_NegativeLookalike(t *testing.T) { + cases := []string{ + "Today is a great day.", + "His date is 2026-07-01.", + "Yesterday's date was 2026-06-30.", + "'s date is 2026-07-01.", + } + for _, c := range cases { + out, hits := NormalizeText(c) + if out != c { + t.Fatalf("input %q should not be modified, got %q", c, out) + } + if len(hits) != 0 { + t.Fatalf("input %q should produce no hits", c) + } + } +} + +func TestNormalizeText_Idempotent(t *testing.T) { + in := "Today’s date is 2026/07/01." + out1, _ := NormalizeText(in) + out2, hits2 := NormalizeText(out1) + if out1 != out2 { + t.Fatalf("normalization not idempotent: %q vs %q", out1, out2) + } + if len(hits2) != 0 { + t.Fatalf("second pass should produce no hits, got %d", len(hits2)) + } +} + +func TestNormalizeText_MultipleOccurrences(t *testing.T) { + in := "First line.\nToday’s date is 2026/07/01.\nMore text.\nTodayʼs date is 2026-07-01.\nEnd." + out, hits := NormalizeText(in) + if strings.Count(out, "Today's date is") != 2 { + t.Fatalf("expected two canonicalized sentences, got: %q", out) + } + if strings.Contains(out, "’") || strings.Contains(out, "ʼ") || strings.Contains(out, "2026/07/01") { + t.Fatalf("fingerprint characters must be gone, got: %q", out) + } + if len(hits) != 2 { + t.Fatalf("expected 2 hits, got %d", len(hits)) + } +} + +func TestNormalizeDateline_SystemString(t *testing.T) { + body := []byte(`{"system":"You are helpful.\nToday’s date is 2026/07/01.\nBe brief.","messages":[]}`) + out, hits, changed := NormalizeDateline(body) + if !changed { + t.Fatalf("expected changed=true") + } + if len(hits) != 1 { + t.Fatalf("expected 1 hit, got %d", len(hits)) + } + if !bytes.Contains(out, []byte("Today's date is 2026-07-01.")) { + t.Fatalf("output missing canonical dateline: %s", string(out)) + } + if bytes.Contains(out, []byte("2026/07/01")) { + t.Fatalf("output should not contain slash date: %s", string(out)) + } +} + +func TestNormalizeDateline_SystemBlocksArray(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"You are helpful."},{"type":"text","text":"Todayʼs date is 2026/07/01."}],"messages":[]}`) + out, hits, changed := NormalizeDateline(body) + if !changed { + t.Fatalf("expected changed=true") + } + if len(hits) != 1 || hits[0].ApostropheVariant != "u02bc" { + t.Fatalf("unexpected hits: %+v", hits) + } + if !bytes.Contains(out, []byte("Today's date is 2026-07-01.")) { + t.Fatalf("output missing canonical dateline: %s", string(out)) + } +} + +func TestNormalizeDateline_MessagesContentStringInSystemReminder(t *testing.T) { + body := []byte(`{"messages":[{"role":"user","content":"\n# currentDate\nToday’s date is 2026/07/01.\n\nHello, please help."}]}`) + out, hits, changed := NormalizeDateline(body) + if !changed { + t.Fatalf("expected changed=true") + } + if len(hits) != 1 { + t.Fatalf("expected 1 hit, got %d", len(hits)) + } + if !bytes.Contains(out, []byte("Today's date is 2026-07-01.")) { + t.Fatalf("canonical dateline missing: %s", string(out)) + } +} + +func TestNormalizeDateline_MessagesContentBlocksInSystemReminder(t *testing.T) { + body := []byte(`{"messages":[{"role":"user","content":[{"type":"text","text":"\nToday’s date is 2026/07/01.\n"},{"type":"text","text":"do X"}]}]}`) + out, hits, changed := NormalizeDateline(body) + if !changed { + t.Fatalf("expected changed=true") + } + if len(hits) != 1 { + t.Fatalf("expected 1 hit, got %d", len(hits)) + } + if !bytes.Contains(out, []byte("Today's date is 2026-07-01.")) { + t.Fatalf("canonical dateline missing: %s", string(out)) + } +} + +func TestNormalizeDateline_LeavesOutOfScopeUntouched(t *testing.T) { + // User prose outside that mentions today's date must + // not be modified. tool_use.input / tool_result.content are never scanned. + body := []byte(`{"messages":[` + + `{"role":"user","content":"Today’s date is 2026/07/01. Please help."},` + + `{"role":"assistant","content":[{"type":"tool_use","id":"x","name":"y","input":{"note":"Today’s date is 2026/07/01."}}]},` + + `{"role":"user","content":[{"type":"tool_result","tool_use_id":"x","content":"log: Today’s date is 2026/07/01."}]}` + + `]}`) + out, hits, changed := NormalizeDateline(body) + if changed { + t.Fatalf("expected changed=false, hits=%v out=%s", hits, string(out)) + } + if !bytes.Equal(out, body) { + t.Fatalf("output should equal input byte-for-byte") + } +} + +func TestNormalizeDateline_Idempotent(t *testing.T) { + body := []byte(`{"messages":[{"role":"user","content":"\nToday’s date is 2026/07/01.\n"}]}`) + first, _, changed1 := NormalizeDateline(body) + if !changed1 { + t.Fatalf("expected first pass to change body") + } + second, _, changed2 := NormalizeDateline(first) + if changed2 { + t.Fatalf("second pass should not report changes") + } + if !bytes.Equal(first, second) { + t.Fatalf("second pass diverged: %s vs %s", string(first), string(second)) + } +} + +func TestNormalizeDateline_EmptyBody(t *testing.T) { + out, hits, changed := NormalizeDateline(nil) + if changed || out != nil || hits != nil { + t.Fatalf("empty body should be no-op") + } +} + +func TestNormalizeDateline_NoDateline(t *testing.T) { + body := []byte(`{"messages":[{"role":"user","content":"hello"}],"system":"just a system prompt"}`) + out, hits, changed := NormalizeDateline(body) + if changed || len(hits) != 0 { + t.Fatalf("expected no changes; changed=%v hits=%v", changed, hits) + } + if &out[0] != &body[0] { + // Identity is a bonus but not strict; verify content equality at minimum + if !bytes.Equal(out, body) { + t.Fatalf("output must byte-match input when no changes needed") + } + } +} + +func TestNormalizeDateline_MultipleSystemReminderBlocksInSameText(t *testing.T) { + body := []byte(`{"messages":[{"role":"user","content":"\nToday’s date is 2026/07/01.\n\nsome prose\n\nAlso Todayʼs date is 2026/07/01.\n"}]}`) + out, hits, changed := NormalizeDateline(body) + if !changed { + t.Fatalf("expected changed=true") + } + if len(hits) != 2 { + t.Fatalf("expected 2 hits, got %d", len(hits)) + } + if bytes.Contains(out, []byte("2026/07/01")) { + t.Fatalf("slash separator must be gone: %s", string(out)) + } +} diff --git a/backend/internal/pkg/claude/constants.go b/backend/internal/pkg/claude/constants.go index 06335aeda9..b159f8f5a5 100644 --- a/backend/internal/pkg/claude/constants.go +++ b/backend/internal/pkg/claude/constants.go @@ -146,6 +146,12 @@ var DefaultModels = []Model{ DisplayName: "Claude Opus 4.8", CreatedAt: "2026-05-29T00:00:00Z", }, + { + ID: "claude-sonnet-5", + Type: "model", + DisplayName: "Claude Sonnet 5", + CreatedAt: "2026-07-01T00:00:00Z", + }, { ID: "claude-sonnet-4-6", Type: "model", diff --git a/backend/internal/pkg/usagestats/usage_log_types.go b/backend/internal/pkg/usagestats/usage_log_types.go index 7e0d12fdd2..a96ea670c4 100644 --- a/backend/internal/pkg/usagestats/usage_log_types.go +++ b/backend/internal/pkg/usagestats/usage_log_types.go @@ -261,17 +261,19 @@ type PlatformDashboardStats struct { // UsageLogFilters represents filters for usage log queries type UsageLogFilters struct { - UserID int64 - APIKeyID int64 - AccountID int64 - GroupID int64 - Model string - RequestType *int16 - Stream *bool - BillingType *int8 - BillingMode string - StartTime *time.Time - EndTime *time.Time + UserID int64 + APIKeyID int64 + AccountID int64 + GroupID int64 + Model string + // ModelFilterSource controls how Model is matched. Empty preserves raw usage_logs.model semantics. + ModelFilterSource string + RequestType *int16 + Stream *bool + BillingType *int8 + BillingMode string + StartTime *time.Time + EndTime *time.Time // ExactTotal requests exact COUNT(*) for pagination. Default false for fast large-table paging. ExactTotal bool } diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index 3fb9e3f737..a2a1e49770 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -143,23 +143,53 @@ func appendRawUsageLogModelWhereCondition(conditions []string, args []any, model } func appendUsageLogBillingModeWhereCondition(conditions []string, args []any, billingMode string) ([]string, []any) { + return appendUsageLogBillingModeWhereConditionWithAlias(conditions, args, billingMode, "") +} + +func appendUsageLogBillingModeWhereConditionWithAlias(conditions []string, args []any, billingMode string, alias string) ([]string, []any) { mode := strings.TrimSpace(billingMode) if mode == "" { return conditions, args } + column := func(name string) string { + if alias == "" { + return name + } + return alias + "." + name + } placeholder := fmt.Sprintf("$%d", len(args)+1) switch service.BillingMode(mode) { case service.BillingModeImage: - conditions = append(conditions, fmt.Sprintf("(billing_mode = %s OR COALESCE(image_count, 0) > 0)", placeholder)) + conditions = append(conditions, fmt.Sprintf("(%s = %s OR COALESCE(%s, 0) > 0)", column("billing_mode"), placeholder, column("image_count"))) case service.BillingModeToken: - conditions = append(conditions, fmt.Sprintf("(billing_mode = %s OR ((billing_mode IS NULL OR billing_mode = '') AND COALESCE(image_count, 0) <= 0))", placeholder)) + conditions = append(conditions, fmt.Sprintf("(%s = %s OR ((%s IS NULL OR %s = '') AND COALESCE(%s, 0) <= 0))", column("billing_mode"), placeholder, column("billing_mode"), column("billing_mode"), column("image_count"))) default: - conditions = append(conditions, fmt.Sprintf("billing_mode = %s", placeholder)) + conditions = append(conditions, fmt.Sprintf("%s = %s", column("billing_mode"), placeholder)) } args = append(args, mode) return conditions, args } +func appendUsageLogBillingModeQueryFilter(query string, args []any, billingMode string, alias string) (string, []any) { + conditions, args := appendUsageLogBillingModeWhereConditionWithAlias(nil, args, billingMode, alias) + if len(conditions) == 0 { + return query, args + } + return query + " AND " + conditions[0], args +} + +func appendUsageLogModelWhereCondition(conditions []string, args []any, model string, source string) ([]string, []any) { + if strings.TrimSpace(source) == "" { + return appendRawUsageLogModelWhereCondition(conditions, args, model) + } + if strings.TrimSpace(model) == "" { + return conditions, args + } + conditions = append(conditions, fmt.Sprintf("%s = $%d", resolveModelDimensionExpression(source), len(args)+1)) + args = append(args, model) + return conditions, args +} + // appendRawUsageLogModelQueryFilter keeps direct model filters on the raw model column for backward // compatibility with historical rows. Requested/upstream analytics must use // resolveModelDimensionExpression instead. @@ -172,6 +202,18 @@ func appendRawUsageLogModelQueryFilter(query string, args []any, model string) ( return query, args } +func appendUsageLogModelQueryFilter(query string, args []any, model string, source string) (string, []any) { + if strings.TrimSpace(source) == "" { + return appendRawUsageLogModelQueryFilter(query, args, model) + } + if strings.TrimSpace(model) == "" { + return query, args + } + query += fmt.Sprintf(" AND %s = $%d", resolveModelDimensionExpression(source), len(args)+1) + args = append(args, model) + return query, args +} + type usageLogRepository struct { client *dbent.Client sql sqlExecutor @@ -2806,7 +2848,7 @@ func (r *usageLogRepository) ListWithFilters(ctx context.Context, params paginat conditions = append(conditions, fmt.Sprintf("group_id = $%d", len(args)+1)) args = append(args, filters.GroupID) } - conditions, args = appendRawUsageLogModelWhereCondition(conditions, args, filters.Model) + conditions, args = appendUsageLogModelWhereCondition(conditions, args, filters.Model, filters.ModelFilterSource) conditions, args = appendRequestTypeOrStreamWhereCondition(conditions, args, filters.RequestType, filters.Stream) if filters.BillingType != nil { conditions = append(conditions, fmt.Sprintf("billing_type = $%d", len(args)+1)) @@ -3018,7 +3060,15 @@ func (r *usageLogRepository) GetBatchAPIKeyUsageStats(ctx context.Context, apiKe // GetUsageTrendWithFilters returns usage trend data with optional filters func (r *usageLogRepository) GetUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) (results []TrendDataPoint, err error) { - if shouldUsePreaggregatedTrend(granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType) { + return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "") +} + +func (r *usageLogRepository) GetUsageTrendWithUsageFilters(ctx context.Context, startTime, endTime time.Time, granularity string, filters UsageLogFilters) (results []TrendDataPoint, err error) { + return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) +} + +func (r *usageLogRepository) getUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []TrendDataPoint, err error) { + if shouldUsePreaggregatedTrend(granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType, billingMode) { aggregated, aggregatedErr := r.getUsageTrendFromAggregates(ctx, startTime, endTime, granularity) if aggregatedErr == nil && len(aggregated) > 0 { return aggregated, nil @@ -3059,12 +3109,13 @@ func (r *usageLogRepository) GetUsageTrendWithFilters(ctx context.Context, start query += fmt.Sprintf(" AND group_id = $%d", len(args)+1) args = append(args, groupID) } - query, args = appendRawUsageLogModelQueryFilter(query, args, model) + query, args = appendUsageLogModelQueryFilter(query, args, model, modelSource) query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream) if billingType != nil { query += fmt.Sprintf(" AND billing_type = $%d", len(args)+1) args = append(args, int16(*billingType)) } + query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") query += " GROUP BY date ORDER BY date ASC" rows, err := r.sql.QueryContext(ctx, query, args...) @@ -3087,7 +3138,7 @@ func (r *usageLogRepository) GetUsageTrendWithFilters(ctx context.Context, start return results, nil } -func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) bool { +func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string) bool { if granularity != "day" && granularity != "hour" { return false } @@ -3098,7 +3149,8 @@ func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID model == "" && requestType == nil && stream == nil && - billingType == nil + billingType == nil && + billingMode == "" } func (r *usageLogRepository) getUsageTrendFromAggregates(ctx context.Context, startTime, endTime time.Time, granularity string) (results []TrendDataPoint, err error) { @@ -3163,16 +3215,20 @@ func (r *usageLogRepository) getUsageTrendFromAggregates(ctx context.Context, st // GetModelStatsWithFilters returns model statistics with optional filters func (r *usageLogRepository) GetModelStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType, usagestats.ModelSourceRequested) + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, usagestats.ModelSourceRequested, "") } // GetModelStatsWithFiltersBySource returns model statistics with optional filters and model source dimension. // source: requested | upstream | mapping. func (r *usageLogRepository) GetModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8, source string) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType, source) + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, source, "") } -func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8, source string) (results []ModelStat, err error) { +func (r *usageLogRepository) GetModelStatsWithUsageFiltersBySource(ctx context.Context, startTime, endTime time.Time, filters UsageLogFilters, source string) (results []ModelStat, err error) { + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, source, filters.BillingMode) +} + +func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, source string, billingMode string) (results []ModelStat, err error) { actualCostExpr := "COALESCE(SUM(actual_cost), 0) as actual_cost" // 当仅按 account_id 聚合时,实际费用使用账号倍率(total_cost * account_rate_multiplier)。 if accountID > 0 && userID == 0 && apiKeyID == 0 { @@ -3214,11 +3270,16 @@ func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Contex query += fmt.Sprintf(" AND group_id = $%d", len(args)+1) args = append(args, groupID) } + if strings.TrimSpace(model) != "" { + query += fmt.Sprintf(" AND %s = $%d", modelExpr, len(args)+1) + args = append(args, model) + } query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream) if billingType != nil { query += fmt.Sprintf(" AND billing_type = $%d", len(args)+1) args = append(args, int16(*billingType)) } + query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") query += fmt.Sprintf(" GROUP BY %s ORDER BY total_tokens DESC", modelExpr) rows, err := r.sql.QueryContext(ctx, query, args...) @@ -3243,6 +3304,14 @@ func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Contex // GetGroupStatsWithFilters returns group usage statistics with optional filters func (r *usageLogRepository) GetGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) (results []usagestats.GroupStat, err error) { + return r.getGroupStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, "") +} + +func (r *usageLogRepository) GetGroupStatsWithUsageFilters(ctx context.Context, startTime, endTime time.Time, filters UsageLogFilters) (results []usagestats.GroupStat, err error) { + return r.getGroupStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) +} + +func (r *usageLogRepository) getGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []usagestats.GroupStat, err error) { query := ` SELECT COALESCE(ul.group_id, 0) as group_id, @@ -3274,11 +3343,17 @@ func (r *usageLogRepository) GetGroupStatsWithFilters(ctx context.Context, start query += fmt.Sprintf(" AND ul.group_id = $%d", len(args)+1) args = append(args, groupID) } + if strings.TrimSpace(model) != "" { + modelExpr := resolveModelDimensionExpressionWithAlias(usagestats.ModelSourceRequested, "ul") + query += fmt.Sprintf(" AND %s = $%d", modelExpr, len(args)+1) + args = append(args, model) + } query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream) if billingType != nil { query += fmt.Sprintf(" AND ul.billing_type = $%d", len(args)+1) args = append(args, int16(*billingType)) } + query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "ul") query += " GROUP BY ul.group_id, g.name ORDER BY total_tokens DESC" rows, err := r.sql.QueryContext(ctx, query, args...) @@ -3444,12 +3519,22 @@ func (r *usageLogRepository) GetAllGroupUsageSummary(ctx context.Context, todayS // resolveModelDimensionExpression maps model source type to a safe SQL expression. func resolveModelDimensionExpression(modelType string) string { - requestedExpr := "COALESCE(NULLIF(TRIM(requested_model), ''), model)" + return resolveModelDimensionExpressionWithAlias(modelType, "") +} + +func resolveModelDimensionExpressionWithAlias(modelType, alias string) string { + column := func(name string) string { + if alias == "" { + return name + } + return alias + "." + name + } + requestedExpr := fmt.Sprintf("COALESCE(NULLIF(TRIM(%s), ''), %s)", column("requested_model"), column("model")) switch usagestats.NormalizeModelSource(modelType) { case usagestats.ModelSourceUpstream: - return fmt.Sprintf("COALESCE(NULLIF(TRIM(upstream_model), ''), %s)", requestedExpr) + return fmt.Sprintf("COALESCE(NULLIF(TRIM(%s), ''), %s)", column("upstream_model"), requestedExpr) case usagestats.ModelSourceMapping: - return fmt.Sprintf("(%s || ' -> ' || COALESCE(NULLIF(TRIM(upstream_model), ''), %s))", requestedExpr, requestedExpr) + return fmt.Sprintf("(%s || ' -> ' || COALESCE(NULLIF(TRIM(%s), ''), %s))", requestedExpr, column("upstream_model"), requestedExpr) default: return requestedExpr } @@ -3523,7 +3608,7 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us conditions = append(conditions, fmt.Sprintf("group_id = $%d", len(args)+1)) args = append(args, filters.GroupID) } - conditions, args = appendRawUsageLogModelWhereCondition(conditions, args, filters.Model) + conditions, args = appendUsageLogModelWhereCondition(conditions, args, filters.Model, filters.ModelFilterSource) conditions, args = appendRequestTypeOrStreamWhereCondition(conditions, args, filters.RequestType, filters.Stream) if filters.BillingType != nil { conditions = append(conditions, fmt.Sprintf("billing_type = $%d", len(args)+1)) @@ -3587,7 +3672,7 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us } // endpoint 明细:best-effort(失败 log + 返空),不致命。 runEndpoints := func(c context.Context) { - res, err := r.GetEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType) + res, err := r.getEndpointStatsByColumnWithFilters(c, "inbound_endpoint", start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) if err != nil { if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err) @@ -3597,7 +3682,7 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us endpoints = res } runUpstream := func(c context.Context) { - res, err := r.GetUpstreamEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType) + res, err := r.getEndpointStatsByColumnWithFilters(c, "upstream_endpoint", start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) if err != nil { if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err) @@ -3607,7 +3692,7 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us upstreamEndpoints = res } runPaths := func(c context.Context) { - res, err := r.getEndpointPathStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType) + res, err := r.getEndpointPathStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) if err != nil { if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", err) @@ -3658,7 +3743,7 @@ type AccountUsageStatsResponse = usagestats.AccountUsageStatsResponse // EndpointStat represents endpoint usage statistics row. type EndpointStat = usagestats.EndpointStat -func (r *usageLogRepository) getEndpointStatsByColumnWithFilters(ctx context.Context, endpointColumn string, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) (results []EndpointStat, err error) { +func (r *usageLogRepository) getEndpointStatsByColumnWithFilters(ctx context.Context, endpointColumn string, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []EndpointStat, err error) { actualCostExpr := "COALESCE(SUM(actual_cost), 0) as actual_cost" if accountID > 0 && userID == 0 && apiKeyID == 0 { actualCostExpr = "COALESCE(SUM(COALESCE(account_stats_cost, total_cost) * COALESCE(account_rate_multiplier, 1)), 0) as actual_cost" @@ -3692,12 +3777,13 @@ func (r *usageLogRepository) getEndpointStatsByColumnWithFilters(ctx context.Con query += fmt.Sprintf(" AND group_id = $%d", len(args)+1) args = append(args, groupID) } - query, args = appendRawUsageLogModelQueryFilter(query, args, model) + query, args = appendUsageLogModelQueryFilter(query, args, model, modelSource) query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream) if billingType != nil { query += fmt.Sprintf(" AND billing_type = $%d", len(args)+1) args = append(args, int16(*billingType)) } + query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") query += " GROUP BY endpoint ORDER BY requests DESC" rows, err := r.sql.QueryContext(ctx, query, args...) @@ -3725,7 +3811,7 @@ func (r *usageLogRepository) getEndpointStatsByColumnWithFilters(ctx context.Con return results, nil } -func (r *usageLogRepository) getEndpointPathStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) (results []EndpointStat, err error) { +func (r *usageLogRepository) getEndpointPathStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []EndpointStat, err error) { actualCostExpr := "COALESCE(SUM(actual_cost), 0) as actual_cost" if accountID > 0 && userID == 0 && apiKeyID == 0 { actualCostExpr = "COALESCE(SUM(COALESCE(account_stats_cost, total_cost) * COALESCE(account_rate_multiplier, 1)), 0) as actual_cost" @@ -3763,12 +3849,13 @@ func (r *usageLogRepository) getEndpointPathStatsWithFilters(ctx context.Context query += fmt.Sprintf(" AND group_id = $%d", len(args)+1) args = append(args, groupID) } - query, args = appendRawUsageLogModelQueryFilter(query, args, model) + query, args = appendUsageLogModelQueryFilter(query, args, model, modelSource) query, args = appendRequestTypeOrStreamQueryFilter(query, args, requestType, stream) if billingType != nil { query += fmt.Sprintf(" AND billing_type = $%d", len(args)+1) args = append(args, int16(*billingType)) } + query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") query += " GROUP BY endpoint ORDER BY requests DESC" rows, err := r.sql.QueryContext(ctx, query, args...) @@ -3798,12 +3885,12 @@ func (r *usageLogRepository) getEndpointPathStatsWithFilters(ctx context.Context // GetEndpointStatsWithFilters returns inbound endpoint statistics with optional filters. func (r *usageLogRepository) GetEndpointStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) ([]EndpointStat, error) { - return r.getEndpointStatsByColumnWithFilters(ctx, "inbound_endpoint", startTime, endTime, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType) + return r.getEndpointStatsByColumnWithFilters(ctx, "inbound_endpoint", startTime, endTime, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "") } // GetUpstreamEndpointStatsWithFilters returns upstream endpoint statistics with optional filters. func (r *usageLogRepository) GetUpstreamEndpointStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) ([]EndpointStat, error) { - return r.getEndpointStatsByColumnWithFilters(ctx, "upstream_endpoint", startTime, endTime, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType) + return r.getEndpointStatsByColumnWithFilters(ctx, "upstream_endpoint", startTime, endTime, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "") } // GetAccountUsageStats returns comprehensive usage statistics for an account over a time range diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index 62ce892c5f..19710bca93 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -306,6 +306,20 @@ func TestAppendUsageLogBillingModeWhereCondition(t *testing.T) { } } +func TestAppendUsageLogBillingModeWhereConditionWithAlias(t *testing.T) { + conditions, args := appendUsageLogBillingModeWhereConditionWithAlias(nil, nil, string(service.BillingModeImage), "ul") + + require.Equal(t, []string{"(ul.billing_mode = $1 OR COALESCE(ul.image_count, 0) > 0)"}, conditions) + require.Equal(t, []any{string(service.BillingModeImage)}, args) +} + +func TestAppendUsageLogBillingModeQueryFilter(t *testing.T) { + query, args := appendUsageLogBillingModeQueryFilter("SELECT * FROM usage_logs WHERE user_id = $1", []any{int64(42)}, string(service.BillingModeToken), "") + + require.Equal(t, "SELECT * FROM usage_logs WHERE user_id = $1 AND (billing_mode = $2 OR ((billing_mode IS NULL OR billing_mode = '') AND COALESCE(image_count, 0) <= 0))", query) + require.Equal(t, []any{int64(42), string(service.BillingModeToken)}, args) +} + func anySliceToDriverValues(values []any) []driver.Value { out := make([]driver.Value, 0, len(values)) for _, value := range values { @@ -341,6 +355,26 @@ func TestUsageLogRepositoryListWithFiltersRequestTypePriority(t *testing.T) { require.NoError(t, mock.ExpectationsWereMet()) } +func TestUsageLogRepositoryListWithFiltersRequestedModelSource(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + + filters := usagestats.UsageLogFilters{ + Model: "gpt-5", + ModelFilterSource: usagestats.ModelSourceRequested, + } + + mock.ExpectQuery("SELECT .* FROM usage_logs WHERE COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$1 ORDER BY id DESC LIMIT \\$2 OFFSET \\$3"). + WithArgs("gpt-5", 21, 0). + WillReturnRows(sqlmock.NewRows([]string{"id"})) + + logs, page, err := repo.ListWithFilters(context.Background(), pagination.PaginationParams{Page: 1, PageSize: 20}, filters) + require.NoError(t, err) + require.Empty(t, logs) + require.NotNil(t, page) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUsageLogRepositoryGetUsageTrendWithFiltersRequestTypePriority(t *testing.T) { db, mock := newSQLMock(t) repo := &usageLogRepository{sql: db} @@ -360,6 +394,27 @@ func TestUsageLogRepositoryGetUsageTrendWithFiltersRequestTypePriority(t *testin require.NoError(t, mock.ExpectationsWereMet()) } +func TestUsageLogRepositoryGetUsageTrendWithUsageFiltersRequestedModelSource(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + + start := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(24 * time.Hour) + filters := usagestats.UsageLogFilters{ + Model: "gpt-5", + ModelFilterSource: usagestats.ModelSourceRequested, + } + + mock.ExpectQuery("AND COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$3"). + WithArgs(start, end, "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{"date", "requests", "input_tokens", "output_tokens", "cache_creation_tokens", "cache_read_tokens", "total_tokens", "cost", "actual_cost"})) + + trend, err := repo.GetUsageTrendWithUsageFilters(context.Background(), start, end, "day", filters) + require.NoError(t, err) + require.Empty(t, trend) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUsageLogRepositoryGetModelStatsWithFiltersRequestTypePriority(t *testing.T) { db, mock := newSQLMock(t) repo := &usageLogRepository{sql: db} @@ -379,6 +434,45 @@ func TestUsageLogRepositoryGetModelStatsWithFiltersRequestTypePriority(t *testin require.NoError(t, mock.ExpectationsWereMet()) } +func TestUsageLogRepositoryGetStatsWithFiltersRequestedModelSource(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + + filters := usagestats.UsageLogFilters{ + Model: "gpt-5", + ModelFilterSource: usagestats.ModelSourceRequested, + } + + mock.ExpectQuery("FROM usage_logs\\s+WHERE COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$1"). + WithArgs("gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{ + "total_requests", + "total_input_tokens", + "total_output_tokens", + "total_cache_tokens", + "total_cache_creation_tokens", + "total_cache_read_tokens", + "total_cost", + "total_actual_cost", + "total_account_cost", + "avg_duration_ms", + }).AddRow(int64(1), int64(2), int64(3), int64(4), int64(1), int64(3), 1.2, 1.0, 1.2, 20.0)) + mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(inbound_endpoint\\), ''\\), 'unknown'\\) AS endpoint"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) + mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(upstream_endpoint\\), ''\\), 'unknown'\\) AS endpoint"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) + mock.ExpectQuery("SELECT CONCAT\\("). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) + + stats, err := repo.GetStatsWithFilters(context.Background(), filters) + require.NoError(t, err) + require.Equal(t, int64(1), stats.TotalRequests) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUsageLogRepositoryGetStatsWithFiltersRequestTypePriority(t *testing.T) { db, mock := newSQLMock(t) repo := &usageLogRepository{sql: db} @@ -452,6 +546,29 @@ func TestUsageLogRepositoryGetModelStatsAccountCostColumn(t *testing.T) { require.NoError(t, mock.ExpectationsWereMet()) } +func TestUsageLogRepositoryGetModelStatsWithUsageFiltersAppliesRequestedModelFilter(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + + start := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(24 * time.Hour) + filters := usagestats.UsageLogFilters{Model: "gpt-5"} + + mock.ExpectQuery("AND COALESCE\\(NULLIF\\(TRIM\\(requested_model\\), ''\\), model\\) = \\$3"). + WithArgs(start, end, "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{ + "model", "requests", "input_tokens", "output_tokens", + "cache_creation_tokens", "cache_read_tokens", "total_tokens", + "cost", "actual_cost", "account_cost", + }).AddRow("gpt-5", int64(1), int64(10), int64(20), int64(0), int64(0), int64(30), 0.1, 0.08, 0.07)) + + results, err := repo.GetModelStatsWithUsageFiltersBySource(context.Background(), start, end, filters, usagestats.ModelSourceRequested) + require.NoError(t, err) + require.Len(t, results, 1) + require.Equal(t, "gpt-5", results[0].Model) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUsageLogRepositoryGetGroupStatsAccountCostColumn(t *testing.T) { db, mock := newSQLMock(t) repo := &usageLogRepository{sql: db} @@ -481,6 +598,28 @@ func TestUsageLogRepositoryGetGroupStatsAccountCostColumn(t *testing.T) { require.NoError(t, mock.ExpectationsWereMet()) } +func TestUsageLogRepositoryGetGroupStatsWithUsageFiltersAppliesRequestedModelFilter(t *testing.T) { + db, mock := newSQLMock(t) + repo := &usageLogRepository{sql: db} + + start := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(24 * time.Hour) + filters := usagestats.UsageLogFilters{Model: "gpt-5"} + + mock.ExpectQuery("AND COALESCE\\(NULLIF\\(TRIM\\(ul.requested_model\\), ''\\), ul.model\\) = \\$3"). + WithArgs(start, end, "gpt-5"). + WillReturnRows(sqlmock.NewRows([]string{ + "group_id", "group_name", "requests", "total_tokens", + "cost", "actual_cost", "account_cost", + }).AddRow(int64(1), "default", int64(1), int64(30), 0.1, 0.08, 0.07)) + + results, err := repo.GetGroupStatsWithUsageFilters(context.Background(), start, end, filters) + require.NoError(t, err) + require.Len(t, results, 1) + require.Equal(t, int64(1), results[0].GroupID) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestUsageLogRepositoryGetStatsWithFiltersAlwaysReturnsAccountCost(t *testing.T) { db, mock := newSQLMock(t) repo := &usageLogRepository{sql: db} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index e925e54121..4b4e5e0b04 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -846,6 +846,7 @@ func TestAPIContracts(t *testing.T) { "claude_oauth_system_prompt_blocks": "", "enable_anthropic_cache_ttl_1h_injection": false, "rewrite_message_cache_control": false, + "enable_client_dateline_normalization": true, "antigravity_user_agent_version": "", "enable_fingerprint_unification": true, "enable_metadata_passthrough": false, @@ -1090,6 +1091,7 @@ func TestAPIContracts(t *testing.T) { "claude_oauth_system_prompt_blocks": "", "enable_anthropic_cache_ttl_1h_injection": false, "rewrite_message_cache_control": false, + "enable_client_dateline_normalization": true, "antigravity_user_agent_version": "", "min_codex_version": "", "max_codex_version": "", @@ -2506,7 +2508,7 @@ func (r *stubUsageLogRepo) ListWithFilters(ctx context.Context, params paginatio continue } // Apply Model filter - if filters.Model != "" && log.Model != filters.Model { + if filters.Model != "" && stubUsageLogFilterModel(log, filters.ModelFilterSource) != filters.Model { continue } // Apply Stream filter @@ -2532,6 +2534,13 @@ func (r *stubUsageLogRepo) ListWithFilters(ctx context.Context, params paginatio return out, paginationResult(total, params), nil } +func stubUsageLogFilterModel(log service.UsageLog, source string) string { + if source == usagestats.ModelSourceRequested && log.RequestedModel != "" { + return log.RequestedModel + } + return log.Model +} + func (r *stubUsageLogRepo) GetGlobalStats(ctx context.Context, startTime, endTime time.Time) (*usagestats.UsageStats, error) { return nil, errors.New("not implemented") } @@ -2541,7 +2550,55 @@ func (r *stubUsageLogRepo) GetAccountUsageStats(ctx context.Context, accountID i } func (r *stubUsageLogRepo) GetStatsWithFilters(ctx context.Context, filters usagestats.UsageLogFilters) (*usagestats.UsageStats, error) { - return nil, errors.New("not implemented") + logs, _, err := r.ListWithFilters(ctx, pagination.PaginationParams{Page: 1, PageSize: 100000}, filters) + if err != nil { + return nil, err + } + + var totalRequests int64 + var totalInputTokens int64 + var totalOutputTokens int64 + var totalCacheTokens int64 + var totalCacheCreationTokens int64 + var totalCacheReadTokens int64 + var totalCost float64 + var totalActualCost float64 + var totalDuration int64 + var durationCount int64 + + for _, log := range logs { + totalRequests++ + totalInputTokens += int64(log.InputTokens) + totalOutputTokens += int64(log.OutputTokens) + totalCacheTokens += int64(log.CacheCreationTokens + log.CacheReadTokens) + totalCacheCreationTokens += int64(log.CacheCreationTokens) + totalCacheReadTokens += int64(log.CacheReadTokens) + totalCost += log.TotalCost + totalActualCost += log.ActualCost + if log.DurationMs != nil { + totalDuration += int64(*log.DurationMs) + durationCount++ + } + } + + var avgDuration float64 + if durationCount > 0 { + avgDuration = float64(totalDuration) / float64(durationCount) + } + + return &usagestats.UsageStats{ + TotalRequests: totalRequests, + TotalInputTokens: totalInputTokens, + TotalOutputTokens: totalOutputTokens, + TotalCacheTokens: totalCacheTokens, + TotalCacheCreationTokens: totalCacheCreationTokens, + TotalCacheReadTokens: totalCacheReadTokens, + TotalTokens: totalInputTokens + totalOutputTokens + totalCacheTokens, + TotalCost: totalCost, + TotalActualCost: totalActualCost, + AverageDurationMs: avgDuration, + Endpoints: []usagestats.EndpointStat{}, + }, nil } func (r *stubUsageLogRepo) GetAllGroupUsageSummary(ctx context.Context, todayStart time.Time) ([]usagestats.GroupUsageSummary, error) { return nil, errors.New("not implemented") diff --git a/backend/internal/server/routes/user.go b/backend/internal/server/routes/user.go index 0f3758f764..9f89687ad1 100644 --- a/backend/internal/server/routes/user.go +++ b/backend/internal/server/routes/user.go @@ -90,6 +90,7 @@ func RegisterUserRoutes( usage.GET("/dashboard/stats", h.Usage.DashboardStats) usage.GET("/dashboard/trend", h.Usage.DashboardTrend) usage.GET("/dashboard/models", h.Usage.DashboardModels) + usage.GET("/dashboard/snapshot-v2", h.Usage.DashboardSnapshotV2) usage.POST("/dashboard/api-keys-usage", h.Usage.DashboardAPIKeysUsage) } diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index e2046dd979..a781936598 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -276,8 +276,9 @@ func (s *BillingService) initFallbackPricing() { LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, } - // GPT-5.5 暂无独立定价,回退到 GPT-5.4 + // GPT-5.5 / GPT-5.5 Pro 暂无独立定价,回退到 GPT-5.4。 s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"] + s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"] s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{ InputPricePerToken: 7.5e-7, @@ -666,6 +667,8 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { // OpenAI(GPT-5 / Codex 族):仅匹配已知型号,避免未知 OpenAI 型号误计价。 if normalized := normalizeKnownOpenAICodexModel(modelLower); normalized != "" { switch normalized { + case "gpt-5.5-pro": + return s.fallbackPrices["gpt-5.5-pro"] case "gpt-5.5": return s.fallbackPrices["gpt-5.5"] case "gpt-5.4-mini": @@ -1057,7 +1060,7 @@ func isOpenAIGPT54Model(model string) bool { // normalizeCodexModel 的默认兜底把非 OpenAI 模型(claude-*、gemini-*、gpt-4o) // 误识别为 gpt-5.4。 normalized := normalizeKnownOpenAICodexModel(model) - return normalized == "gpt-5.4" || normalized == "gpt-5.5" + return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" } // CalculateCostWithConfig 使用配置中的默认倍率计算费用 diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 62792dc5df..92c143c6ff 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -263,6 +263,25 @@ func TestCalculateCost_OpenAIGPT54LongContextAppliesWholeSessionMultipliers(t *t require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10) } +func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) { + svc := newTestBillingService() + + tokens := UsageTokens{ + InputTokens: 300000, + OutputTokens: 4000, + } + + cost, err := svc.CalculateCost("gpt-5.5-pro", tokens, 1.0) + require.NoError(t, err) + + expectedInput := float64(tokens.InputTokens) * 2.5e-6 * 2.0 + expectedOutput := float64(tokens.OutputTokens) * 15e-6 * 1.5 + require.InDelta(t, expectedInput, cost.InputCost, 1e-10) + require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10) + require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10) + require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10) +} + // 回归测试 #2293:长上下文计费触发时,cache_read_tokens 也应应用 LongContextInputMultiplier。 // 修复前:CacheReadCost = tokens * 0.25e-6 (漏乘倍率,少计费用)。 // 修复后:CacheReadCost = tokens * 0.25e-6 * LongContextInputMultiplier(=2.0)。 diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index c83b683670..15e9ec73bf 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -450,6 +450,13 @@ const ( SettingKeyClaudeOAuthSystemPromptBlocks = "claude_oauth_system_prompt_blocks" // SettingKeyEnableAnthropicCacheTTL1hInjection 是否对 Anthropic OAuth/SetupToken 请求体注入 1h cache_control ttl(默认 false) SettingKeyEnableAnthropicCacheTTL1hInjection = "enable_anthropic_cache_ttl_1h_injection" + // SettingKeyEnableClientDatelineNormalization 是否对 Anthropic OAuth/SetupToken 账号 + // 的 /v1/messages 请求体做客户端 dateline 归一化(默认 true)。 + // 归一化把 system prompt / 块中 "Today's date is …" 语句里的 + // 非 ASCII 撇号与 "/" 日期分隔符还原为 ASCII 撇号 + "-" 分隔符,抹除某些客户端 + // 在检测到非官方 base URL 时注入的 3 bit 隐写指纹。仅适用于 Anthropic OAuth/SetupToken + // 账号;API Key 账号不受影响。 + SettingKeyEnableClientDatelineNormalization = "enable_client_dateline_normalization" // SettingKeyRewriteMessageCacheControl 是否改写 messages[*].content[*].cache_control(默认 false) SettingKeyRewriteMessageCacheControl = "rewrite_message_cache_control" // SettingKeyAntigravityUserAgentVersion Antigravity 上游 User-Agent 版本号(空值使用环境变量/默认值) diff --git a/backend/internal/service/gateway_beta_test.go b/backend/internal/service/gateway_beta_test.go index 6919c148f3..294b634eb7 100644 --- a/backend/internal/service/gateway_beta_test.go +++ b/backend/internal/service/gateway_beta_test.go @@ -242,3 +242,72 @@ func TestIsCountTokensUnsupported404(t *testing.T) { }) } } + +// TestDefaultBetaPolicy_Context1M_Sonnet5Whitelist 验证默认策略下 context-1m-2025-08-07 的分模型行为: +// - claude-sonnet-5 及后续版本:pass(放行),保留 1M 上下文能力 +// - 其他 sonnet 版本(4.x 及以下)、opus、haiku:filter(过滤),因为上游不支持 +func TestDefaultBetaPolicy_Context1M_Sonnet5Whitelist(t *testing.T) { + settings := DefaultBetaPolicySettings() + + // 找到 context-1m-2025-08-07 规则 + var rule *BetaPolicyRule + for i := range settings.Rules { + if settings.Rules[i].BetaToken == "context-1m-2025-08-07" { + rule = &settings.Rules[i] + break + } + } + require.NotNil(t, rule, "default policy must include context-1m-2025-08-07 rule") + require.Equal(t, BetaPolicyActionPass, rule.Action, "primary action for whitelisted models is pass") + require.Equal(t, BetaPolicyActionFilter, rule.FallbackAction, "non-whitelisted models must be filtered") + require.NotEmpty(t, rule.ModelWhitelist, "context-1m must be scoped to sonnet-5+ via whitelist") + + // 表驱动:模型 → 期望 action + // 覆盖每种上游路径下的模型 ID 变形:直连 Anthropic API、Vertex AI("@YYYYMMDD" 后缀)、 + // AWS Bedrock 跨区域推理(us./eu./apac./jp./au./us-gov./global./anthropic. 前缀)。 + cases := []struct { + model string + wantAction string + desc string + }{ + // —— 直连 Anthropic API —— sonnet-5 系列应放行 + {"claude-sonnet-5", BetaPolicyActionPass, "sonnet-5 canonical"}, + {"claude-sonnet-5-20260701", BetaPolicyActionPass, "sonnet-5 dated variant matches wildcard"}, + {"claude-sonnet-5-thinking", BetaPolicyActionPass, "sonnet-5 thinking variant matches wildcard"}, + // —— Vertex AI 归一化后的 sonnet-5 —— 也应放行 + {"claude-sonnet-5@20260701", BetaPolicyActionPass, "sonnet-5 Vertex-normalized dated form"}, + // —— AWS Bedrock 各跨区域前缀 sonnet-5 —— 也应放行 + {"us.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock us. sonnet-5"}, + {"eu.anthropic.claude-sonnet-5-20260701-v1:0", BetaPolicyActionPass, "bedrock eu. sonnet-5 dated"}, + {"apac.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock apac. sonnet-5"}, + {"jp.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock jp. sonnet-5"}, + {"au.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock au. sonnet-5"}, + {"us-gov.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock us-gov. sonnet-5"}, + {"global.anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock global. sonnet-5"}, + {"anthropic.claude-sonnet-5-v1", BetaPolicyActionPass, "bedrock no-region sonnet-5"}, + + // —— sonnet-4.x 及以下必须过滤 —— + {"claude-sonnet-4-6", BetaPolicyActionFilter, "sonnet-4.6 must be filtered"}, + {"claude-sonnet-4-5-20250929", BetaPolicyActionFilter, "sonnet-4.5 dated must be filtered"}, + {"claude-sonnet-4", BetaPolicyActionFilter, "sonnet-4 must be filtered"}, + {"claude-sonnet-4-5@20250929", BetaPolicyActionFilter, "sonnet-4.5 Vertex format must be filtered"}, + {"us.anthropic.claude-sonnet-4-6", BetaPolicyActionFilter, "bedrock us. sonnet-4.6 must be filtered"}, + {"us.anthropic.claude-sonnet-4-5-20250929-v1:0", BetaPolicyActionFilter, "bedrock us. sonnet-4.5 must be filtered"}, + // —— Opus / Haiku 必须过滤(无 1M) —— + {"claude-opus-4-8", BetaPolicyActionFilter, "opus must be filtered"}, + {"claude-opus-4-7", BetaPolicyActionFilter, "opus 4.7 must be filtered"}, + {"us.anthropic.claude-opus-4-8-v1", BetaPolicyActionFilter, "bedrock opus 4.8 must be filtered"}, + {"claude-haiku-4-5", BetaPolicyActionFilter, "haiku must be filtered"}, + {"us.anthropic.claude-haiku-4-5-20251001-v1:0", BetaPolicyActionFilter, "bedrock haiku must be filtered"}, + {"claude-3-5-sonnet-20241022", BetaPolicyActionFilter, "legacy sonnet 3.5 must be filtered"}, + // —— 特殊边界:不应把 "claude-sonnet-50" / "claude-sonnet-5.1" 之类意外命名误放行 —— + {"claude-sonnet-50", BetaPolicyActionFilter, "must not over-match a hypothetical sonnet-50"}, + } + + for _, tc := range cases { + t.Run(tc.model, func(t *testing.T) { + action, _ := resolveRuleAction(*rule, tc.model) + require.Equal(t, tc.wantAction, action, tc.desc) + }) + } +} diff --git a/backend/internal/service/gateway_dateline_normalization_test.go b/backend/internal/service/gateway_dateline_normalization_test.go new file mode 100644 index 0000000000..125e03f733 --- /dev/null +++ b/backend/internal/service/gateway_dateline_normalization_test.go @@ -0,0 +1,123 @@ +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/anthropicfp" + "github.com/stretchr/testify/require" +) + +// TestGatewayClientDatelineNormalization_Scope covers the account/switch matrix +// for the shouldNormalizeClientDateline gate: Anthropic OAuth/SetupToken pass +// only when the switch is on; API-Key and non-Anthropic platforms are excluded +// unconditionally. +func TestGatewayClientDatelineNormalization_Scope(t *testing.T) { + repo := &gatewayTTLSettingRepo{data: map[string]string{}} + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + svc := &GatewayService{ + settingService: NewSettingService(repo, &config.Config{}), + } + ctx := context.Background() + + // Default (missing key): fallback in parseSettings/cache loader is true. + require.True(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth})) + require.True(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeSetupToken})) + require.False(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey})) + require.False(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth})) + + // Switch off: no account qualifies. + repo.data[SettingKeyEnableClientDatelineNormalization] = "false" + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + require.False(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth})) + require.False(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeSetupToken})) + + // Switch back on: OAuth qualifies again. + repo.data[SettingKeyEnableClientDatelineNormalization] = "true" + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + require.True(t, svc.shouldNormalizeClientDateline(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth})) +} + +// TestGatewayClientDatelineNormalization_HelperNoRewrite exercises the code +// path used by Forward: the helper must return ok=false when the switch is +// off, when the account is API-Key, when the account is nil, and when the +// body carries no fingerprinted dateline. It must return ok=true and a +// rewritten body when both the switch is on and the account is Anthropic +// OAuth/SetupToken and a rewrite actually happened. +func TestGatewayClientDatelineNormalization_HelperNoRewrite(t *testing.T) { + repo := &gatewayTTLSettingRepo{data: map[string]string{ + SettingKeyEnableClientDatelineNormalization: "true", + }} + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + svc := &GatewayService{ + settingService: NewSettingService(repo, &config.Config{}), + } + ctx := context.Background() + + dirty := []byte(`{"messages":[{"role":"user","content":"\nToday’s date is 2026/07/01.\n"}]}`) + clean := []byte(`{"messages":[{"role":"user","content":"just hello"}]}`) + + // API-Key account: never rewrites, even with dirty payload. + next, ok := svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey}, dirty) + require.False(t, ok) + require.Nil(t, next) + + // Nil account: safe no-op. + next, ok = svc.normalizeClientDatelineIfEnabled(ctx, nil, dirty) + require.False(t, ok) + require.Nil(t, next) + + // OAuth account + clean body: no changes, ok=false. + next, ok = svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth}, clean) + require.False(t, ok) + require.Nil(t, next) + + // OAuth account + dirty body: rewritten, ok=true. + next, ok = svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth}, dirty) + require.True(t, ok) + require.NotNil(t, next) + require.Contains(t, string(next), "Today's date is 2026-07-01.") + require.NotContains(t, string(next), "2026/07/01") + require.NotContains(t, string(next), "Today’s date is") + + // SetupToken account + dirty body: rewritten, ok=true. + next, ok = svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeSetupToken}, dirty) + require.True(t, ok) + require.Contains(t, string(next), "Today's date is 2026-07-01.") + + // Switch off: even OAuth account is not rewritten. + repo.data[SettingKeyEnableClientDatelineNormalization] = "false" + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + next, ok = svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth}, dirty) + require.False(t, ok) + require.Nil(t, next) +} + +// TestGatewayClientDatelineNormalization_LeavesUserProseUntouched double-checks +// that the pure normalizer never touches content outside +// blocks. This is an integration guard between the switch-gated helper and +// the pkg/anthropicfp scope contract, tripped by anyone who broadens scope. +func TestGatewayClientDatelineNormalization_LeavesUserProseUntouched(t *testing.T) { + repo := &gatewayTTLSettingRepo{data: map[string]string{ + SettingKeyEnableClientDatelineNormalization: "true", + }} + gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{}) + svc := &GatewayService{ + settingService: NewSettingService(repo, &config.Config{}), + } + ctx := context.Background() + + // User prose that happens to include a fingerprint-looking sentence + // (outside ) must be preserved byte-for-byte. + body := []byte(`{"messages":[{"role":"user","content":"I wrote: Today’s date is 2026/07/01. What do you think?"}]}`) + next, ok := svc.normalizeClientDatelineIfEnabled(ctx, &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth}, body) + require.False(t, ok, "must not rewrite user prose outside ") + require.Nil(t, next) + + // Direct pure-fn check for redundancy. + out, hits, changed := anthropicfp.NormalizeDateline(body) + require.False(t, changed) + require.Empty(t, hits) + require.Equal(t, body, out) +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 4781da2435..9a6b06a28e 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -26,6 +26,7 @@ import ( "unsafe" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/anthropicfp" "github.com/Wei-Shaw/sub2api/internal/pkg/claude" "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" @@ -4782,6 +4783,32 @@ func (s *GatewayService) shouldInjectAnthropicCacheTTL1h(ctx context.Context, ac return s.settingService.IsAnthropicCacheTTL1hInjectionEnabled(ctx) } +// shouldNormalizeClientDateline reports whether the request body's client +// dateline should be normalized before forwarding to Anthropic. The switch is +// scoped to Anthropic OAuth/SetupToken accounts only; API-Key accounts and +// non-Anthropic platforms bypass this step entirely. +func (s *GatewayService) shouldNormalizeClientDateline(ctx context.Context, account *Account) bool { + if account == nil || !account.IsAnthropicOAuthOrSetupToken() || s == nil || s.settingService == nil { + return false + } + return s.settingService.IsClientDatelineNormalizationEnabled(ctx) +} + +// normalizeClientDatelineIfEnabled applies dateline normalization to body when +// the switch is on and the account qualifies. Returns (nextBody, true) only +// when the body actually changed; otherwise returns (nil, false) so callers +// can skip the writeback. +func (s *GatewayService) normalizeClientDatelineIfEnabled(ctx context.Context, account *Account, body []byte) ([]byte, bool) { + if !s.shouldNormalizeClientDateline(ctx, account) { + return nil, false + } + next, _, changed := anthropicfp.NormalizeDateline(body) + if !changed { + return nil, false + } + return next, true +} + func (s *GatewayService) claudeOAuthSystemPromptInjectionSettings(ctx context.Context) (bool, string, string) { if s == nil || s.settingService == nil { return true, "", "" @@ -4931,6 +4958,16 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A } } + // 客户端 dateline 归一化:仅对 Anthropic OAuth/SetupToken 账号生效。 + // 抹除 "Today's date is …" 语句里可能被注入的隐写指纹(4 种撇号 × 2 种日期 + // 分隔符),还原为 ASCII 撇号 + "-" 分隔符。运行在 mimicry 分支之外, + // 保证真实 Claude Code 客户端注入的指纹同样被清洗。 + if next, ok := s.normalizeClientDatelineIfEnabled(ctx, account, body); ok { + if err := replaceBody(next); err != nil { + return nil, err + } + } + // 强制执行 cache_control 块数量限制(最多 4 个) if err := replaceBody(enforceCacheControlLimit(body)); err != nil { return nil, err diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 01b5d3e526..c551ba54f1 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -10,6 +10,7 @@ import ( var codexModelMap = map[string]string{ "gpt-5.5": "gpt-5.5", + "gpt-5.5-pro": "gpt-5.5-pro", "codex-auto-review": "codex-auto-review", "gpt-5.4": "gpt-5.4", "gpt-5.4-mini": "gpt-5.4-mini", @@ -61,6 +62,7 @@ var codexVersionModelPrefixes = []struct { {prefix: "gpt-5.3-codex", target: "gpt-5.3-codex"}, {prefix: "gpt-5.4-mini", target: "gpt-5.4-mini"}, {prefix: "gpt-5.4-nano", target: "gpt-5.4-nano"}, + {prefix: "gpt-5.5-pro", target: "gpt-5.5-pro"}, {prefix: "gpt-5.5", target: "gpt-5.5"}, {prefix: "gpt-5.4", target: "gpt-5.4"}, {prefix: "gpt-5.2", target: "gpt-5.2"}, @@ -1157,11 +1159,41 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an } typ, _ := m["type"].(string) - // chatgpt.com codex backend (OAuth path) does not persist reasoning - // items because applyCodexOAuthTransform forces store=false. Any rs_* - // reference replayed in input is guaranteed to 404 upstream - // ("Item with id 'rs_...' not found"). Drop reasoning items entirely. + // chatgpt.com codex (OAuth path) runs with store=false (forced by + // applyCodexOAuthTransform). Replaying a reasoning item with its rs_* + // id but no encrypted_content 404s upstream ("Item with id 'rs_...' + // not found") — the 404 is triggered by the id lookup, not by the + // reasoning item itself. So strip the id (always, independent of + // PreserveReferences) yet keep the item: under store=false + // encrypted_content is the official channel for carrying reasoning + // context across turns, and dropping the whole item silently degrades + // multi-turn agent reasoning. Preserve encrypted_content/content/ + // summary and every other field verbatim. Upstream additionally + // requires a summary field — a missing one is rejected with 400 + // "Missing required parameter 'input[N].summary'" — so backfill an + // empty array when it is absent. Contracts verified end-to-end against + // chatgpt.com codex (gpt-5.5); see issue #1957. + // compaction_summary items (cmp_*) are the other encrypted_content + // carrier. Verified against the live backend: they require + // encrypted_content (a missing one is rejected with 400), and with it + // present the cmp_* id does not 404 whether kept or stripped. Being + // neither reasoning nor tool calls, they flow through the generic path + // below (id stripped when !PreserveReferences, encrypted_content + // preserved either way), which is safe and needs no special-casing. if typ == "reasoning" { + newItem := make(map[string]any, len(m)) + for key, value := range m { + if key == "id" { + // rs_* id replayed under store=false 404s; strip it. + continue + } + newItem[key] = value + } + if summary, ok := newItem["summary"]; !ok || summary == nil { + // Upstream requires a summary field; an empty array satisfies it. + newItem["summary"] = []any{} + } + filtered = append(filtered, newItem) continue } diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index e9157377f6..bb27c81352 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -898,6 +898,10 @@ func TestNormalizeCodexModel_Gpt53(t *testing.T) { "gpt-5.4": "gpt-5.4", "gpt5.5": "gpt-5.5", "openai/gpt5.5": "gpt-5.5", + "gpt-5.5-pro": "gpt-5.5-pro", + "gpt5.5-pro": "gpt-5.5-pro", + "openai/gpt5.5-pro": "gpt-5.5-pro", + "gpt-5.5-pro-high": "gpt-5.5-pro", "codex-auto-review": "codex-auto-review", "gpt5.4": "gpt-5.4", "gpt-5.4-high": "gpt-5.4", @@ -1343,23 +1347,24 @@ func TestIsInstructionsEmpty(t *testing.T) { } } -func TestFilterCodexInput_DropsReasoningItemsRegardlessOfPreserveReferences(t *testing.T) { - // Reasoning items in input[] reference rs_* IDs that were emitted by - // chatgpt.com under store=false (forced by applyCodexOAuthTransform). - // They are never persisted upstream, so forwarding them produces a - // guaranteed 404 ("Item with id 'rs_...' not found"). Drop them - // regardless of preserveReferences. See: Wei-Shaw/sub2api issue #1957. - +// TestFilterCodexInput_PreservesReasoningStripsID covers the core OAuth-path +// reasoning contract (replaces the earlier "drops reasoning" test, whose +// premise was wrong). A reasoning item carrying encrypted_content is the +// official channel for replaying reasoning context across turns under +// store=false, so it must survive the filter with encrypted_content intact; +// only its rs_* id is stripped (always, independent of PreserveReferences) +// because a bare rs_* id replayed under store=false 404s upstream. Contracts +// 1/2/3, verified end-to-end against chatgpt.com codex (gpt-5.5). See issue +// #1957. +func TestFilterCodexInput_PreservesReasoningStripsID(t *testing.T) { build := func() []any { return []any{ - map[string]any{"type": "message", "id": "msg_0", "role": "user", "content": "hi"}, map[string]any{ - "type": "reasoning", - "id": "rs_0672f12450da0b9c0169f07220a6c08198b68c2455ced99344", - "summary": []any{}, + "type": "reasoning", + "id": "rs_0672f12450da0b9c0169f07220a6c08198b68c2455ced99344", + "encrypted_content": "gAAAAAB-enc-payload", + "summary": []any{}, }, - map[string]any{"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "tool"}, - map[string]any{"type": "function_call_output", "call_id": "call_1", "output": "{}"}, } } @@ -1367,31 +1372,175 @@ func TestFilterCodexInput_DropsReasoningItemsRegardlessOfPreserveReferences(t *t preserve := preserve t.Run(fmt.Sprintf("preserveReferences=%v", preserve), func(t *testing.T) { filtered := filterCodexInput(build(), preserve) + require.Len(t, filtered, 1) + item, ok := filtered[0].(map[string]any) + require.True(t, ok) + // Contract 2: the reasoning item survives the filter. + require.Equal(t, "reasoning", item["type"]) + // Contract 2: encrypted_content (cross-turn channel) preserved verbatim. + require.Equal(t, "gAAAAAB-enc-payload", item["encrypted_content"]) + // Contract 1/3: rs_* id stripped unconditionally, even when + // PreserveReferences=true (id lookup, not the item, triggers the 404). + _, hasID := item["id"] + require.False(t, hasID) + // summary passed through untouched. + summary, ok := item["summary"].([]any) + require.True(t, ok) + require.Len(t, summary, 0) + }) + } +} + +// TestFilterCodexInput_BareReasoningStripsIDBackfillsSummary covers contract 1 +// plus 5: a reasoning item carrying only an rs_* id (no encrypted_content) is +// kept as an empty shell with the id stripped, and a missing summary is +// backfilled to [] so upstream does not reject it with 400 "Missing required +// parameter 'input[N].summary'". Verified against chatgpt.com codex (gpt-5.5). +func TestFilterCodexInput_BareReasoningStripsIDBackfillsSummary(t *testing.T) { + input := []any{ + map[string]any{ + "type": "reasoning", + "id": "rs_0672f12450da0b9c0169f07220a6c08198b68c2455ced99344", + }, + } + + filtered := filterCodexInput(input, false) + require.Len(t, filtered, 1) + + item, ok := filtered[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "reasoning", item["type"]) + // Contract 1: id stripped. + _, hasID := item["id"] + require.False(t, hasID) + // Contract 5: summary backfilled to an empty array. + summary, ok := item["summary"].([]any) + require.True(t, ok) + require.Len(t, summary, 0) +} + +// TestFilterCodexInput_ReasoningBackfillsMissingSummary isolates contract 5: +// even when a reasoning item carries other content (here encrypted_content), +// a missing summary field is always added as [] before forwarding upstream. +func TestFilterCodexInput_ReasoningBackfillsMissingSummary(t *testing.T) { + input := []any{ + map[string]any{ + "type": "reasoning", + "id": "rs_abc", + "encrypted_content": "gAAAAAB-enc", + }, + } + + filtered := filterCodexInput(input, false) + require.Len(t, filtered, 1) + + item, ok := filtered[0].(map[string]any) + require.True(t, ok) + summary, ok := item["summary"].([]any) + require.True(t, ok) + require.Len(t, summary, 0) + // encrypted_content still preserved alongside the backfilled summary. + require.Equal(t, "gAAAAAB-enc", item["encrypted_content"]) +} + +// TestFilterCodexInput_PreservesReasoningSummaryAndContent verifies that a +// non-empty summary is not overwritten and that arbitrary reasoning fields +// (e.g. content) survive verbatim — only the id is removed. +func TestFilterCodexInput_PreservesReasoningSummaryAndContent(t *testing.T) { + summary := []any{ + map[string]any{"type": "summary_text", "text": "Considered the options."}, + } + content := []any{ + map[string]any{"type": "reasoning_text", "text": "internal chain"}, + } + input := []any{ + map[string]any{ + "type": "reasoning", + "id": "rs_abc", + "summary": summary, + "content": content, + "encrypted_content": "gAAAAAB-enc", + }, + } + + filtered := filterCodexInput(input, false) + require.Len(t, filtered, 1) + + item, ok := filtered[0].(map[string]any) + require.True(t, ok) + // Non-empty summary preserved verbatim (not replaced with []). + require.Equal(t, summary, item["summary"]) + // content preserved verbatim. + require.Equal(t, content, item["content"]) + require.Equal(t, "gAAAAAB-enc", item["encrypted_content"]) + _, hasID := item["id"] + require.False(t, hasID) +} + +// TestFilterCodexInput_PreservesReasoningInMixedInput exercises contract 7: +// reasoning items are stripped of their rs_* ids but kept (with +// encrypted_content) while message / function_call / function_call_output +// items flow through unchanged, with tool-call pairing (call_id) intact. +func TestFilterCodexInput_PreservesReasoningInMixedInput(t *testing.T) { + build := func() []any { + return []any{ + map[string]any{"type": "message", "id": "msg_0", "role": "user", "content": "hi"}, + map[string]any{ + "type": "reasoning", + "id": "rs_1", + "encrypted_content": "gAAAAAB-enc-1", + "summary": []any{}, + }, + map[string]any{ + "type": "reasoning", + "id": "rs_2", + "summary": []any{}, + }, + // call_id already in fc_ form so the unrelated call_->fc_ + // normalization does not obscure the pairing assertion. + map[string]any{"type": "function_call", "id": "fc_1", "call_id": "fc_1", "name": "tool", "arguments": "{}"}, + map[string]any{"type": "function_call_output", "call_id": "fc_1", "output": "{}"}, + } + } + + for _, preserve := range []bool{true, false} { + preserve := preserve + t.Run(fmt.Sprintf("preserveReferences=%v", preserve), func(t *testing.T) { + filtered := filterCodexInput(build(), preserve) + // Nothing is dropped: both reasoning items are now preserved. + require.Len(t, filtered, 5) + + byType := make(map[string][]map[string]any) for _, raw := range filtered { item, ok := raw.(map[string]any) require.True(t, ok) - require.NotEqual(t, "reasoning", item["type"], - "reasoning items must be dropped from input on the OAuth path") + typ, _ := item["type"].(string) + byType[typ] = append(byType[typ], item) + // No surviving item may carry an rs_* id. if id, ok := item["id"].(string); ok { require.False(t, strings.HasPrefix(id, "rs_"), "no item carrying an rs_* id should survive the filter") } } - // Sanity check: the non-reasoning items should still be present. - gotTypes := make(map[string]int) - for _, raw := range filtered { - item, ok := raw.(map[string]any) - require.True(t, ok) - typ, ok := item["type"].(string) - require.True(t, ok) - gotTypes[typ]++ + // Both reasoning items kept, ids stripped, summary present. + require.Len(t, byType["reasoning"], 2) + for _, r := range byType["reasoning"] { + _, hasID := r["id"] + require.False(t, hasID) + _, hasSummary := r["summary"] + require.True(t, hasSummary) } - require.Equal(t, 1, gotTypes["message"]) - require.Equal(t, 1, gotTypes["function_call"]) - require.Equal(t, 1, gotTypes["function_call_output"]) - require.Equal(t, 0, gotTypes["reasoning"]) + require.Equal(t, "gAAAAAB-enc-1", byType["reasoning"][0]["encrypted_content"]) + + // message / function_call(+output) untouched by reasoning handling. + require.Len(t, byType["message"], 1) + // Contract 7: tool-call pairing by call_id is unaffected. + require.Len(t, byType["function_call"], 1) + require.Equal(t, "fc_1", byType["function_call"][0]["call_id"]) + require.Len(t, byType["function_call_output"], 1) + require.Equal(t, "fc_1", byType["function_call_output"][0]["call_id"]) }) } } diff --git a/backend/internal/service/openai_compat_prompt_cache_key_test.go b/backend/internal/service/openai_compat_prompt_cache_key_test.go index 3fe7db6ef6..ce1c68a2b7 100644 --- a/backend/internal/service/openai_compat_prompt_cache_key_test.go +++ b/backend/internal/service/openai_compat_prompt_cache_key_test.go @@ -16,6 +16,7 @@ func mustRawJSON(t *testing.T, s string) json.RawMessage { func TestShouldAutoInjectPromptCacheKeyForCompat(t *testing.T) { require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.5")) + require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.5-pro")) require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.4")) require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.4-mini")) require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.2")) diff --git a/backend/internal/service/openai_model_alias.go b/backend/internal/service/openai_model_alias.go index 2fa2c90efd..ac2a8cf942 100644 --- a/backend/internal/service/openai_model_alias.go +++ b/backend/internal/service/openai_model_alias.go @@ -65,6 +65,8 @@ func normalizeKnownOpenAICodexModel(model string) string { } switch { + case strings.Contains(normalized, "gpt-5.5-pro"): + return "gpt-5.5-pro" case strings.Contains(normalized, "gpt-5.5"): return "gpt-5.5" case strings.Contains(normalized, "gpt-5.4-mini"): diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index 020e887528..0f65740be7 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -94,6 +94,15 @@ func TestResolveOpenAIForwardModel(t *testing.T) { defaultMappedModel: "gpt-5.4", expectedModel: "gpt-5.5", }, + { + name: "preserves gpt-5.5-pro instead of group default", + account: &Account{ + Credentials: map[string]any{}, + }, + requestedModel: "gpt-5.5-pro", + defaultMappedModel: "gpt-5.5", + expectedModel: "gpt-5.5-pro", + }, { name: "preserves compact-spelled gpt5.5 instead of group default", account: &Account{ @@ -261,6 +270,12 @@ func TestNormalizeOpenAIModelForUpstream(t *testing.T) { model: "gpt-5.4-high", want: "gpt-5.4", }, + { + name: "oauth preserves GPT-5.5 Pro model", + account: &Account{Type: AccountTypeOAuth}, + model: "openai/gpt-5.5-pro", + want: "gpt-5.5-pro", + }, { name: "oauth preserves codex auto review model", account: &Account{Type: AccountTypeOAuth}, @@ -303,3 +318,17 @@ func TestUsageBillingModelCandidatesPreserveCodexAutoReviewModel(t *testing.T) { } } } + +func TestUsageBillingModelCandidatesPreserveGPT55ProModel(t *testing.T) { + candidates := usageBillingModelCandidates("openai/gpt-5.5-pro") + + expected := []string{"openai/gpt-5.5-pro", "gpt-5.5-pro"} + if len(candidates) != len(expected) { + t.Fatalf("usageBillingModelCandidates(openai/gpt-5.5-pro) = %#v, want %#v", candidates, expected) + } + for i := range expected { + if candidates[i] != expected[i] { + t.Fatalf("usageBillingModelCandidates(openai/gpt-5.5-pro) = %#v, want %#v", candidates, expected) + } + } +} diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 81415d457f..00e9d9029b 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -112,6 +112,7 @@ type cachedGatewayForwardingSettings struct { claudeOAuthSystemPromptBlocks string anthropicCacheTTL1hInjection bool rewriteMessageCacheControl bool + clientDatelineNormalization bool expiresAt int64 // unix nano } @@ -2207,6 +2208,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting updates[SettingKeyClaudeOAuthSystemPromptBlocks] = settings.ClaudeOAuthSystemPromptBlocks updates[SettingKeyEnableAnthropicCacheTTL1hInjection] = strconv.FormatBool(settings.EnableAnthropicCacheTTL1hInjection) updates[SettingKeyRewriteMessageCacheControl] = strconv.FormatBool(settings.RewriteMessageCacheControl) + updates[SettingKeyEnableClientDatelineNormalization] = strconv.FormatBool(settings.EnableClientDatelineNormalization) updates[SettingKeyAntigravityUserAgentVersion] = antigravity.NormalizeUserAgentVersion(settings.AntigravityUserAgentVersion) updates[SettingKeyOpenAICodexUserAgent] = strings.TrimSpace(settings.OpenAICodexUserAgent) // codex_cli_only 加固 @@ -2345,6 +2347,7 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { claudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks, anthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection, rewriteMessageCacheControl: settings.RewriteMessageCacheControl, + clientDatelineNormalization: settings.EnableClientDatelineNormalization, expiresAt: time.Now().Add(gatewayForwardingCacheTTL).UnixNano(), }) s.antigravityUAVersionSF.Forget("antigravity_user_agent_version") @@ -2548,6 +2551,7 @@ func (s *SettingService) IsBackendModeEnabled(ctx context.Context) bool { type gatewayForwardingSettingsResult struct { fp, mp, cch, claudeOAuthSystemPromptInjection, cacheTTL1h, rewriteMessageCacheControl bool + clientDatelineNormalization bool claudeOAuthSystemPrompt, claudeOAuthSystemPromptBlocks string } @@ -2563,6 +2567,7 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) claudeOAuthSystemPromptBlocks: cached.claudeOAuthSystemPromptBlocks, cacheTTL1h: cached.anthropicCacheTTL1hInjection, rewriteMessageCacheControl: cached.rewriteMessageCacheControl, + clientDatelineNormalization: cached.clientDatelineNormalization, } } } @@ -2578,6 +2583,7 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) claudeOAuthSystemPromptBlocks: cached.claudeOAuthSystemPromptBlocks, cacheTTL1h: cached.anthropicCacheTTL1hInjection, rewriteMessageCacheControl: cached.rewriteMessageCacheControl, + clientDatelineNormalization: cached.clientDatelineNormalization, }, nil } } @@ -2592,6 +2598,7 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) SettingKeyClaudeOAuthSystemPromptBlocks, SettingKeyEnableAnthropicCacheTTL1hInjection, SettingKeyRewriteMessageCacheControl, + SettingKeyEnableClientDatelineNormalization, }) if err != nil { slog.Warn("failed to get gateway forwarding settings", "error", err) @@ -2602,9 +2609,10 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) claudeOAuthSystemPromptInjection: true, anthropicCacheTTL1hInjection: false, rewriteMessageCacheControl: s.defaultRewriteMessageCacheControl(), + clientDatelineNormalization: true, expiresAt: time.Now().Add(gatewayForwardingErrorTTL).UnixNano(), }) - return gatewayForwardingSettingsResult{fp: true, claudeOAuthSystemPromptInjection: true, rewriteMessageCacheControl: s.defaultRewriteMessageCacheControl()}, nil + return gatewayForwardingSettingsResult{fp: true, claudeOAuthSystemPromptInjection: true, rewriteMessageCacheControl: s.defaultRewriteMessageCacheControl(), clientDatelineNormalization: true}, nil } fp := true if v, ok := values[SettingKeyEnableFingerprintUnification]; ok && v != "" { @@ -2623,6 +2631,10 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) if v, ok := values[SettingKeyRewriteMessageCacheControl]; ok && v != "" { rewriteMessageCacheControl = v == "true" } + clientDatelineNormalization := true + if v, ok := values[SettingKeyEnableClientDatelineNormalization]; ok && v != "" { + clientDatelineNormalization = v == "true" + } gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{ fingerprintUnification: fp, metadataPassthrough: mp, @@ -2632,6 +2644,7 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) claudeOAuthSystemPromptBlocks: systemPromptBlocks, anthropicCacheTTL1hInjection: cacheTTL1h, rewriteMessageCacheControl: rewriteMessageCacheControl, + clientDatelineNormalization: clientDatelineNormalization, expiresAt: time.Now().Add(gatewayForwardingCacheTTL).UnixNano(), }) return gatewayForwardingSettingsResult{ @@ -2643,12 +2656,13 @@ func (s *SettingService) getGatewayForwardingSettingsCached(ctx context.Context) claudeOAuthSystemPromptBlocks: systemPromptBlocks, cacheTTL1h: cacheTTL1h, rewriteMessageCacheControl: rewriteMessageCacheControl, + clientDatelineNormalization: clientDatelineNormalization, }, nil }) if r, ok := val.(gatewayForwardingSettingsResult); ok { return r } - return gatewayForwardingSettingsResult{fp: true, claudeOAuthSystemPromptInjection: true} + return gatewayForwardingSettingsResult{fp: true, claudeOAuthSystemPromptInjection: true, clientDatelineNormalization: true} } // GetGatewayForwardingSettings returns cached gateway forwarding settings. @@ -2669,6 +2683,12 @@ func (s *SettingService) IsRewriteMessageCacheControlEnabled(ctx context.Context return s.getGatewayForwardingSettingsCached(ctx).rewriteMessageCacheControl } +// IsClientDatelineNormalizationEnabled 检查是否启用 Anthropic OAuth/SetupToken 请求体 +// 的客户端 dateline 归一化。默认开启。 +func (s *SettingService) IsClientDatelineNormalizationEnabled(ctx context.Context) bool { + return s.getGatewayForwardingSettingsCached(ctx).clientDatelineNormalization +} + // GetClaudeOAuthSystemPromptInjectionSettings returns the Claude OAuth mimic // system block switch, legacy custom expansion prompt, and configurable blocks JSON. // Empty values mean use the built-in Claude Code default blocks. @@ -3173,6 +3193,7 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { SettingKeyAllowUngroupedKeyScheduling: "false", SettingKeyEnableAnthropicCacheTTL1hInjection: "false", SettingKeyRewriteMessageCacheControl: strconv.FormatBool(s.defaultRewriteMessageCacheControl()), + SettingKeyEnableClientDatelineNormalization: "true", SettingKeyAntigravityUserAgentVersion: "", SettingKeyOpenAICodexUserAgent: "", SettingPaymentVisibleMethodAlipaySource: "", @@ -3711,6 +3732,11 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin } else { result.RewriteMessageCacheControl = s.defaultRewriteMessageCacheControl() } + if v, ok := settings[SettingKeyEnableClientDatelineNormalization]; ok && v != "" { + result.EnableClientDatelineNormalization = v == "true" + } else { + result.EnableClientDatelineNormalization = true + } result.AntigravityUserAgentVersion = antigravity.NormalizeUserAgentVersion(settings[SettingKeyAntigravityUserAgentVersion]) result.OpenAICodexUserAgent = strings.TrimSpace(settings[SettingKeyOpenAICodexUserAgent]) // codex_cli_only 加固 diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 17d2d741ca..ac225e1b14 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -197,6 +197,7 @@ type SystemSettings struct { ClaudeOAuthSystemPrompt string // Claude OAuth mimic 路径注入的通用扩展 system prompt;空值使用内置默认 ClaudeOAuthSystemPromptBlocks string // Claude OAuth mimic 路径注入的 system blocks JSON 配置;空值使用内置默认 EnableAnthropicCacheTTL1hInjection bool // 是否对 Anthropic OAuth/SetupToken 请求体注入 1h cache_control ttl(默认 false) + EnableClientDatelineNormalization bool // 是否对 Anthropic OAuth/SetupToken 请求体做客户端 dateline 归一化(默认 true) RewriteMessageCacheControl bool // 是否改写 messages[*].content[*].cache_control(默认 false) AntigravityUserAgentVersion string // Antigravity 上游 User-Agent 版本号;空值使用配置/默认值 OpenAICodexUserAgent string // OpenAI Codex 上游完整 User-Agent;空值使用内置默认 @@ -488,6 +489,23 @@ func DefaultRateLimit429CooldownSettings() *RateLimit429CooldownSettings { } // DefaultBetaPolicySettings 返回默认的 Beta 策略配置 +// +// context-1m-2025-08-07 的默认策略: +// - 仅 claude-sonnet-5 及后续版本(如 claude-sonnet-5-*)在上游默认支持 1M 上下文。 +// - Sonnet 4.x 及以下、Opus、Haiku 上游都不支持该 beta,透传上去会被上游 400 或降级。 +// - 因此默认对 sonnet-5* 放行、其余全部过滤,与上游能力保持一致。 +// +// 白名单需要覆盖每个上游路径的模型 ID 变形: +// - 直连 Anthropic API(OAuth mimic / API Key / SetupToken):模型保持客户端原样 +// (如 "claude-sonnet-5"、"claude-sonnet-5-YYYYMMDD"、"claude-sonnet-5-thinking")。 +// - Vertex AI:normalizeVertexAnthropicModelID 会把 "-YYYYMMDD" 后缀转成 "@YYYYMMDD" +// (如 "claude-sonnet-5@YYYYMMDD")。 +// - AWS Bedrock:ResolveBedrockModelID 会输出带跨区域前缀的模型 ID +// (us./eu./apac./jp./au./us-gov./global. 或无前缀的 "anthropic." 形式)。 +// +// 白名单只用后缀通配符(matchModelPattern 语义),因此每个路径都需要显式列出前缀。 +// 精确匹配 "claude-sonnet-5" + 后缀 "-*" 与 "@*",可覆盖直连/Vertex 场景,同时避免误伤 +// 未来可能出现的 "claude-sonnet-50" 或 "claude-sonnet-5.x" 之类的意外命名。 func DefaultBetaPolicySettings() *BetaPolicySettings { return &BetaPolicySettings{ Rules: []BetaPolicyRule{ @@ -498,8 +516,26 @@ func DefaultBetaPolicySettings() *BetaPolicySettings { }, { BetaToken: "context-1m-2025-08-07", - Action: BetaPolicyActionFilter, + Action: BetaPolicyActionPass, Scope: BetaPolicyScopeAll, + ModelWhitelist: []string{ + // 直连 Anthropic API(客户端请求 model 原样) + "claude-sonnet-5", + "claude-sonnet-5-*", + // Vertex AI 走 normalizeVertexAnthropicModelID 后 "@YYYYMMDD" 格式 + "claude-sonnet-5@*", + // AWS Bedrock cross-region inference profile + "us.anthropic.claude-sonnet-5*", + "eu.anthropic.claude-sonnet-5*", + "apac.anthropic.claude-sonnet-5*", + "jp.anthropic.claude-sonnet-5*", + "au.anthropic.claude-sonnet-5*", + "us-gov.anthropic.claude-sonnet-5*", + "global.anthropic.claude-sonnet-5*", + // AWS Bedrock 无 cross-region 前缀 + "anthropic.claude-sonnet-5*", + }, + FallbackAction: BetaPolicyActionFilter, }, }, } diff --git a/backend/internal/service/usage_service.go b/backend/internal/service/usage_service.go index b56f96edf7..a085a71371 100644 --- a/backend/internal/service/usage_service.go +++ b/backend/internal/service/usage_service.go @@ -316,6 +316,25 @@ func (s *UsageService) GetUserUsageTrendByUserID(ctx context.Context, userID int return trend, nil } +// GetUsageTrendWithFilters returns trend data using the shared usage filter shape. +func (s *UsageService) GetUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, filters usagestats.UsageLogFilters) ([]usagestats.TrendDataPoint, error) { + type usageTrendWithFiltersRepo interface { + GetUsageTrendWithUsageFilters(ctx context.Context, startTime, endTime time.Time, granularity string, filters usagestats.UsageLogFilters) ([]usagestats.TrendDataPoint, error) + } + if filterRepo, ok := s.usageRepo.(usageTrendWithFiltersRepo); ok { + trend, err := filterRepo.GetUsageTrendWithUsageFilters(ctx, startTime, endTime, granularity, filters) + if err != nil { + return nil, fmt.Errorf("get usage trend with filters: %w", err) + } + return trend, nil + } + trend, err := s.usageRepo.GetUsageTrendWithFilters(ctx, startTime, endTime, granularity, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType) + if err != nil { + return nil, fmt.Errorf("get usage trend with filters: %w", err) + } + return trend, nil +} + // GetUserModelStats returns per-user model usage stats. func (s *UsageService) GetUserModelStats(ctx context.Context, userID int64, startTime, endTime time.Time) ([]usagestats.ModelStat, error) { stats, err := s.usageRepo.GetUserModelStats(ctx, userID, startTime, endTime) @@ -325,6 +344,55 @@ func (s *UsageService) GetUserModelStats(ctx context.Context, userID int64, star return stats, nil } +// GetModelStatsWithFiltersBySource returns model stats using the shared usage filter shape. +func (s *UsageService) GetModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters, modelSource string) ([]usagestats.ModelStat, error) { + normalizedSource := usagestats.NormalizeModelSource(modelSource) + type modelStatsWithUsageFiltersRepo interface { + GetModelStatsWithUsageFiltersBySource(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters, source string) ([]usagestats.ModelStat, error) + } + if filterRepo, ok := s.usageRepo.(modelStatsWithUsageFiltersRepo); ok { + stats, err := filterRepo.GetModelStatsWithUsageFiltersBySource(ctx, startTime, endTime, filters, normalizedSource) + if err != nil { + return nil, fmt.Errorf("get model stats with filters by source: %w", err) + } + return stats, nil + } + type modelStatsBySourceRepo interface { + GetModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8, source string) ([]usagestats.ModelStat, error) + } + if sourceRepo, ok := s.usageRepo.(modelStatsBySourceRepo); ok { + stats, err := sourceRepo.GetModelStatsWithFiltersBySource(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.RequestType, filters.Stream, filters.BillingType, normalizedSource) + if err != nil { + return nil, fmt.Errorf("get model stats with filters by source: %w", err) + } + return stats, nil + } + stats, err := s.usageRepo.GetModelStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.RequestType, filters.Stream, filters.BillingType) + if err != nil { + return nil, fmt.Errorf("get model stats with filters: %w", err) + } + return stats, nil +} + +// GetGroupStatsWithFilters returns group stats using the shared usage filter shape. +func (s *UsageService) GetGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters) ([]usagestats.GroupStat, error) { + type groupStatsWithUsageFiltersRepo interface { + GetGroupStatsWithUsageFilters(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters) ([]usagestats.GroupStat, error) + } + if filterRepo, ok := s.usageRepo.(groupStatsWithUsageFiltersRepo); ok { + stats, err := filterRepo.GetGroupStatsWithUsageFilters(ctx, startTime, endTime, filters) + if err != nil { + return nil, fmt.Errorf("get group stats with filters: %w", err) + } + return stats, nil + } + stats, err := s.usageRepo.GetGroupStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.RequestType, filters.Stream, filters.BillingType) + if err != nil { + return nil, fmt.Errorf("get group stats with filters: %w", err) + } + return stats, nil +} + // GetAPIKeyModelStats returns per-model usage stats for a specific API Key. func (s *UsageService) GetAPIKeyModelStats(ctx context.Context, apiKeyID int64, startTime, endTime time.Time) ([]usagestats.ModelStat, error) { stats, err := s.usageRepo.GetModelStatsWithFilters(ctx, startTime, endTime, 0, apiKeyID, 0, 0, nil, nil, nil) diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 9210e8c84f..44fbe29187 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -561,6 +561,7 @@ export interface SystemSettings { claude_oauth_system_prompt_blocks: string; enable_anthropic_cache_ttl_1h_injection: boolean; rewrite_message_cache_control: boolean; + enable_client_dateline_normalization: boolean; antigravity_user_agent_version: string; openai_codex_user_agent: string; // codex_cli_only 加固 @@ -811,6 +812,7 @@ export interface UpdateSettingsRequest { claude_oauth_system_prompt_blocks?: string; enable_anthropic_cache_ttl_1h_injection?: boolean; rewrite_message_cache_control?: boolean; + enable_client_dateline_normalization?: boolean; antigravity_user_agent_version?: string; openai_codex_user_agent?: string; // codex_cli_only 加固 diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts index f0aec3d6a6..5e43157c65 100644 --- a/frontend/src/api/usage.ts +++ b/frontend/src/api/usage.ts @@ -11,6 +11,8 @@ import type { PaginatedResponse, TrendDataPoint, ModelStat, + GroupStat, + UsageRequestType, UserErrorRequest, UserErrorRequestDetail, UserErrorListParams @@ -57,6 +59,14 @@ export interface TrendParams { start_date?: string end_date?: string granularity?: 'day' | 'hour' + api_key_id?: number + model?: string + group_id?: number + request_type?: UsageRequestType + stream?: boolean + billing_type?: number | null + billing_mode?: string | null + timezone?: string } export interface TrendResponse { @@ -91,6 +101,22 @@ export interface ApiKeyDailyUsageResponse { end_date: string } +export interface UsageDashboardSnapshotV2Params extends TrendParams { + include_trend?: boolean + include_model_stats?: boolean + include_group_stats?: boolean +} + +export interface UsageDashboardSnapshotV2Response { + generated_at: string + start_date: string + end_date: string + granularity: string + trend?: TrendDataPoint[] + models?: ModelStat[] + groups?: GroupStat[] +} + /** * List usage logs with optional filters * @param page - Page number (default: 1) @@ -141,10 +167,12 @@ export async function query( * @returns Usage statistics */ export async function getStats( - period: string = 'today', + paramsOrPeriod: (UsageQueryParams & { period?: string; timezone?: string }) | string = 'today', apiKeyId?: number ): Promise { - const params: Record = { period } + const params: Record = typeof paramsOrPeriod === 'string' + ? { period: paramsOrPeriod } + : { ...paramsOrPeriod } if (apiKeyId !== undefined) { params.api_key_id = apiKeyId @@ -251,6 +279,15 @@ export async function getDashboardTrend(params?: TrendParams): Promise { const { data } = await apiClient.get('/usage/dashboard/models', { params }) return data @@ -273,6 +310,16 @@ export async function getMyApiKeyDailyUsage( return data } +export async function getDashboardSnapshotV2( + params?: UsageDashboardSnapshotV2Params +): Promise { + const { data } = await apiClient.get( + '/usage/dashboard/snapshot-v2', + { params } + ) + return data +} + export interface BatchApiKeyUsageStats { api_key_id: number today_actual_cost: number @@ -308,11 +355,9 @@ export async function getDashboardApiKeysUsage( } export async function listMyErrorRequests( - params: UserErrorListParams, - config: { signal?: AbortSignal } = {} + params: UserErrorListParams ): Promise> { const { data } = await apiClient.get>('/usage/errors', { - ...config, params }) return data @@ -335,10 +380,11 @@ export const usageAPI = { getDashboardTrend, getDashboardModels, getMyApiKeyDailyUsage, + getDashboardSnapshotV2, getDashboardApiKeysUsage, // Error requests listMyErrorRequests, - getMyErrorDetail, + getMyErrorDetail } export default usageAPI diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 4df73c5309..d6cf0d807b 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -560,18 +560,18 @@ Google One - 个人账号,享受 Google One 订阅配额 + {{ t('admin.accounts.gemini.oauthType.googleOneDesc') }}
- 推荐个人用户 + {{ t('admin.accounts.gemini.oauthType.badges.individuals') }} - 无需 GCP + {{ t('admin.accounts.gemini.oauthType.badges.noGcp') }}
@@ -603,10 +603,10 @@ GCP Code Assist - 企业级,需要 GCP 项目 + {{ t('admin.accounts.gemini.oauthType.codeAssistDesc') }}
- 需要激活 GCP 项目并绑定信用卡 + {{ t('admin.accounts.gemini.oauthType.codeAssistRequirement') }} - 企业用户 + {{ t('admin.accounts.gemini.oauthType.badges.enterprise') }} - 高并发 + {{ t('admin.accounts.gemini.oauthType.badges.highConcurrency') }}
@@ -648,7 +648,13 @@ > - {{ showAdvancedOAuth ? '隐藏' : '显示' }}高级选项(自建 OAuth Client) + + {{ + showAdvancedOAuth + ? t('admin.accounts.gemini.oauthType.hideAdvanced') + : t('admin.accounts.gemini.oauthType.showAdvanced') + }} + @@ -3072,7 +3078,7 @@ rel="noreferrer" class="text-sm text-blue-600 hover:underline dark:text-blue-400" > - 修改归属地 + {{ t('admin.accounts.gemini.setupGuide.links.countryChange') }} · {{ - t('admin.accounts.oauth.openai.mobileRefreshTokenAuth', '手动输入 Mobile RT') + t('admin.accounts.oauth.openai.mobileRefreshTokenAuth') }}