Files
NianAIGC/backend/internal/httpapi/usage.go
2026-08-21 13:44:56 +08:00

118 lines
3.6 KiB
Go

package httpapi
import (
"errors"
"net/http"
"strings"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/usage"
)
type UsageReporter = usage.Reporter
type usageHandler struct {
authorizer *PlatformAuthorizer
reporter UsageReporter
now func() time.Time
}
func NewUsageHandler(authorizer *PlatformAuthorizer, reporter UsageReporter, now func() time.Time) http.Handler {
if now == nil {
now = time.Now
}
return &usageHandler{authorizer: authorizer, reporter: reporter, now: now}
}
func (h *usageHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/api/usage" && r.URL.Path != "/api/admin/usage" {
http.NotFound(w, r)
return
}
if !allow(w, r, http.MethodGet) {
return
}
if h.authorizer == nil || h.reporter == nil {
writeAPIError(w, 500, "服务器内部错误。")
return
}
if r.URL.Path == "/api/usage" {
h.personal(w, r)
return
}
h.admin(w, r)
}
func (h *usageHandler) personal(w http.ResponseWriter, r *http.Request) {
session, err := h.authorizer.Authorize(r, PlatformApp)
if err != nil {
writeAuthError(w, err)
return
}
preset := usage.Preset(r.URL.Query().Get("preset"))
switch preset {
case usage.PresetToday, usage.Preset7Days, usage.Preset30Days, usage.PresetMonth:
default:
preset = usage.PresetMonth
}
report, err := h.reporter.Personal(r.Context(), usage.PersonalRequest{AccountID: session.User.ID, Preset: preset, Now: h.now()})
if err != nil {
writeAPIError(w, 500, "服务器内部错误。")
return
}
writeJSON(w, 200, report)
}
func (h *usageHandler) admin(w http.ResponseWriter, r *http.Request) {
session, err := h.authorizer.Authorize(r, PlatformAdmin)
if err != nil {
writeAuthError(w, err)
return
}
query := r.URL.Query()
capability, provider := strings.TrimSpace(query.Get("capability")), strings.TrimSpace(query.Get("provider"))
if capability != "" && capability != "image.generate" && capability != "video.generate" {
writeAPIError(w, 400, "不支持的功能类型。")
return
}
if provider != "" && provider != "volcengine-visual" && provider != "evolink" && provider != "seedance" && provider != "seedream" && provider != "bailian" {
writeAPIError(w, 400, "不支持的服务商。")
return
}
startDate, endDate := strings.TrimSpace(query.Get("startDate")), strings.TrimSpace(query.Get("endDate"))
if !validDate(startDate) || !validDate(endDate) {
writeAPIError(w, 400, "日期格式无效。")
return
}
request := usage.AdminRequest{Requester: usage.Requester{AccountID: session.User.ID, OrganizationID: session.User.OrganizationID, Role: session.User.Role}, StartDate: startDate, EndDate: endDate, Capability: capability, Provider: provider, Now: h.now()}
if session.User.Role == "super_admin" {
request.OrganizationID = strings.TrimSpace(query.Get("organizationId"))
request.OwnerID = strings.TrimSpace(query.Get("ownerId"))
} else {
request.OrganizationID = session.User.OrganizationID
request.RedactAccounts = true
}
report, err := h.reporter.Admin(r.Context(), request)
if err != nil {
if errors.Is(err, usage.ErrDateRangeTooLong) {
writeAPIError(w, 400, "单次查询最多支持 10 年。")
return
}
if errors.Is(err, usage.ErrInvalidDateRange) {
writeAPIError(w, 400, "开始日期不能晚于结束日期。")
return
}
writeAPIError(w, 500, "服务器内部错误。")
return
}
if request.RedactAccounts {
report.Accounts = []any{}
report.Recent = []any{}
report.Options.Accounts = []usage.Option{}
}
writeJSON(w, 200, report)
}
func validDate(value string) bool {
if value == "" {
return true
}
_, err := time.Parse("2006-01-02", value)
return err == nil
}