mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:37:52 +08:00
审计修复,逐条如下。 H1 利润策略泄露给所有普通用户 profit_control_enabled / profit_min_margin / profit_safety_buffer 从 dto.Group 移到 dto.AdminGroup(后者内嵌前者),赋值相应从 groupFromServiceBase 移到 GroupFromServiceAdmin;前端 TS 同步从 Group 移到 AdminGroup。dto.Group 是 GET /api/v1/groups/available 的响应体,该响应本就带 rate_multiplier,相乘即可反推运营方上游采购成本上限。 api_contract_test.go 的 /groups/available golden JSON 回滚这三个字段,并把 fixture 改成非零值(require.JSONEq 是精确比对,缺字段即失败)。 新增 dto 层边界测试:普通用户 DTO 不含三字段、管理员 DTO 仍含。 M1 利润终检 continue 与 failover 503 退避互动产生活锁 FailoverState 新增 profitVetoedAccountIDs / profitVetoCount 与 RecordProfitVeto():加入排除集 + 计数,达 maxProfitVetoAttempts(10) 返回 FailoverExhausted。HandleSelectionExhausted 的 503 清空分支改为清空后把利润 否决的账号放回排除集;若排除集已全部由利润否决贡献,清空不会带来任何新候选, 直接判定耗尽(否则 SwitchCount 永不前进、退避条件永远成立,每 2s 空转一轮)。 五个 handler 否决点(gateway_handler ×2 / responses / chat_completions / gemini_v1beta)改为经 RecordProfitVeto 决策,耗尽时按无可用账号终止。 回归测试钉死:503 之后持续利润否决必须有限步终止且不 spin;未启用利润控制的 请求退避语义完全不变。 M2 排队等槽后才终检,延迟可放大到 N × WaitPlan.Timeout OpenAI 侧选号循环(自有 failedAccountIDs map,非 FailoverState)新增 recordOpenAIProfitVeto + handleOpenAIProfitVetoExhausted,共用同一上限语义。 覆盖 responses / messages-dispatch / chat_completions / alpha_search / embeddings / images / grok_media 七处,以及 WS 两处否决分支。 M4 ws_v2 透传 ingress 绕过 per-turn 重定价(选方案 B:最小止血) 透传 relay 只回调 AfterTurn、没有任何 turn 起始回调,hooks.BeforeTurn 永远 不触发,而 handler 把 turnPricingAt 初始化成建连时刻 ⇒ 透传连接全部 turn 按 建连时刻的高峰因子结算,客户端峰前建连保活即可全程谷价——正是本 PR 想堵的 漏洞。改为 openAIWSTurnPricing 零值起步、只由 BeforeTurn 冻结;透传路径保持 零值,RecordUsage 回退记录时刻,与引入利润控制前的基线一致。 未选方案 A(给透传补 turn 起始回调):passthrough_relay.go 是 #5167 刚修过的 取消传播/close frame 时序敏感区;且 BeforeTurn 还承担 turn>1 的并发槽位抢占, 接进去等于给透传连接引入 per-turn 抢槽,风险远超本次修复范围。透传仍有建连时 的准入门,只是没有 turn 级复核,已在两处注释写明。 测试:service 层钉死透传 ingress 不触发 BeforeTurn(含失败时的复核指引), handler 层钉死零值语义与逐 turn 覆盖。 M5 装门读分组走了带账号计数聚合的 GetByID SchedulerSnapshotService 新增 GetGroupByIDLite,openai/gateway 两处装门改用 之。门只需要平台/倍率/利润/高峰字段,且该查询发生在「是否启用利润控制」判定 之前,未启用的分组同样付代价。两个测试 stub 的 GetByID 改成 panic 守卫。 M6 认证快照注释与真实读取路径相反 门解析优先取 ctxkey.Group,而它就是本快照物化出来的对象,直连流量走的正是这 条路。改正注释,与 api_key_repo.go 投影处的说明对齐,避免后人照旧注释删列。 M3 rate_multiplier 为 nil 时利润门 fail-closed(不改行为,加护栏) 保留 fail-closed。补 repository 层测试钉死账号调度快照的 full/metadata 两份 payload 都必须保留 RateMultiplier(含 0 值),漏列在 CI 就红。 L1 迁移号注释 191 / 191-192 改为实际的 192/193。 L2 admin group Create 的利润配置预校验改用与 CreateGroup 一致的归一化平台 (新增 service.NormalizeGroupPlatform,两边共用)。保留预校验而非删除: service 层返回的是无类型 error,经 ErrorFrom 会变成 500,删掉会把合法的 400 降级成 500。 L3 前端利润校验的上界改为判定换算后的小数(后端按小数校验 [0,1)), 99.999% 会四舍五入进位成 1.0 而被后端 400;i18n en/zh 同步改为 0-99.99。 L4 clampProfitControlThreshold / profitControlOverThreshold 抽为共用函数, 线上装门/否决点与 profit-preview 不再各自实现,附边界语义测试。 L5 profit-preview 补「默认 D 有账号但最低有效 D 归零」的告警(两档都为 0 由 既有告警覆盖,不重复)。
295 lines
10 KiB
Go
295 lines
10 KiB
Go
package handler
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
|
|
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
|
|
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
|
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
|
"github.com/Wei-Shaw/sub2api/internal/service"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/tidwall/gjson"
|
|
"go.uber.org/zap"
|
|
)
|
|
|
|
// Embeddings handles the OpenAI-compatible Embeddings API.
|
|
// POST /v1/embeddings
|
|
func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
|
streamStarted := false
|
|
requestStart := time.Now()
|
|
|
|
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
|
if !ok {
|
|
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
|
|
return
|
|
}
|
|
|
|
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
|
if !ok {
|
|
h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
|
|
return
|
|
}
|
|
reqLog := requestLogger(
|
|
c,
|
|
"handler.openai_gateway.embeddings",
|
|
zap.Int64("user_id", subject.UserID),
|
|
zap.Int64("api_key_id", apiKey.ID),
|
|
zap.Any("group_id", apiKey.GroupID),
|
|
)
|
|
if !h.ensureResponsesDependencies(c, reqLog) {
|
|
return
|
|
}
|
|
|
|
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
|
|
if err != nil {
|
|
if maxErr, ok := extractMaxBytesError(err); ok {
|
|
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
|
|
return
|
|
}
|
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
|
|
return
|
|
}
|
|
if len(body) == 0 {
|
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
|
|
return
|
|
}
|
|
if !gjson.ValidBytes(body) {
|
|
logRequestBodyParseFailure(reqLog, body, nil)
|
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
|
return
|
|
}
|
|
|
|
modelResult := gjson.GetBytes(body, "model")
|
|
if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" {
|
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
|
|
return
|
|
}
|
|
reqModel := modelResult.String()
|
|
ensureCompositeTargetPlatform(c, apiKey, reqModel)
|
|
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) {
|
|
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
|
|
return
|
|
}
|
|
reqLog = reqLog.With(zap.String("model", reqModel))
|
|
setOpsRequestContext(c, reqModel, false)
|
|
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
|
|
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, "openai_embeddings", reqModel, body); decision != nil && !decision.AllowNextStage {
|
|
h.openAISecurityAuditError(c, decision)
|
|
return
|
|
}
|
|
|
|
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
|
|
|
|
subscription, _ := middleware2.GetSubscriptionFromContext(c)
|
|
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
|
|
|
|
userReleaseFunc, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog)
|
|
if !acquired {
|
|
return
|
|
}
|
|
if userReleaseFunc != nil {
|
|
defer userReleaseFunc()
|
|
}
|
|
|
|
if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
|
|
reqLog.Info("openai_embeddings.billing_check_failed", zap.Error(err))
|
|
status, code, message, retryAfter := billingErrorDetails(err)
|
|
if retryAfter > 0 {
|
|
c.Header("Retry-After", strconv.Itoa(retryAfter))
|
|
}
|
|
h.errorResponse(c, status, code, message)
|
|
return
|
|
}
|
|
|
|
profitVetoCount := 0
|
|
failedAccountIDs := make(map[int64]struct{})
|
|
var lastFailoverErr *service.UpstreamFailoverError
|
|
switchCount := 0
|
|
maxAccountSwitches := h.maxAccountSwitches
|
|
if maxAccountSwitches <= 0 {
|
|
maxAccountSwitches = 3
|
|
}
|
|
routingStart := time.Now()
|
|
|
|
// 分组利润控制:embeddings 文本入口请求级装门并固定 pricingAt。
|
|
embPricingCtx, pricingAt := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID)
|
|
c.Request = c.Request.WithContext(embPricingCtx)
|
|
|
|
for {
|
|
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
|
|
c.Request.Context(),
|
|
apiKey.GroupID,
|
|
"",
|
|
"",
|
|
reqModel,
|
|
failedAccountIDs,
|
|
service.OpenAIUpstreamTransportHTTPSSE,
|
|
service.OpenAIEndpointCapabilityEmbeddings,
|
|
false,
|
|
false,
|
|
true,
|
|
)
|
|
if err != nil {
|
|
if failoverClientGone(c) {
|
|
reqLog.Info("openai_embeddings.account_select_aborted_client_disconnected", zap.Error(err))
|
|
return
|
|
}
|
|
reqLog.Warn("openai_embeddings.account_select_failed",
|
|
zap.Error(err),
|
|
zap.Int("excluded_account_count", len(failedAccountIDs)),
|
|
)
|
|
if len(failedAccountIDs) == 0 {
|
|
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
|
|
if !cls.ModelNotFound {
|
|
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
|
}
|
|
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
|
return
|
|
}
|
|
if lastFailoverErr != nil {
|
|
h.handleFailoverExhausted(c, lastFailoverErr, false)
|
|
} else {
|
|
h.errorResponse(c, http.StatusBadGateway, "api_error", "Upstream request failed")
|
|
}
|
|
return
|
|
}
|
|
if selection == nil || selection.Account == nil {
|
|
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
|
|
if !cls.ModelNotFound {
|
|
markOpsRoutingCapacityLimited(c)
|
|
}
|
|
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
|
return
|
|
}
|
|
account := selection.Account
|
|
setOpsSelectedAccount(c, account.ID, account.Platform)
|
|
|
|
accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog)
|
|
if slotResult == openAISlotAcquireProfitVetoed {
|
|
// 利润终检否决:排除该账号重新选号;否决次数达上限则按无可用账号终止。
|
|
if !recordOpenAIProfitVeto(failedAccountIDs, account.ID, &profitVetoCount) {
|
|
h.handleOpenAIProfitVetoExhausted(c, streamStarted, reqLog, profitVetoCount)
|
|
return
|
|
}
|
|
continue
|
|
}
|
|
if slotResult != openAISlotAcquireOK {
|
|
return
|
|
}
|
|
|
|
service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
|
|
forwardStart := time.Now()
|
|
|
|
forwardBody := body
|
|
if channelMapping.Mapped {
|
|
forwardBody = h.gatewayService.ReplaceModelInBody(body, channelMapping.MappedModel)
|
|
}
|
|
writerSizeBeforeForward := c.Writer.Size()
|
|
result, err := func() (*service.OpenAIForwardResult, error) {
|
|
defer func() {
|
|
if accountReleaseFunc != nil {
|
|
accountReleaseFunc()
|
|
}
|
|
}()
|
|
return h.gatewayService.ForwardEmbeddings(c.Request.Context(), c, account, forwardBody, "")
|
|
}()
|
|
|
|
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
|
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
|
responseLatencyMs := forwardDurationMs
|
|
if upstreamLatencyMs > 0 && forwardDurationMs > upstreamLatencyMs {
|
|
responseLatencyMs = forwardDurationMs - upstreamLatencyMs
|
|
}
|
|
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, responseLatencyMs)
|
|
|
|
if err != nil {
|
|
var failoverErr *service.UpstreamFailoverError
|
|
if errors.As(err, &failoverErr) {
|
|
if c.Writer.Size() != writerSizeBeforeForward {
|
|
h.handleFailoverExhausted(c, failoverErr, true)
|
|
return
|
|
}
|
|
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
|
if failoverClientGone(c) {
|
|
reqLog.Info("openai_embeddings.failover_aborted_client_disconnected",
|
|
zap.Int64("account_id", account.ID),
|
|
zap.Int("upstream_status", failoverErr.StatusCode),
|
|
)
|
|
return
|
|
}
|
|
h.gatewayService.RecordOpenAIAccountSwitch()
|
|
failedAccountIDs[account.ID] = struct{}{}
|
|
lastFailoverErr = failoverErr
|
|
if switchCount >= maxAccountSwitches {
|
|
h.handleFailoverExhausted(c, failoverErr, false)
|
|
return
|
|
}
|
|
switchCount++
|
|
reqLog.Warn("openai_embeddings.upstream_failover_switching",
|
|
zap.Int64("account_id", account.ID),
|
|
zap.Int("upstream_status", failoverErr.StatusCode),
|
|
zap.Int("switch_count", switchCount),
|
|
zap.Int("max_switches", maxAccountSwitches),
|
|
)
|
|
continue
|
|
}
|
|
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), false, nil)
|
|
if c.Writer.Size() == writerSizeBeforeForward {
|
|
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
|
|
}
|
|
reqLog.Warn("openai_embeddings.forward_failed",
|
|
zap.Int64("account_id", account.ID),
|
|
zap.Error(err),
|
|
)
|
|
return
|
|
}
|
|
|
|
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, account.GetMappedModel(reqModel), true, nil)
|
|
userAgent := c.GetHeader("User-Agent")
|
|
clientIP := ip.GetClientIP(c)
|
|
inboundEndpoint := GetInboundEndpoint(c)
|
|
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
|
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
|
|
sessionID := service.ExtractClientSessionID(c)
|
|
|
|
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
|
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
|
Result: result,
|
|
APIKey: apiKey,
|
|
User: apiKey.User,
|
|
Account: account,
|
|
Subscription: subscription,
|
|
InboundEndpoint: inboundEndpoint,
|
|
UpstreamEndpoint: upstreamEndpoint,
|
|
UserAgent: userAgent,
|
|
IPAddress: clientIP,
|
|
APIKeyService: h.apiKeyService,
|
|
QuotaPlatform: quotaPlatform,
|
|
SessionID: sessionID,
|
|
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
|
|
PricingAt: pricingAt,
|
|
}); err != nil {
|
|
logger.L().With(
|
|
zap.String("component", "handler.openai_gateway.embeddings"),
|
|
zap.Int64("user_id", subject.UserID),
|
|
zap.Int64("api_key_id", apiKey.ID),
|
|
zap.Any("group_id", apiKey.GroupID),
|
|
zap.String("model", reqModel),
|
|
zap.Int64("account_id", account.ID),
|
|
).Error("openai_embeddings.record_usage_failed", zap.Error(err))
|
|
}
|
|
})
|
|
reqLog.Debug("openai_embeddings.request_completed",
|
|
zap.Int64("account_id", account.ID),
|
|
zap.Int("switch_count", switchCount),
|
|
)
|
|
return
|
|
}
|
|
}
|