mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge branch 'Wei-Shaw:main' into main
This commit is contained in:
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"`
|
||||
|
||||
|
||||
@@ -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 最小账号信息(避免泄露敏感字段)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user