118 lines
3.6 KiB
Go
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
|
|
}
|