Files
sub2api/backend/internal/handler/model_plaza_handler.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

205 lines
7.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package handler
import (
"log/slog"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
// ModelPlazaHandler 处理「模型广场」查询。
//
// 广场路由挂 OptionalJWT 中间件:匿名可访问(除非 require_auth 开启),带 token 则
// 识别用户。可见性规则(橱窗语义,与「可用渠道」的可绑定语义不同):
// - 匿名:仅非专属分组(订阅型照常展示);
// - 登录:非专属分组 + user_allowed_groups 授权的专属分组(不检查订阅有效性)。
type ModelPlazaHandler struct {
channelService *service.ChannelService
apiKeyService *service.APIKeyService
settingService *service.SettingService
}
// NewModelPlazaHandler 创建模型广场 handler。
func NewModelPlazaHandler(
channelService *service.ChannelService,
apiKeyService *service.APIKeyService,
settingService *service.SettingService,
) *ModelPlazaHandler {
return &ModelPlazaHandler{
channelService: channelService,
apiKeyService: apiKeyService,
settingService: settingService,
}
}
// modelPlazaOfficialPricing LiteLLM 官方参考价(USD per token)。
type modelPlazaOfficialPricing struct {
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheWrite1hPrice *float64 `json:"cache_write_1h_price,omitempty"`
CacheReadPrice *float64 `json:"cache_read_price"`
}
// modelPlazaModel 广场模型条目:渠道定价(白名单形态)+ 官方参考价。
type modelPlazaModel struct {
Name string `json:"name"`
Platform string `json:"platform"`
Pricing *userSupportedModelPricing `json:"pricing"`
OfficialPricing *modelPlazaOfficialPricing `json:"official_pricing"`
}
// modelPlazaGroup 广场分组条目(白名单字段)。
type modelPlazaGroup struct {
ID int64 `json:"id"`
Name string `json:"name"`
Description string `json:"description"`
Platform string `json:"platform"`
SubscriptionType string `json:"subscription_type"`
RateMultiplier float64 `json:"rate_multiplier"`
UserRateMultiplier *float64 `json:"user_rate_multiplier,omitempty"`
PeakRateEnabled bool `json:"peak_rate_enabled"`
PeakStart string `json:"peak_start"`
PeakEnd string `json:"peak_end"`
PeakRateMultiplier float64 `json:"peak_rate_multiplier"`
IsExclusive bool `json:"is_exclusive"`
// 生图独立倍率:为 true 时图片计费模型的实付倍率取 ImageRateMultiplier
// 不取分组/用户专属倍率。
ImageRateIndependent bool `json:"image_rate_independent"`
ImageRateMultiplier float64 `json:"image_rate_multiplier"`
Models []modelPlazaModel `json:"models"`
}
// modelPlazaResponse 广场页响应。
type modelPlazaResponse struct {
Description string `json:"description"`
Groups []modelPlazaGroup `json:"groups"`
}
// Get 返回模型广场数据。
// GET /api/v1/model-plaza
func (h *ModelPlazaHandler) Get(c *gin.Context) {
if h.settingService == nil {
response.NotFound(c, "Model plaza is not enabled")
return
}
rt := h.settingService.GetModelPlazaRuntime(c.Request.Context())
if !rt.Enabled {
response.NotFound(c, "Model plaza is not enabled")
return
}
subject, authed := middleware.GetAuthSubjectFromContext(c)
if rt.RequireAuth && !authed {
response.Unauthorized(c, "Authentication required")
return
}
groups, err := h.channelService.ListPlazaGroups(c.Request.Context())
if err != nil {
response.ErrorFrom(c, err)
return
}
// allowedExclusive == nil 表示匿名;登录用户恒为非 nil(可能为空集合)。
var allowedExclusive map[int64]struct{}
var userRates map[int64]float64
if authed {
allowedExclusive, err = h.apiKeyService.GetUserAllowedGroupIDSet(c.Request.Context(), subject.UserID)
if err != nil {
// 可见性数据拿不到时不能静默降级成匿名视图(会错漏专属分组),直接报错。
response.ErrorFrom(c, err)
return
}
userRates, err = h.apiKeyService.GetUserGroupRates(c.Request.Context(), subject.UserID)
if err != nil {
// 专属倍率仅是展示增强,失败降级为分组默认倍率。
slog.Warn("model_plaza_user_rates_failed", "error", err, "user_id", subject.UserID)
userRates = nil
}
}
visible := filterPlazaVisibleGroups(groups, allowedExclusive)
out := make([]modelPlazaGroup, 0, len(visible))
for i := range visible {
out = append(out, toModelPlazaGroupDTO(&visible[i], userRates))
}
response.Success(c, modelPlazaResponse{
Description: rt.Description,
Groups: out,
})
}
// filterPlazaVisibleGroups 按登录态裁剪分组可见性。
// allowedExclusive == nil 表示匿名(仅非专属);非 nil 表示登录(非专属 + 授权专属)。
func filterPlazaVisibleGroups(
groups []service.PlazaGroup,
allowedExclusive map[int64]struct{},
) []service.PlazaGroup {
visible := make([]service.PlazaGroup, 0, len(groups))
for _, g := range groups {
if g.IsExclusive {
if allowedExclusive == nil {
continue
}
if _, ok := allowedExclusive[g.ID]; !ok {
continue
}
}
visible = append(visible, g)
}
return visible
}
// toModelPlazaGroupDTO 将 service 层广场分组映射为白名单 DTO,并合并用户专属倍率。
func toModelPlazaGroupDTO(g *service.PlazaGroup, userRates map[int64]float64) modelPlazaGroup {
models := make([]modelPlazaModel, 0, len(g.Models))
for i := range g.Models {
m := &g.Models[i]
models = append(models, modelPlazaModel{
Name: m.Name,
Platform: m.Platform,
Pricing: toUserPricing(m.Pricing),
OfficialPricing: toModelPlazaOfficialPricing(m.OfficialPricing),
})
}
dto := modelPlazaGroup{
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
Models: models,
}
if rate, ok := userRates[g.ID]; ok {
dto.UserRateMultiplier = &rate
}
return dto
}
// toModelPlazaOfficialPricing 转换官方参考价;nil 透传(前端显示 "-")。
func toModelPlazaOfficialPricing(p *service.PlazaOfficialPricing) *modelPlazaOfficialPricing {
if p == nil {
return nil
}
return &modelPlazaOfficialPricing{
InputPrice: p.InputPrice,
OutputPrice: p.OutputPrice,
CacheWritePrice: p.CacheWritePrice,
CacheWrite1hPrice: p.CacheWrite1hPrice,
CacheReadPrice: p.CacheReadPrice,
}
}