From cafc95c3e277536beb0164607250b4ab8351868b Mon Sep 17 00:00:00 2001
From: PMExtra
Date: Mon, 27 Apr 2026 01:40:26 +0800
Subject: [PATCH] feat: align user usage analytics with admin
---
backend/internal/handler/dto/mappers.go | 10 +-
.../handler/dto/mappers_usage_test.go | 42 +-
backend/internal/handler/dto/types.go | 4 +-
backend/internal/handler/usage_handler.go | 418 +++--
.../usage_handler_request_type_test.go | 252 ++-
.../pkg/usagestats/usage_log_types.go | 24 +-
backend/internal/repository/usage_log_repo.go | 135 +-
.../usage_log_repo_request_type_test.go | 139 ++
backend/internal/server/api_contract_test.go | 59 +-
backend/internal/server/routes/user.go | 1 +
backend/internal/service/usage_service.go | 68 +
frontend/src/api/usage.ts | 58 +-
.../admin/usage/UsageStatsCards.vue | 29 +-
.../src/components/admin/usage/UsageTable.vue | 44 +-
.../charts/EndpointDistributionChart.vue | 15 +-
.../charts/GroupDistributionChart.vue | 23 +-
.../charts/ModelDistributionChart.vue | 25 +-
.../charts/UserBreakdownSubTable.vue | 16 +-
.../__tests__/GroupDistributionChart.spec.ts | 19 +
.../__tests__/ModelDistributionChart.spec.ts | 19 +
frontend/src/types/index.ts | 13 +-
frontend/src/views/user/UsageView.vue | 1500 +++++++----------
.../views/user/__tests__/UsageView.spec.ts | 636 +++----
23 files changed, 1910 insertions(+), 1639 deletions(-)
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.accountCost') }} ${{ totalAccountCost.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()
})
})