Merge pull request #2007 from PMExtra/feature/user-usage-admin-parity

feat: align user usage analytics with admin
This commit is contained in:
Wesley Liddick
2026-06-30 17:10:09 +08:00
committed by GitHub
23 changed files with 1910 additions and 1639 deletions
+6 -4
View File
@@ -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,
@@ -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()
+3 -1
View File
@@ -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 最小账号信息(避免泄露敏感字段)
+257 -161
View File
@@ -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"`
@@ -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)
}
}
@@ -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
}
+111 -24
View File
@@ -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
@@ -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}
+57 -2
View File
@@ -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")
+1
View File
@@ -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)
}
+68
View File
@@ -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)
+52 -6
View File
@@ -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<UsageStatsResponse> {
const params: Record<string, unknown> = { period }
const params: Record<string, unknown> = 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<TrendResp
export async function getDashboardModels(params?: {
start_date?: string
end_date?: string
api_key_id?: number
model?: string
model_source?: 'requested'
group_id?: number
request_type?: UsageRequestType
stream?: boolean
billing_type?: number | null
billing_mode?: string | null
timezone?: string
}): Promise<ModelStatsResponse> {
const { data } = await apiClient.get<ModelStatsResponse>('/usage/dashboard/models', { params })
return data
@@ -273,6 +310,16 @@ export async function getMyApiKeyDailyUsage(
return data
}
export async function getDashboardSnapshotV2(
params?: UsageDashboardSnapshotV2Params
): Promise<UsageDashboardSnapshotV2Response> {
const { data } = await apiClient.get<UsageDashboardSnapshotV2Response>(
'/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<PaginatedResponse<UserErrorRequest>> {
const { data } = await apiClient.get<PaginatedResponse<UserErrorRequest>>('/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
@@ -68,9 +68,14 @@
${{ (stats?.total_actual_cost || 0).toFixed(4) }}
</p>
<p class="text-xs text-gray-400">
<span class="text-orange-500">{{ t('usage.accountCost') }} ${{ (stats?.total_account_cost || 0).toFixed(4) }}</span>
<span> · </span>
<span>{{ t('usage.standardCost') }} ${{ (stats?.total_cost || 0).toFixed(4) }}</span>
<template v-if="showAccountCost && totalAccountCost != null">
<span class="text-orange-500">{{ t('usage.accountCost') }} ${{ totalAccountCost.toFixed(4) }}</span>
<span> · </span>
</template>
<span>
{{ t('usage.standardCost') }}
<span :class="{ 'line-through': strikeStandardCost }">${{ (stats?.total_cost || 0).toFixed(4) }}</span>
</span>
</p>
</div>
</div>
@@ -84,14 +89,30 @@
</template>
<script setup lang="ts">
import { computed } from 'vue'
import { useI18n } from 'vue-i18n'
import type { AdminUsageStatsResponse } from '@/api/admin/usage'
import type { UsageStatsResponse } from '@/types'
import Icon from '@/components/icons/Icon.vue'
defineProps<{ stats: AdminUsageStatsResponse | null }>()
const props = withDefaults(defineProps<{
stats: (AdminUsageStatsResponse | UsageStatsResponse) | null
showAccountCost?: boolean
strikeStandardCost?: boolean
}>(), {
showAccountCost: true,
strikeStandardCost: false,
})
const { t } = useI18n()
const totalAccountCost = computed(() => {
const stats = props.stats as (AdminUsageStatsResponse & { total_account_cost?: number }) | null
return stats?.total_account_cost ?? null
})
const showAccountCost = computed(() => props.showAccountCost)
const strikeStandardCost = computed(() => props.strikeStandardCost)
const formatDuration = (ms: number) =>
ms < 1000 ? `${ms.toFixed(0)}ms` : `${(ms / 1000).toFixed(2)}s`
@@ -68,7 +68,7 @@
<span class="font-medium text-gray-500 dark:text-gray-400">{{ t('usage.inbound') }}:</span>
<span class="ml-1">{{ row.inbound_endpoint?.trim() || '-' }}</span>
</div>
<div class="break-all text-gray-700 dark:text-gray-300">
<div v-if="showUpstreamEndpoint" class="break-all text-gray-700 dark:text-gray-300">
<span class="font-medium text-gray-500 dark:text-gray-400">{{ t('usage.upstream') }}:</span>
<span class="ml-1">{{ row.upstream_endpoint?.trim() || '-' }}</span>
</div>
@@ -163,7 +163,7 @@
</div>
</div>
</div>
<div v-if="row.account_rate_multiplier != null" class="mt-0.5 text-[11px] text-orange-500 dark:text-orange-400">
<div v-if="showAccountBilling && row.account_rate_multiplier != null" class="mt-0.5 text-[11px] text-orange-500 dark:text-orange-400">
A ${{ accountBilled(row).toFixed(6) }}
</div>
</div>
@@ -380,20 +380,22 @@
<span class="font-semibold text-green-400">${{ tooltipData?.actual_cost?.toFixed(6) || '0.000000' }}</span>
</div>
<!-- Account billing (separated from user billing) -->
<div class="flex items-center justify-between gap-6 border-t border-gray-700 pt-1.5">
<span class="text-gray-400">{{ t('usage.accountMultiplier') }}</span>
<span class="font-semibold text-blue-400">{{ formatMultiplier(tooltipData?.account_rate_multiplier ?? 1) }}x</span>
</div>
<div class="flex items-center justify-between gap-6">
<span class="text-gray-400">{{ t('usage.accountBilled') }}</span>
<span class="font-semibold text-green-400">
${{ accountBilled({
total_cost: tooltipData?.total_cost,
account_stats_cost: tooltipData?.account_stats_cost,
account_rate_multiplier: tooltipData?.account_rate_multiplier,
}).toFixed(6) }}
</span>
</div>
<template v-if="showAccountBilling">
<div class="flex items-center justify-between gap-6 border-t border-gray-700 pt-1.5">
<span class="text-gray-400">{{ t('usage.accountMultiplier') }}</span>
<span class="font-semibold text-blue-400">{{ formatMultiplier(tooltipData?.account_rate_multiplier ?? 1) }}x</span>
</div>
<div class="flex items-center justify-between gap-6">
<span class="text-gray-400">{{ t('usage.accountBilled') }}</span>
<span class="font-semibold text-green-400">
${{ accountBilled({
total_cost: tooltipData?.total_cost,
account_stats_cost: tooltipData?.account_stats_cost,
account_rate_multiplier: tooltipData?.account_rate_multiplier,
}).toFixed(6) }}
</span>
</div>
</template>
</div>
<div class="absolute right-full top-1/2 h-0 w-0 -translate-y-1/2 border-b-[6px] border-r-[6px] border-t-[6px] border-b-transparent border-r-gray-900 border-t-transparent dark:border-r-gray-800"></div>
</div>
@@ -449,19 +451,25 @@ interface Props {
serverSideSort?: boolean
defaultSortKey?: string
defaultSortOrder?: 'asc' | 'desc'
showAccountBilling?: boolean
showUpstreamEndpoint?: boolean
}
withDefaults(defineProps<Props>(), {
const props = withDefaults(defineProps<Props>(), {
loading: false,
serverSideSort: false,
defaultSortKey: '',
defaultSortOrder: 'asc'
defaultSortOrder: 'asc',
showAccountBilling: true,
showUpstreamEndpoint: true
})
defineEmits<{
userClick: [userID: number, email?: string]
sort: [key: string, order: 'asc' | 'desc']
}>()
const { t } = useI18n()
const showAccountBilling = props.showAccountBilling
const showUpstreamEndpoint = props.showUpstreamEndpoint
// Tooltip state - cost
const tooltipVisible = ref(false)
@@ -89,13 +89,14 @@
<tbody>
<template v-for="item in displayEndpointStats" :key="item.endpoint">
<tr
class="border-t border-gray-100 cursor-pointer transition-colors hover:bg-gray-50 dark:border-gray-700 dark:hover:bg-dark-700/40"
@click="toggleBreakdown(item.endpoint)"
class="border-t border-gray-100 transition-colors dark:border-gray-700"
:class="enableBreakdown ? 'cursor-pointer hover:bg-gray-50 dark:hover:bg-dark-700/40' : ''"
@click="enableBreakdown && toggleBreakdown(item.endpoint)"
>
<td class="max-w-[180px] truncate py-1.5 font-medium text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300" :title="item.endpoint">
<td class="max-w-[180px] truncate py-1.5 font-medium" :class="enableBreakdown ? 'text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300' : 'text-gray-900 dark:text-white'" :title="item.endpoint">
<span class="inline-flex items-center gap-1">
<svg v-if="expandedKey === item.endpoint" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
<svg v-if="enableBreakdown && expandedKey === item.endpoint" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else-if="enableBreakdown" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
{{ item.endpoint }}
</span>
</td>
@@ -159,6 +160,7 @@ const props = withDefaults(
source?: EndpointSource
showMetricToggle?: boolean
showSourceToggle?: boolean
enableBreakdown?: boolean
startDate?: string
endDate?: string
filters?: Record<string, any>
@@ -171,7 +173,8 @@ const props = withDefaults(
metric: 'tokens',
source: 'inbound',
showMetricToggle: false,
showSourceToggle: false
showSourceToggle: false,
enableBreakdown: true
}
)
@@ -45,7 +45,7 @@
<th class="pb-2 text-right">{{ t('admin.dashboard.requests') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.tokens') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.actual') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.accountCost') }}</th>
<th v-if="showAccountCost" class="pb-2 text-right">{{ t('admin.dashboard.accountCost') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.standard') }}</th>
</tr>
</thead>
@@ -53,17 +53,17 @@
<template v-for="group in displayGroupStats" :key="group.group_id">
<tr
class="border-t border-gray-100 transition-colors dark:border-gray-700"
:class="group.group_id > 0 ? 'cursor-pointer hover:bg-gray-50 dark:hover:bg-dark-700/40' : ''"
@click="group.group_id > 0 && toggleBreakdown('group', group.group_id)"
:class="enableBreakdown && group.group_id > 0 ? 'cursor-pointer hover:bg-gray-50 dark:hover:bg-dark-700/40' : ''"
@click="enableBreakdown && group.group_id > 0 && toggleBreakdown('group', group.group_id)"
>
<td
class="max-w-[100px] truncate py-1.5 font-medium"
:class="group.group_id > 0 ? 'text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300' : 'text-gray-900 dark:text-white'"
:class="enableBreakdown && group.group_id > 0 ? 'text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300' : 'text-gray-900 dark:text-white'"
:title="group.group_name || String(group.group_id)"
>
<span class="inline-flex items-center gap-1">
<svg v-if="group.group_id > 0 && expandedKey === `group-${group.group_id}`" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else-if="group.group_id > 0" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
<svg v-if="enableBreakdown && group.group_id > 0 && expandedKey === `group-${group.group_id}`" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else-if="enableBreakdown && group.group_id > 0" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
{{ group.group_name || t('admin.dashboard.noGroup') }}
</span>
</td>
@@ -76,7 +76,7 @@
<td class="py-1.5 text-right text-green-600 dark:text-green-400">
${{ formatCost(group.actual_cost) }}
</td>
<td class="py-1.5 text-right text-orange-500 dark:text-orange-400">
<td v-if="showAccountCost" class="py-1.5 text-right text-orange-500 dark:text-orange-400">
${{ formatCost(group.account_cost) }}
</td>
<td class="py-1.5 text-right text-gray-400 dark:text-gray-500">
@@ -85,10 +85,11 @@
</tr>
<!-- User breakdown sub-rows -->
<tr v-if="expandedKey === `group-${group.group_id}`">
<td colspan="6" class="p-0">
<td :colspan="distributionColspan" class="p-0">
<UserBreakdownSubTable
:items="breakdownItems"
:loading="breakdownLoading"
:show-account-cost="showAccountCost"
/>
</td>
</tr>
@@ -127,6 +128,8 @@ const props = withDefaults(defineProps<{
loading?: boolean
metric?: DistributionMetric
showMetricToggle?: boolean
enableBreakdown?: boolean
showAccountCost?: boolean
startDate?: string
endDate?: string
filters?: Record<string, any>
@@ -134,6 +137,8 @@ const props = withDefaults(defineProps<{
loading: false,
metric: 'tokens',
showMetricToggle: false,
enableBreakdown: true,
showAccountCost: true,
})
const emit = defineEmits<{
@@ -143,6 +148,8 @@ const emit = defineEmits<{
const expandedKey = ref<string | null>(null)
const breakdownItems = ref<UserBreakdownItem[]>([])
const breakdownLoading = ref(false)
const showAccountCost = computed(() => props.showAccountCost)
const distributionColspan = computed(() => showAccountCost.value ? 6 : 5)
const toggleBreakdown = async (type: string, id: number | string) => {
const key = `${type}-${id}`
@@ -114,23 +114,25 @@
<th class="pb-2 text-right">{{ t('admin.dashboard.requests') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.tokens') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.actual') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.accountCost') }}</th>
<th v-if="showAccountCost" class="pb-2 text-right">{{ t('admin.dashboard.accountCost') }}</th>
<th class="pb-2 text-right">{{ t('admin.dashboard.standard') }}</th>
</tr>
</thead>
<tbody>
<template v-for="model in displayModelStats" :key="model.model">
<tr
class="border-t border-gray-100 cursor-pointer transition-colors hover:bg-gray-50 dark:border-gray-700 dark:hover:bg-dark-700/40"
@click="toggleBreakdown('model', model.model)"
class="border-t border-gray-100 transition-colors dark:border-gray-700"
:class="enableBreakdown ? 'cursor-pointer hover:bg-gray-50 dark:hover:bg-dark-700/40' : ''"
@click="enableBreakdown && toggleBreakdown('model', model.model)"
>
<td
class="max-w-[100px] truncate py-1.5 font-medium text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300"
class="max-w-[100px] truncate py-1.5 font-medium"
:class="enableBreakdown ? 'text-blue-600 hover:text-blue-800 dark:text-blue-400 dark:hover:text-blue-300' : 'text-gray-900 dark:text-white'"
:title="model.model"
>
<span class="inline-flex items-center gap-1">
<svg v-if="expandedKey === `model-${model.model}`" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
<svg v-if="enableBreakdown && expandedKey === `model-${model.model}`" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M19 9l-7 7-7-7"/></svg>
<svg v-else-if="enableBreakdown" class="h-3 w-3 shrink-0" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path stroke-linecap="round" stroke-linejoin="round" stroke-width="2" d="M9 5l7 7-7 7"/></svg>
{{ model.model }}
</span>
</td>
@@ -143,7 +145,7 @@
<td class="py-1.5 text-right text-green-600 dark:text-green-400">
${{ formatCost(model.actual_cost) }}
</td>
<td class="py-1.5 text-right text-orange-500 dark:text-orange-400">
<td v-if="showAccountCost" class="py-1.5 text-right text-orange-500 dark:text-orange-400">
${{ formatCost(model.account_cost) }}
</td>
<td class="py-1.5 text-right text-gray-400 dark:text-gray-500">
@@ -151,10 +153,11 @@
</td>
</tr>
<tr v-if="expandedKey === `model-${model.model}`">
<td colspan="6" class="p-0">
<td :colspan="distributionColspan" class="p-0">
<UserBreakdownSubTable
:items="breakdownItems"
:loading="breakdownLoading"
:show-account-cost="showAccountCost"
/>
</td>
</tr>
@@ -270,6 +273,8 @@ const props = withDefaults(defineProps<{
metric?: DistributionMetric
showSourceToggle?: boolean
showMetricToggle?: boolean
enableBreakdown?: boolean
showAccountCost?: boolean
rankingLoading?: boolean
rankingError?: boolean
startDate?: string
@@ -288,6 +293,8 @@ const props = withDefaults(defineProps<{
metric: 'tokens',
showSourceToggle: false,
showMetricToggle: false,
enableBreakdown: true,
showAccountCost: true,
rankingLoading: false,
rankingError: false
})
@@ -328,6 +335,8 @@ const emit = defineEmits<{
}>()
const enableRankingView = computed(() => props.enableRankingView)
const showAccountCost = computed(() => props.showAccountCost)
const distributionColspan = computed(() => showAccountCost.value ? 6 : 5)
const activeView = ref<'model_distribution' | 'spending_ranking'>('model_distribution')
const chartColors = [
@@ -25,7 +25,7 @@
<td class="py-1 text-right text-green-600 dark:text-green-400">
${{ formatCost(user.actual_cost) }}
</td>
<td class="py-1 text-right text-orange-500 dark:text-orange-400">
<td v-if="showAccountCost" class="py-1 text-right text-orange-500 dark:text-orange-400">
${{ formatCost(user.account_cost) }}
</td>
<td class="py-1 pr-1 text-right text-gray-400 dark:text-gray-500">
@@ -38,16 +38,23 @@
</template>
<script setup lang="ts">
import { computed } from 'vue'
import { useI18n } from 'vue-i18n'
import LoadingSpinner from '@/components/common/LoadingSpinner.vue'
import type { UserBreakdownItem } from '@/types'
const { t } = useI18n()
defineProps<{
const props = withDefaults(defineProps<{
items: UserBreakdownItem[]
loading?: boolean
}>()
showAccountCost?: boolean
}>(), {
loading: false,
showAccountCost: true,
})
const showAccountCost = computed(() => props.showAccountCost)
const formatTokens = (value: number): string => {
if (value >= 1_000_000_000) return `${(value / 1_000_000_000).toFixed(2)}B`
@@ -56,7 +63,8 @@ const formatTokens = (value: number): string => {
return value.toLocaleString()
}
const formatCost = (value: number): string => {
const formatCost = (value: number | undefined | null): string => {
if (value == null) return '0.0000'
if (value >= 1000) return (value / 1000).toFixed(2) + 'K'
if (value >= 1) return value.toFixed(2)
if (value >= 0.01) return value.toFixed(3)
@@ -10,6 +10,7 @@ const messages: Record<string, string> = {
'admin.dashboard.requests': 'Requests',
'admin.dashboard.tokens': 'Tokens',
'admin.dashboard.actual': 'Actual',
'admin.dashboard.accountCost': 'Account Cost',
'admin.dashboard.standard': 'Standard',
'admin.dashboard.metricTokens': 'By Tokens',
'admin.dashboard.metricActualCost': 'By Actual Cost',
@@ -111,4 +112,22 @@ describe('GroupDistributionChart', () => {
})
expect(label).toBe('group-b: $0.900 (90.0%)')
})
it('can hide account cost for user usage stats without account_cost', () => {
const wrapper = mount(GroupDistributionChart, {
props: {
groupStats,
showAccountCost: false,
},
global: {
stubs: {
LoadingSpinner: true,
},
},
})
expect(wrapper.text()).not.toContain('Account Cost')
expect(wrapper.findAll('thead th')).toHaveLength(5)
expect(wrapper.findAll('tbody tr')[0].findAll('td')).toHaveLength(5)
})
})
@@ -17,6 +17,7 @@ const messages: Record<string, string> = {
'admin.dashboard.requests': 'Requests',
'admin.dashboard.tokens': 'Tokens',
'admin.dashboard.actual': 'Actual',
'admin.dashboard.accountCost': 'Account Cost',
'admin.dashboard.standard': 'Standard',
'admin.dashboard.metricTokens': 'By Tokens',
'admin.dashboard.metricActualCost': 'By Actual Cost',
@@ -126,6 +127,24 @@ describe('ModelDistributionChart', () => {
expect(label).toBe('model-b: $1.40 (87.5%)')
})
it('can hide account cost for user usage stats without account_cost', () => {
const wrapper = mount(ModelDistributionChart, {
props: {
modelStats,
showAccountCost: false,
},
global: {
stubs: {
LoadingSpinner: true,
},
},
})
expect(wrapper.text()).not.toContain('Account Cost')
expect(wrapper.findAll('thead th')).toHaveLength(5)
expect(wrapper.findAll('tbody tr')[0].findAll('td')).toHaveLength(5)
})
it('renders Others in the spending ranking table and uses a dedicated chart color', async () => {
const wrapper = mount(ModelDistributionChart, {
props: {
+8 -5
View File
@@ -1290,6 +1290,7 @@ export interface UsageLog {
// User-Agent
user_agent: string | null
ip_address?: string | null
// Cache TTL Override
cache_ttl_overridden: boolean
@@ -1323,9 +1324,6 @@ export interface AdminUsageLog extends UsageLog {
channel_id?: number | null
billing_tier?: string | null
// 用户请求 IP(仅管理员可见)
ip_address?: string | null
// 最小账号信息(仅管理员接口返回)
account?: UsageLogAccountSummary
}
@@ -1468,6 +1466,9 @@ export interface UsageStatsResponse {
total_actual_cost: number // 实际扣除
average_duration_ms: number
models?: Record<string, number>
endpoints?: EndpointStat[]
upstream_endpoints?: EndpointStat[]
endpoint_paths?: EndpointStat[]
}
// ==================== Trend & Chart Types ====================
@@ -1494,7 +1495,7 @@ export interface ModelStat {
total_tokens: number
cost: number // 标准计费
actual_cost: number // 实际扣除
account_cost: number // 账号成本
account_cost?: number // 账号成本(仅管理员接口返回)
}
export interface EndpointStat {
@@ -1512,7 +1513,7 @@ export interface GroupStat {
total_tokens: number
cost: number // 标准计费
actual_cost: number // 实际扣除
account_cost: number // 账号成本
account_cost?: number // 账号成本(仅管理员接口返回)
}
export interface UserBreakdownItem {
@@ -1687,8 +1688,10 @@ export interface UsageQueryParams {
request_type?: UsageRequestType
stream?: boolean
billing_type?: number | null
billing_mode?: string | null
start_date?: string
end_date?: string
timezone?: string
sort_by?: string
sort_order?: 'asc' | 'desc'
}
File diff suppressed because it is too large Load Diff
@@ -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<string, string> = {
'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: '<div><slot /></div>' }
const TablePageLayoutStub = {
template: '<div><slot name="actions" /><slot name="filters" /><slot name="table" /><slot /></div>',
}
const DataTableStub = {
props: ['data'],
template: `
<div>
<div v-for="row in data" :key="row.request_id">
<slot name="cell-billing_mode" :row="row" />
<slot name="cell-tokens" :row="row" />
<slot name="cell-cost" :row="row" />
</div>
</div>
`,
const simpleStub = { template: '<div><slot /></div>' }
const chartStub = { template: '<div />' }
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<string, unknown> | 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<string>((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()
})
})