diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 415d18dcaa..896e5b6beb 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -573,7 +573,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 @@ -590,7 +590,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, @@ -622,6 +621,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, @@ -633,7 +633,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 @@ -648,8 +648,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/types.go b/backend/internal/handler/dto/types.go index d8748baa34..8659a1e788 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -480,6 +480,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"` @@ -515,7 +517,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/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..41bbc524e4 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -2506,7 +2506,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 +2532,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 +2548,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/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/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/admin/usage/UsageStatsCards.vue b/frontend/src/components/admin/usage/UsageStatsCards.vue index 415c7cc186..c8c9b287fe 100644 --- a/frontend/src/components/admin/usage/UsageStatsCards.vue +++ b/frontend/src/components/admin/usage/UsageStatsCards.vue @@ -68,9 +68,14 @@ ${{ (stats?.total_actual_cost || 0).toFixed(4) }}

- {{ t('usage.accountCost') }} ${{ (stats?.total_account_cost || 0).toFixed(4) }} - · - {{ t('usage.standardCost') }} ${{ (stats?.total_cost || 0).toFixed(4) }} + + + {{ t('usage.standardCost') }} + ${{ (stats?.total_cost || 0).toFixed(4) }} +

@@ -84,14 +89,30 @@ diff --git a/frontend/src/views/user/__tests__/UsageView.spec.ts b/frontend/src/views/user/__tests__/UsageView.spec.ts index 011b96c8bf..555ea4b0f6 100644 --- a/frontend/src/views/user/__tests__/UsageView.spec.ts +++ b/frontend/src/views/user/__tests__/UsageView.spec.ts @@ -1,13 +1,26 @@ -import { describe, expect, it, vi, beforeEach } from 'vitest' +import { beforeEach, describe, expect, it, vi } from 'vitest' import { flushPromises, mount } from '@vue/test-utils' -import { nextTick } from 'vue' import UsageView from '../UsageView.vue' -const { query, getStatsByDateRange, list, showError, showWarning, showSuccess, showInfo } = vi.hoisted(() => ({ +const { + query, + getStats, + getDashboardModels, + getDashboardSnapshotV2, + list, + getAvailable, + showError, + showWarning, + showSuccess, + showInfo, +} = vi.hoisted(() => ({ query: vi.fn(), - getStatsByDateRange: vi.fn(), + getStats: vi.fn(), + getDashboardModels: vi.fn(), + getDashboardSnapshotV2: vi.fn(), list: vi.fn(), + getAvailable: vi.fn(), showError: vi.fn(), showWarning: vi.fn(), showSuccess: vi.fn(), @@ -15,62 +28,55 @@ const { query, getStatsByDateRange, list, showError, showWarning, showSuccess, s })) const messages: Record = { - 'usage.costDetails': 'Cost Breakdown', - 'admin.usage.inputCost': 'Input Cost', - 'admin.usage.outputCost': 'Output Cost', - 'admin.usage.cacheCreationCost': 'Cache Creation Cost', - 'admin.usage.cacheReadCost': 'Cache Read Cost', - 'usage.inputTokenPrice': 'Input price', - 'usage.outputTokenPrice': 'Output price', - 'usage.perMillionTokens': '/ 1M tokens', - 'usage.serviceTier': 'Service tier', - 'usage.serviceTierPriority': 'Fast', - 'usage.serviceTierFlex': 'Flex', - 'usage.serviceTierStandard': 'Standard', - 'usage.rate': 'Rate', - 'usage.original': 'Original', - 'usage.billed': 'Billed', - 'usage.allApiKeys': 'All API Keys', - 'usage.apiKeyFilter': 'API Key', - 'usage.model': 'Model', - 'usage.reasoningEffort': 'Reasoning Effort', - 'usage.type': 'Type', - 'usage.tokens': 'Tokens', - 'usage.cost': 'Cost', - 'usage.firstToken': 'First Token', - 'usage.duration': 'Duration', - 'usage.time': 'Time', - 'usage.userAgent': 'User Agent', - 'usage.imageUnit': ' images', - 'usage.imageCount': 'Image count', - 'usage.imageBillingSize': 'Billing size', - 'usage.imageInputSize': 'Input size', - 'usage.imageOutputSize': 'Output size', - 'usage.imageSizeSource': 'Size source', - 'usage.imageSizeBreakdown': 'Size breakdown', - 'usage.imageSizeSourceOutput': 'Upstream output', - 'usage.imageSizeSourceInput': 'Request input', - 'usage.imageSizeSourceDefault': 'Default billing tier', - 'usage.imageSizeSourceLegacy': 'Legacy record', - 'usage.imageSizeSourceMissing': 'Not recorded', - 'usage.imageSizeNotRecorded': 'not recorded', - 'usage.imageSizeLegacyUnstandardized': 'legacy unstandardized', - 'usage.imageSizeUnknown': 'unknown', - 'usage.imageUnitPrice': 'Per-image price', - 'usage.imageTotalPrice': 'Image total price', + 'admin.dashboard.timeRange': 'Time range', + 'admin.dashboard.granularity': 'Granularity', + 'admin.dashboard.day': 'Day', + 'admin.dashboard.hour': 'Hour', + 'admin.users.columnSettings': 'Columns', + 'admin.usage.group': 'Group', + 'admin.usage.billingType': 'Billing type', + 'admin.usage.billingMode': 'Billing mode', + 'admin.usage.allTypes': 'All types', + 'admin.usage.allBillingTypes': 'All billing types', + 'admin.usage.billingTypeBalance': 'Balance', + 'admin.usage.billingTypeSubscription': 'Subscription', + 'admin.usage.allBillingModes': 'All billing modes', 'admin.usage.billingModeToken': 'Token', 'admin.usage.billingModePerRequest': 'Per request', 'admin.usage.billingModeImage': 'Image', + 'admin.usage.allGroups': 'All groups', + 'admin.usage.allModels': 'All models', + 'usage.allApiKeys': 'All API Keys', + 'usage.apiKeyFilter': 'API Key', + 'usage.model': 'Model', + 'usage.type': 'Type', + 'usage.ws': 'WS', + 'usage.stream': 'Stream', + 'usage.sync': 'Sync', + 'usage.exporting': 'Exporting', + 'usage.exportCsv': 'Export CSV', + 'usage.failedToLoad': 'Failed to load', + 'usage.noDataToExport': 'No data', + 'usage.preparingExport': 'Preparing export', + 'usage.exportSuccess': 'Export success', + 'usage.exportFailed': 'Export failed', + 'common.refresh': 'Refresh', + 'common.reset': 'Reset', } vi.mock('@/api', () => ({ usageAPI: { query, - getStatsByDateRange, + getStats, + getDashboardModels, + getDashboardSnapshotV2, }, keysAPI: { list, }, + userGroupsAPI: { + getAvailable, + }, })) vi.mock('@/stores/app', () => ({ @@ -87,178 +93,131 @@ vi.mock('vue-i18n', async () => { } }) -const AppLayoutStub = { template: '
' } -const TablePageLayoutStub = { - template: '
', -} -const DataTableStub = { - props: ['data'], - template: ` -
-
- - - -
-
- `, +const simpleStub = { template: '
' } +const chartStub = { template: '
' } + +const usageLog = { + id: 1, + request_id: 'req-user-export', + actual_cost: 0.092883, + total_cost: 0.092883, + rate_multiplier: 1, + service_tier: 'priority', + input_cost: 0.020285, + output_cost: 0.00303, + cache_creation_cost: 0.000001, + cache_read_cost: 0.069568, + input_tokens: 4057, + output_tokens: 101, + cache_creation_tokens: 4, + cache_read_tokens: 278272, + cache_creation_5m_tokens: 0, + cache_creation_1h_tokens: 0, + image_count: 0, + image_size: null, + first_token_ms: 12, + duration_ms: 345, + created_at: '2026-03-08T00:00:00Z', + model: 'gpt-5.4', + reasoning_effort: null, + ip_address: '203.0.113.10', + api_key: { name: 'demo-key' }, + billing_mode: 'token', + request_type: 'sync', + stream: false, } -describe('user UsageView tooltip', () => { +function mountUsageView() { + return mount(UsageView, { + global: { + stubs: { + AppLayout: simpleStub, + Pagination: true, + Select: true, + DateRangePicker: true, + Icon: true, + UsageStatsCards: chartStub, + UsageTable: chartStub, + ModelDistributionChart: chartStub, + GroupDistributionChart: chartStub, + EndpointDistributionChart: chartStub, + TokenUsageTrend: chartStub, + }, + }, + }) +} + +describe('user UsageView', () => { beforeEach(() => { query.mockReset() - getStatsByDateRange.mockReset() + getStats.mockReset() + getDashboardModels.mockReset() + getDashboardSnapshotV2.mockReset() list.mockReset() + getAvailable.mockReset() showError.mockReset() showWarning.mockReset() showSuccess.mockReset() showInfo.mockReset() - vi.spyOn(HTMLElement.prototype, 'getBoundingClientRect').mockReturnValue({ - x: 0, - y: 0, - top: 20, - left: 20, - right: 120, - bottom: 40, - width: 100, - height: 20, - toJSON: () => ({}), - } as DOMRect) - - ;(globalThis as any).ResizeObserver = class { - observe() {} - disconnect() {} - } + query.mockResolvedValue({ items: [usageLog], total: 1, pages: 1 }) + getStats.mockResolvedValue({ + total_requests: 1, + total_input_tokens: 10, + total_output_tokens: 20, + total_cache_tokens: 0, + total_tokens: 30, + total_cost: 0.1, + total_actual_cost: 0.08, + average_duration_ms: 12, + endpoints: [], + upstream_endpoints: [], + endpoint_paths: [], + }) + getDashboardModels.mockResolvedValue({ + models: [{ model: 'gpt-5.4', requests: 1, input_tokens: 10, output_tokens: 20, cache_creation_tokens: 0, cache_read_tokens: 0, total_tokens: 30, cost: 0.1, actual_cost: 0.08 }], + start_date: '2026-03-08', + end_date: '2026-03-08', + }) + getDashboardSnapshotV2.mockResolvedValue({ + generated_at: '2026-03-08T00:00:00Z', + start_date: '2026-03-08', + end_date: '2026-03-08', + granularity: 'hour', + trend: [], + groups: [], + }) + list.mockResolvedValue({ items: [{ id: 1, name: 'demo-key' }] }) + getAvailable.mockResolvedValue([{ id: 1, name: 'default' }]) }) - it('shows fast service tier and unit prices in user tooltip', async () => { - query.mockResolvedValue({ - items: [ - { - request_id: 'req-user-1', - actual_cost: 0.092883, - total_cost: 0.092883, - rate_multiplier: 1, - service_tier: 'priority', - input_cost: 0.020285, - output_cost: 0.00303, - cache_creation_cost: 0, - cache_read_cost: 0.069568, - input_tokens: 4057, - output_tokens: 101, - cache_creation_tokens: 0, - cache_read_tokens: 278272, - cache_creation_5m_tokens: 0, - cache_creation_1h_tokens: 0, - image_count: 0, - image_size: null, - first_token_ms: null, - duration_ms: 1, - created_at: '2026-03-08T00:00:00Z', - }, - ], - total: 1, - pages: 1, - }) - getStatsByDateRange.mockResolvedValue({ - total_requests: 1, - total_tokens: 100, - total_cost: 0.1, - avg_duration_ms: 1, - }) - list.mockResolvedValue({ items: [] }) - - const wrapper = mount(UsageView, { - global: { - stubs: { - AppLayout: AppLayoutStub, - TablePageLayout: TablePageLayoutStub, - Pagination: true, - EmptyState: true, - Select: true, - DateRangePicker: true, - DataTable: DataTableStub, - Icon: true, - Teleport: true, - }, - }, - }) - + it('loads logs, stats, model stats, and snapshot on first render', async () => { + mountUsageView() await flushPromises() - await nextTick() - const setupState = (wrapper.vm as any).$?.setupState - setupState.tooltipData = { - request_id: 'req-user-1', - actual_cost: 0.092883, - total_cost: 0.092883, - rate_multiplier: 1, - service_tier: 'priority', - input_cost: 0.020285, - output_cost: 0.00303, - cache_creation_cost: 0, - cache_read_cost: 0.069568, - input_tokens: 4057, - output_tokens: 101, - } - setupState.tooltipVisible = true - await nextTick() - - const text = wrapper.text() - expect(text).toContain('Service tier') - expect(text).toContain('Fast') - expect(text).toContain('Rate') - expect(text).toContain('1.00x') - expect(text).toContain('Billed') - expect(text).toContain('$0.092883') - expect(text).toContain('$5.0000 / 1M tokens') - expect(text).toContain('$30.0000 / 1M tokens') + expect(query).toHaveBeenCalled() + expect(getStats).toHaveBeenCalled() + expect(getDashboardModels).toHaveBeenCalled() + expect(getDashboardSnapshotV2).toHaveBeenCalledWith(expect.objectContaining({ + include_trend: true, + include_model_stats: false, + include_group_stats: true, + })) + expect(list).toHaveBeenCalledWith(1, 100) + expect(getAvailable).toHaveBeenCalled() }) - it('exports csv with input and output unit price columns', async () => { - const exportedLogs = [ - { - request_id: 'req-user-export', - actual_cost: 0.092883, - total_cost: 0.092883, - rate_multiplier: 1, - service_tier: 'priority', - input_cost: 0.020285, - output_cost: 0.00303, - cache_creation_cost: 0.000001, - cache_read_cost: 0.069568, - input_tokens: 4057, - output_tokens: 101, - cache_creation_tokens: 4, - cache_read_tokens: 278272, - cache_creation_5m_tokens: 0, - cache_creation_1h_tokens: 0, - image_count: 0, - image_size: null, - first_token_ms: 12, - duration_ms: 345, - created_at: '2026-03-08T00:00:00Z', - model: 'gpt-5.4', - reasoning_effort: null, - api_key: { name: 'demo-key' }, - }, - ] - - query.mockResolvedValue({ - items: exportedLogs, - total: 1, - pages: 1, - }) - getStatsByDateRange.mockResolvedValue({ - total_requests: 1, - total_tokens: 100, - total_cost: 0.1, - avg_duration_ms: 1, - }) - list.mockResolvedValue({ items: [] }) + it('exports csv with current filters and without admin-only fields', async () => { + const wrapper = mountUsageView() + await flushPromises() let exportedBlob: Blob | null = null + let csvContent = '' + const OriginalBlob = globalThis.Blob + vi.stubGlobal('Blob', vi.fn((parts: BlobPart[], options?: BlobPropertyBag) => { + csvContent = parts.map((part) => String(part)).join('') + return new OriginalBlob(parts, options) + })) const originalCreateObjectURL = window.URL.createObjectURL const originalRevokeObjectURL = window.URL.revokeObjectURL window.URL.createObjectURL = vi.fn((blob: Blob | MediaSource) => { @@ -268,146 +227,38 @@ describe('user UsageView tooltip', () => { window.URL.revokeObjectURL = vi.fn(() => {}) as typeof window.URL.revokeObjectURL const clickSpy = vi.spyOn(HTMLAnchorElement.prototype, 'click').mockImplementation(() => {}) - const wrapper = mount(UsageView, { - global: { - stubs: { - AppLayout: AppLayoutStub, - TablePageLayout: TablePageLayoutStub, - Pagination: true, - EmptyState: true, - Select: true, - DateRangePicker: true, - DataTable: DataTableStub, - Icon: true, - Teleport: true, - }, - }, - }) - - await flushPromises() - - const setupState = (wrapper.vm as any).$?.setupState - await setupState.exportToCSV() + await (wrapper.vm as any).exportToCSV() expect(exportedBlob).not.toBeNull() - const hasSortedExportQuery = query.mock.calls.some((call) => { - const params = call[0] as Record | undefined - const config = call[1] - return ( - params?.page_size === 100 && - params?.sort_by === 'created_at' && - params?.sort_order === 'desc' && - config === undefined - ) - }) - expect(hasSortedExportQuery).toBe(true) + expect(query).toHaveBeenCalledWith(expect.objectContaining({ + page_size: 100, + sort_by: 'created_at', + sort_order: 'desc', + })) expect(clickSpy).toHaveBeenCalled() expect(showSuccess).toHaveBeenCalled() + expect(csvContent).toContain('IP Address') + expect(csvContent).toContain('203.0.113.10') + expect(csvContent).toContain('Billed Cost') + expect(csvContent).toContain('Original Cost') + expect(csvContent).not.toContain('Upstream Endpoint') + expect(csvContent).not.toContain('account_cost') + expect(csvContent).not.toContain('account_rate_multiplier') window.URL.createObjectURL = originalCreateObjectURL window.URL.revokeObjectURL = originalRevokeObjectURL + vi.unstubAllGlobals() clickSpy.mockRestore() }) it('exports historical image rows with image billing mode derived from image_count', async () => { - const exportedLogs = [ - { - request_id: 'req-user-export-legacy-image', - actual_cost: 0.2, - total_cost: 0.2, - rate_multiplier: 1, - service_tier: null, - input_cost: 0, - output_cost: 0, - cache_creation_cost: 0, - cache_read_cost: 0, - input_tokens: 0, - output_tokens: 0, - cache_creation_tokens: 0, - cache_read_tokens: 0, - cache_creation_5m_tokens: 0, - cache_creation_1h_tokens: 0, - image_count: 1, - image_size: null, - billing_mode: null, - first_token_ms: null, - duration_ms: 345, - created_at: '2026-03-08T00:00:00Z', - model: 'gpt-image-2', - reasoning_effort: null, - api_key: { name: 'demo-key' }, - }, - ] - - query.mockResolvedValue({ - items: exportedLogs, - total: 1, - pages: 1, - }) - getStatsByDateRange.mockResolvedValue({ - total_requests: 1, - total_tokens: 0, - total_cost: 0.2, - avg_duration_ms: 1, - }) - list.mockResolvedValue({ items: [] }) - - let exportedBlob: Blob | null = null - const originalCreateObjectURL = window.URL.createObjectURL - const originalRevokeObjectURL = window.URL.revokeObjectURL - window.URL.createObjectURL = vi.fn((blob: Blob | MediaSource) => { - exportedBlob = blob as Blob - return 'blob:usage-export' - }) as typeof window.URL.createObjectURL - window.URL.revokeObjectURL = vi.fn(() => {}) as typeof window.URL.revokeObjectURL - const clickSpy = vi.spyOn(HTMLAnchorElement.prototype, 'click').mockImplementation(() => {}) - - const wrapper = mount(UsageView, { - global: { - stubs: { - AppLayout: AppLayoutStub, - TablePageLayout: TablePageLayoutStub, - Pagination: true, - EmptyState: true, - Select: true, - DateRangePicker: true, - DataTable: DataTableStub, - Icon: true, - Teleport: true, - }, - }, - }) - - await flushPromises() - - const setupState = (wrapper.vm as any).$?.setupState - await setupState.exportToCSV() - - expect(exportedBlob).not.toBeNull() - const csv = await new Promise((resolve, reject) => { - const reader = new FileReader() - reader.onload = () => resolve(String(reader.result)) - reader.onerror = () => reject(reader.error) - reader.readAsText(exportedBlob as Blob) - }) - expect(csv).toContain('Billing Mode') - expect(csv).toContain('Image') - expect(csv).not.toContain(',Token,0,0,0,0,') - - window.URL.createObjectURL = originalCreateObjectURL - window.URL.revokeObjectURL = originalRevokeObjectURL - clickSpy.mockRestore() - }) - - it('does not display a 2K fallback for historical image rows with missing size', async () => { query.mockResolvedValue({ items: [ { - request_id: 'req-user-legacy-missing-image', + ...usageLog, + request_id: 'req-user-export-legacy-image', actual_cost: 0.2, total_cost: 0.2, - rate_multiplier: 1, - service_tier: null, input_cost: 0, output_cost: 0, cache_creation_cost: 0, @@ -416,125 +267,40 @@ describe('user UsageView tooltip', () => { output_tokens: 0, cache_creation_tokens: 0, cache_read_tokens: 0, - cache_creation_5m_tokens: 0, - cache_creation_1h_tokens: 0, image_count: 1, - image_size: null, - image_input_size: null, - image_output_size: null, - image_size_source: null, - image_size_breakdown: null, - billing_mode: null, - first_token_ms: null, - duration_ms: 1, - created_at: '2026-03-08T00:00:00Z', model: 'gpt-image-2', + billing_mode: null, + ip_address: null, }, ], total: 1, pages: 1, }) - getStatsByDateRange.mockResolvedValue({ - total_requests: 1, - total_tokens: 0, - total_cost: 0.2, - avg_duration_ms: 1, - }) - list.mockResolvedValue({ items: [] }) - - const wrapper = mount(UsageView, { - global: { - stubs: { - AppLayout: AppLayoutStub, - TablePageLayout: TablePageLayoutStub, - Pagination: true, - EmptyState: true, - Select: true, - DateRangePicker: true, - DataTable: DataTableStub, - Icon: true, - Teleport: true, - }, - }, - }) - - await flushPromises() - await nextTick() - - const text = wrapper.text() - expect(text).toContain('Image') - expect(text).toContain('not recorded') - expect(text).not.toContain('(2K)') - }) - - it('shows image billing metadata in the user cost tooltip', async () => { - query.mockResolvedValue({ - items: [], - total: 0, - pages: 0, - }) - getStatsByDateRange.mockResolvedValue({ - total_requests: 0, - total_tokens: 0, - total_cost: 0, - avg_duration_ms: 0, - }) - list.mockResolvedValue({ items: [] }) - - const wrapper = mount(UsageView, { - global: { - stubs: { - AppLayout: AppLayoutStub, - TablePageLayout: TablePageLayoutStub, - Pagination: true, - EmptyState: true, - Select: true, - DateRangePicker: true, - DataTable: DataTableStub, - Icon: true, - Teleport: true, - }, - }, - }) + const wrapper = mountUsageView() await flushPromises() - const setupState = (wrapper.vm as any).$?.setupState - setupState.tooltipData = { - request_id: 'req-user-output-image', - actual_cost: 0.8, - total_cost: 0.8, - rate_multiplier: 1, - service_tier: null, - input_cost: 0, - output_cost: 0, - cache_creation_cost: 0, - cache_read_cost: 0, - input_tokens: 0, - output_tokens: 0, - cache_creation_tokens: 0, - cache_read_tokens: 0, - billing_mode: null, - image_count: 2, - image_size: '4K', - image_input_size: '1024x1024', - image_output_size: '3840x2160', - image_size_source: 'output', - image_size_breakdown: { '4K': 2 }, - } - setupState.tooltipVisible = true - await nextTick() + let csvContent = '' + const OriginalBlob = globalThis.Blob + vi.stubGlobal('Blob', vi.fn((parts: BlobPart[], options?: BlobPropertyBag) => { + csvContent = parts.map((part) => String(part)).join('') + return new OriginalBlob(parts, options) + })) + const originalCreateObjectURL = window.URL.createObjectURL + const originalRevokeObjectURL = window.URL.revokeObjectURL + window.URL.createObjectURL = vi.fn(() => 'blob:usage-export') as typeof window.URL.createObjectURL + window.URL.revokeObjectURL = vi.fn(() => {}) as typeof window.URL.revokeObjectURL + const clickSpy = vi.spyOn(HTMLAnchorElement.prototype, 'click').mockImplementation(() => {}) - const text = wrapper.text() - expect(text).toContain('Image count') - expect(text).toContain('Billing size') - expect(text).toContain('4K') - expect(text).toContain('Size source') - expect(text).toContain('Upstream output') - expect(text).toContain('Input size') - expect(text).toContain('1024x1024') - expect(text).toContain('Output size') - expect(text).toContain('3840x2160') - expect(text).toContain('4K x 2') + await (wrapper.vm as any).exportToCSV() + + expect(csvContent).toContain('Billing Mode') + expect(csvContent).toContain('Image') + expect(csvContent).not.toContain(',Token,0,0,0,0,') + + window.URL.createObjectURL = originalCreateObjectURL + window.URL.revokeObjectURL = originalRevokeObjectURL + vi.unstubAllGlobals() + clickSpy.mockRestore() }) })