Merge branch 'Wei-Shaw:main' into main

This commit is contained in:
xueshiji
2026-07-01 14:00:51 +08:00
committed by GitHub
71 changed files with 3620 additions and 1931 deletions
@@ -259,6 +259,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: settings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: settings.OpenAICodexUserAgent,
MinCodexVersion: settings.MinCodexVersion,
@@ -598,6 +599,7 @@ type UpdateSettingsRequest struct {
ClaudeOAuthSystemPromptBlocks *string `json:"claude_oauth_system_prompt_blocks"`
EnableAnthropicCacheTTL1hInjection *bool `json:"enable_anthropic_cache_ttl_1h_injection"`
RewriteMessageCacheControl *bool `json:"rewrite_message_cache_control"`
EnableClientDatelineNormalization *bool `json:"enable_client_dateline_normalization"`
AntigravityUserAgentVersion *string `json:"antigravity_user_agent_version"`
OpenAICodexUserAgent *string `json:"openai_codex_user_agent"`
@@ -1731,6 +1733,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
return previousSettings.RewriteMessageCacheControl
}(),
EnableClientDatelineNormalization: func() bool {
if req.EnableClientDatelineNormalization != nil {
return *req.EnableClientDatelineNormalization
}
return previousSettings.EnableClientDatelineNormalization
}(),
AntigravityUserAgentVersion: func() string {
if req.AntigravityUserAgentVersion != nil {
return *req.AntigravityUserAgentVersion
@@ -2143,6 +2151,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent,
MinCodexVersion: updatedSettings.MinCodexVersion,
@@ -2644,6 +2653,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
if before.RewriteMessageCacheControl != after.RewriteMessageCacheControl {
changed = append(changed, "rewrite_message_cache_control")
}
if before.EnableClientDatelineNormalization != after.EnableClientDatelineNormalization {
changed = append(changed, "enable_client_dateline_normalization")
}
if before.AntigravityUserAgentVersion != after.AntigravityUserAgentVersion {
changed = append(changed, "antigravity_user_agent_version")
}
+6 -4
View File
@@ -577,7 +577,7 @@ func AccountSummaryFromService(a *service.Account) *AccountSummary {
}
func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
// 普通用户 DTO:严禁包含管理员字段(例如 account_rate_multiplier、ip_address、account)。
// 普通用户 DTO:严禁包含管理员字段(例如 account_rate_multiplier、account、upstream_model)。
requestType := l.EffectiveRequestType()
stream, openAIWSMode := service.ApplyLegacyRequestFields(requestType, l.Stream, l.OpenAIWSMode)
requestedModel := l.RequestedModel
@@ -594,7 +594,6 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
ServiceTier: l.ServiceTier,
ReasoningEffort: l.ReasoningEffort,
InboundEndpoint: l.InboundEndpoint,
UpstreamEndpoint: l.UpstreamEndpoint,
GroupID: l.GroupID,
SubscriptionID: l.SubscriptionID,
InputTokens: l.InputTokens,
@@ -626,6 +625,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
ImageSizeBreakdown: l.ImageSizeBreakdown,
MediaType: l.MediaType,
UserAgent: l.UserAgent,
IPAddress: l.IPAddress,
CacheTTLOverridden: l.CacheTTLOverridden,
BillingMode: l.BillingMode,
CreatedAt: l.CreatedAt,
@@ -637,7 +637,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
}
// UsageLogFromService converts a service UsageLog to DTO for regular users.
// It excludes Account details and IP address - users should not see these.
// It excludes admin-only account/upstream internals while keeping user billing and request metadata.
func UsageLogFromService(l *service.UsageLog) *UsageLog {
if l == nil {
return nil
@@ -652,8 +652,10 @@ func UsageLogFromServiceAdmin(l *service.UsageLog) *AdminUsageLog {
if l == nil {
return nil
}
usageLog := usageLogFromServiceUser(l)
usageLog.UpstreamEndpoint = l.UpstreamEndpoint
return &AdminUsageLog{
UsageLog: usageLogFromServiceUser(l),
UsageLog: usageLog,
UpstreamModel: l.UpstreamModel,
ChannelID: l.ChannelID,
ModelMappingChain: l.ModelMappingChain,
@@ -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()
+1
View File
@@ -186,6 +186,7 @@ type SystemSettings struct {
ClaudeOAuthSystemPromptBlocks string `json:"claude_oauth_system_prompt_blocks"`
EnableAnthropicCacheTTL1hInjection bool `json:"enable_anthropic_cache_ttl_1h_injection"`
RewriteMessageCacheControl bool `json:"rewrite_message_cache_control"`
EnableClientDatelineNormalization bool `json:"enable_client_dateline_normalization"`
AntigravityUserAgentVersion string `json:"antigravity_user_agent_version"`
OpenAICodexUserAgent string `json:"openai_codex_user_agent"`
+3 -1
View File
@@ -485,6 +485,8 @@ type UsageLog struct {
// User-Agent
UserAgent *string `json:"user_agent"`
// IPAddress is visible to the owner of the usage record.
IPAddress *string `json:"ip_address,omitempty"`
// Cache TTL Override 标记
CacheTTLOverridden bool `json:"cache_ttl_overridden"`
@@ -520,7 +522,7 @@ type AdminUsageLog struct {
// AccountStatsCost 自定义定价规则计算的账号统计费用(nil 表示使用默认公式)
AccountStatsCost *float64 `json:"account_stats_cost,omitempty"`
// IPAddress 用户请求 IP(仅管理员可见)
// IPAddress 用户请求 IP
IPAddress *string `json:"ip_address,omitempty"`
// Account 最小账号信息(避免泄露敏感字段)
+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)
}
}