mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
被问"还有 auto 调度相关 stage 呢?"问出来的真实缺口。
## 问题
chainDrive 只返回 (resp, src, model, err),调用方只知道**最终哪个槽位赢了**。
遍历过程中算出来又丢掉的东西——哪些档被跳过、为什么跳过、哪些槽位硬失败、
哪档全忙——一律不可见。ChainErr 里其实有这些,但**只在全部失败时**才填,
而它是 error 返回值不是记录。于是:
"tier 1 冷却所以降级到 tier 3" == "tier 1 正常接单"
对插件而言 tier 只是个常量 -2("resolved by the chain"),信息量为零。而这
恰恰是优先级链存在的全部理由,也是"我那个贵模型为什么没被用"的答案。
## 做法(scheduler 侧零新依赖)
新增 TraceEvent / TraceSink,chainDrive 多一个可选 sink 参数:
- TraceEvent 是本包的普通 struct,sink 是 func 参数 ⇒ **不新增 import**,
scheduler 仍然可独立测试
- sink 为 nil 时每次 emit 只多一次 nil 判断;没有插件的网关在 AUTO 热路径上
零开销(gateway 的 chainTraceSink 直接返回 nil)
- 事件是纯观测:scheduler 不基于它做任何分支,gateway 也不把它喂回路由/
冷却/配额
四种 kind:tier_skip / slot_fail / tier_busy / selected,selected 每次成功
遍历恰好一次且是最后一步。顺序保证所有 step 在 routed 之前。
## 暴露给插件
新增 chain_step stage(逐个步骤),并在 request_end 载荷里加三个便于做报表的
字段:chain_walk(上限 12 步,防审计记录膨胀)、degraded、tier_served。
## ★ 计费口径(我按推荐的做,已写进文档,需要你确认)
**按实际服务的模型计费**:降级到 tier 3 仍按 tier 3 的价算,轨迹只作观测。
理由与 §7.5 的边界一致——插件只报表不执法,两套口径混在一起会引出"降级该不该
多收钱"这种无法从代码判断的争议。若要改成"按本该用的档计价",需要在 models
价目里允许按 tier 定价,这我没做,因为那是个产品决策。
## 计费插件同步消费
by_tier_served / skip_reasons / degraded_reqs 三个新维度。skip_reasons 的等待
时长做了归一(`no free slot within <wait>`),否则 busy-wait 文案一变就多一行。
降级次数在 request_end 里计而不是在 chain_step 里计:一次降级的请求要走多步,
按步计会重复计数。
## 判据(346 个测试全绿,新增 15 个)
scheduler 6 个:正常路径只发一个 selected / 跳档+降级可见 / 硬失败与跳档
严格区分(不可混为一谈,否则抖动上游看起来像空闲上游)/
nil sink 安全 / 全失败时轨迹与 ChainErr 并存且不互相破坏 /
空链不发事件
gateway 1 个端到端:tier 1 全 500 → 插件收到 slot_fail(tier 1) +
selected(tier 2),request_end 的 tier_served=2 且 degraded=true
lua 2 个:降级计数与按实际模型计价 / 跳过原因归一聚合
lua 1 个:chain_step 是真 stage 且顺序正确
3 个变异都红:去掉 slot_fail(3 个判据红)/ 去掉 tier_skip(1 个)/
去掉 degraded 字段(1 个)。
1423 lines
49 KiB
Go
1423 lines
49 KiB
Go
package gateway
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"log"
|
||
"net/http"
|
||
"regexp"
|
||
"strconv"
|
||
"strings"
|
||
"sync/atomic"
|
||
"time"
|
||
|
||
"llmsproxy/internal/config"
|
||
"llmsproxy/internal/lua"
|
||
"llmsproxy/internal/provider"
|
||
"llmsproxy/internal/scheduler"
|
||
"llmsproxy/internal/types"
|
||
)
|
||
|
||
// chatRequest mirrors the OpenAI chat completions request the gateway accepts.
|
||
type chatRequest struct {
|
||
Model string `json:"model"`
|
||
Messages []types.ChatMessage `json:"messages"`
|
||
Temperature *float64 `json:"temperature,omitempty"`
|
||
MaxTokens int `json:"max_tokens,omitempty"`
|
||
Stream bool `json:"stream,omitempty"`
|
||
Tools []interface{} `json:"tools,omitempty"`
|
||
ToolChoice interface{} `json:"tool_choice,omitempty"`
|
||
DisableThinking bool `json:"disable_thinking"`
|
||
ExtraBody map[string]interface{} `json:"extra_body,omitempty"`
|
||
|
||
// PromptCacheKey 是 OpenAI 原生的缓存键;部分客户端用它携带会话标识。
|
||
PromptCacheKey string `json:"prompt_cache_key,omitempty"`
|
||
}
|
||
|
||
// ChatCompletion is the non-streaming OpenAI response object.
|
||
type ChatCompletion struct {
|
||
ID string `json:"id"`
|
||
Object string `json:"object"`
|
||
Created int64 `json:"created"`
|
||
Model string `json:"model"`
|
||
Choices []ChatChoice `json:"choices"`
|
||
Usage *types.TokenUsage `json:"usage,omitempty"`
|
||
// Cost is the upstream-reported charge for this request, passed through
|
||
// verbatim (OpenCode reports it as a decimal string, always "0" on the
|
||
// flat-rate Go subscription). Absent when the upstream reports nothing.
|
||
Cost string `json:"cost,omitempty"`
|
||
}
|
||
|
||
type ChatChoice struct {
|
||
Index int `json:"index"`
|
||
Message RespMessage `json:"message"`
|
||
FinishReason string `json:"finish_reason"`
|
||
}
|
||
|
||
type RespMessage struct {
|
||
Role string `json:"role,omitempty"`
|
||
Content string `json:"content"`
|
||
ReasoningContent string `json:"reasoning_content,omitempty"`
|
||
ToolCalls json.RawMessage `json:"tool_calls,omitempty"`
|
||
}
|
||
|
||
type ChatChunk struct {
|
||
ID string `json:"id"`
|
||
Object string `json:"object"`
|
||
Created int64 `json:"created"`
|
||
Model string `json:"model"`
|
||
Choices []ChunkChoice `json:"choices"`
|
||
// Usage is sent in the final chunk of a stream (empty choices) so
|
||
// OpenAI-compatible clients can read token usage.
|
||
Usage *types.TokenUsage `json:"usage,omitempty"`
|
||
// Cost is the upstream-reported charge for this request, passed through
|
||
// verbatim (OpenCode reports it as a decimal string). Emitted on the
|
||
// terminal chunk, mirroring OpenCode's own {"choices":[],"cost":"0"}.
|
||
Cost string `json:"cost,omitempty"`
|
||
}
|
||
|
||
type ChunkChoice struct {
|
||
Index int `json:"index"`
|
||
Delta RespMessage `json:"delta"`
|
||
FinishReason *string `json:"finish_reason"`
|
||
}
|
||
|
||
var seq int64
|
||
|
||
func newID() string {
|
||
n := atomic.AddInt64(&seq, 1)
|
||
return fmt.Sprintf("chatcmpl-%d", n)
|
||
}
|
||
|
||
func isAuto(m string) bool {
|
||
m = strings.TrimSpace(m)
|
||
return m == "" || strings.EqualFold(m, "AUTO")
|
||
}
|
||
|
||
// resolveCands picks the ordered candidate providers for a requested model.
|
||
// toolCalling requests are anchored: they resolve to exactly one provider
|
||
// (highest-priority available) so a tool-call round never switches models.
|
||
func (g *Gateway) resolveCands(ctx context.Context, req *chatRequest) ([]*provider.Provider, string) {
|
||
model := req.Model
|
||
if model == "" {
|
||
model = g.core.DefaultModel()
|
||
}
|
||
cands, effective := g.resolveByModel(model)
|
||
cands = chatOnly(cands)
|
||
allow := g.allowedModels(ctx)
|
||
if allow != nil {
|
||
cands = filterCandsByModels(cands, allow)
|
||
}
|
||
if !toolRequest(req) {
|
||
return cands, effective
|
||
}
|
||
// tool-call request: pin to one provider (no AUTO fallback across models)
|
||
if len(cands) == 0 {
|
||
return nil, effective
|
||
}
|
||
first := cands[0]
|
||
// use the prefix-stripped id (effective), never the raw "src:model" form:
|
||
// ModelFor does exact matching and would fall back to the source's best
|
||
// chat model for an unknown id
|
||
eff := first.ModelFor(effective)
|
||
if eff == "" {
|
||
eff = firstModel(first)
|
||
}
|
||
return []*provider.Provider{first}, eff
|
||
}
|
||
|
||
// filterCandsByModels keeps only providers exposing at least one model of the
|
||
// scope (used for user keys with a restricted model scope). Scope entries with
|
||
// a Source pinned to a specific upstream narrow the candidates to that source
|
||
// for the matching model.
|
||
//
|
||
// An "AUTO" scope entry only allows the AUTO routing mode; it does NOT grant
|
||
// access to specific models.
|
||
func filterCandsByModels(cands []*provider.Provider, allow []config.ModelScope) []*provider.Provider {
|
||
for _, m := range allow {
|
||
if m.Model == "" {
|
||
return cands
|
||
}
|
||
}
|
||
allowed := make(map[string]bool, len(allow))
|
||
byModelSrc := map[string]map[string]bool{}
|
||
for _, m := range allow {
|
||
allowed[m.Model] = true
|
||
if m.Source != "" {
|
||
if byModelSrc[m.Model] == nil {
|
||
byModelSrc[m.Model] = map[string]bool{}
|
||
}
|
||
byModelSrc[m.Model][m.Source] = true
|
||
}
|
||
}
|
||
out := make([]*provider.Provider, 0, len(cands))
|
||
for _, p := range cands {
|
||
for _, id := range p.Models() {
|
||
if !allowed[id] {
|
||
continue
|
||
}
|
||
if srcs := byModelSrc[id]; len(srcs) > 0 && !srcs[p.Name()] {
|
||
continue
|
||
}
|
||
out = append(out, p)
|
||
break
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
|
||
// intersectModels restricts a model list to the scope (preserving order).
|
||
// An "AUTO" scope entry only allows the AUTO routing mode; it does NOT grant
|
||
// every model.
|
||
func intersectModels(models []string, allow []config.ModelScope) []string {
|
||
allowed := make(map[string]bool, len(allow))
|
||
for _, m := range allow {
|
||
if m.Model == "" {
|
||
return models
|
||
}
|
||
if !strings.EqualFold(m.Model, "AUTO") {
|
||
allowed[m.Model] = true
|
||
}
|
||
}
|
||
out := make([]string, 0, len(models))
|
||
seen := map[string]bool{}
|
||
for _, m := range models {
|
||
if allowed[m] && !seen[m] {
|
||
seen[m] = true
|
||
out = append(out, m)
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
|
||
// checkModelScope validates the effective model against the key's model scope
|
||
// and token quota. Returns an error message when rejected.
|
||
//
|
||
// A scope entry with model "AUTO" only allows requests where the effective
|
||
// model is AUTO (the routing mode). It does NOT grant access to specific model
|
||
// ids — that requires an explicit scope entry for the model.
|
||
// checkModelScope keeps its string-returning signature for callers that only
|
||
// need to know whether the request may proceed.
|
||
func (g *Gateway) checkModelScope(ctx context.Context, model string) string {
|
||
if q := g.checkQuota(ctx, model); q != nil {
|
||
return q.msg
|
||
}
|
||
return ""
|
||
}
|
||
|
||
// quotaRejection is a quota verdict: the message plus whether the client
|
||
// should retry. A spent quota is a rate limit (429 + Retry-After), not a
|
||
// permission failure (403): a client that sees 403 gives up on the key, while
|
||
// one that sees 429 with a retry hint waits and resumes when the window rolls
|
||
// over.
|
||
type quotaRejection struct {
|
||
msg string
|
||
retry int64 // seconds until the window resets; 0 = unknown
|
||
}
|
||
|
||
func (q *quotaRejection) Error() string { return q.msg }
|
||
|
||
// checkQuota validates the effective model against the key's model scope and
|
||
// that entry's quota. Returns nil when the request may proceed.
|
||
//
|
||
// Quotas are per scope entry, never key-wide: a model whose budget is spent
|
||
// is refused on its own while the key's other models keep working. The verdict
|
||
// carries the remaining seconds of the reset window so a spent budget answers
|
||
// 429 + Retry-After (come back when it rolls over) instead of 403 (which reads
|
||
// as "this key may never use this model" and makes clients give up).
|
||
func (g *Gateway) checkQuota(ctx context.Context, model string) *quotaRejection {
|
||
allow := g.allowedModels(ctx)
|
||
if allow == nil {
|
||
return nil
|
||
}
|
||
for _, sc := range allow {
|
||
if sc.Model != model {
|
||
continue
|
||
}
|
||
win := AutoPeriodSeconds(sc.Period, sc.Hours)
|
||
k := keyID(reqKey(ctx))
|
||
if sc.TokenQuota > 0 {
|
||
if used := g.scopeTokens(ctx, sc); used >= sc.TokenQuota {
|
||
return "aRejection{
|
||
msg: fmt.Sprintf("token quota exceeded for %q (%d/%d)", model, used, sc.TokenQuota),
|
||
retry: AutoSecondsToReset(sc.Period, sc.Hours),
|
||
}
|
||
}
|
||
}
|
||
if sc.ReqQuota > 0 {
|
||
if used := g.stats.KeyWindowReqs(k, win); used >= sc.ReqQuota {
|
||
return "aRejection{
|
||
msg: fmt.Sprintf("request quota exceeded for %q (%d/%d)", model, used, sc.ReqQuota),
|
||
retry: AutoSecondsToReset(sc.Period, sc.Hours),
|
||
}
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
return "aRejection{msg: fmt.Sprintf("model %q is not allowed for this key", model)}
|
||
}
|
||
|
||
// quotaWindowSuffix describes a quota's reset window for an error message, so
|
||
// a rejected caller can tell a permanent block from one that clears in an hour.
|
||
func quotaWindowSuffix(period string, hours int64) string {
|
||
switch {
|
||
case period == "hour":
|
||
return ", resets hourly"
|
||
case period == "week":
|
||
return ", resets weekly"
|
||
case period == "month":
|
||
return ", resets monthly"
|
||
case period == "nhour" && hours > 1:
|
||
return fmt.Sprintf(", resets every %dh", hours)
|
||
}
|
||
return ""
|
||
} // scopeTokens returns the tokens this key consumed on the scope entry's model
|
||
// within its reset window, isolated per key. For an AUTO entry the cap covers
|
||
// everything the key routed through AUTO; for a model entry it covers that
|
||
// model only.
|
||
//
|
||
// It reads the per-key hourly buckets rather than the key-blind model
|
||
// buckets, so one key's usage can never exhaust another's quota.
|
||
func (g *Gateway) scopeTokens(ctx context.Context, sc config.ModelScope) int64 {
|
||
k := keyID(reqKey(ctx))
|
||
win := AutoPeriodSeconds(sc.Period, sc.Hours)
|
||
if sc.Model != "" && strings.EqualFold(sc.Model, "AUTO") {
|
||
return g.stats.KeyWindowTokens(k, win)
|
||
}
|
||
return g.stats.KeyWindowModelTokens(k, sc.Model, sc.Source, win)
|
||
}
|
||
|
||
// hasScopeModel reports whether a model (possibly with a "source-model" /
|
||
// "source:model" / "source/model" pinning prefix) is allowed by a key's model
|
||
// scope. The prefix is stripped strictly: only when the prefix names a real
|
||
// source that actually serves the bare model (via Registry.EffectiveModel), so
|
||
// model ids that themselves contain separators (e.g. "deepseek-v4-flash-free")
|
||
// are never corrupted (P10-2).
|
||
//
|
||
// A scope entry with model "AUTO" only matches the literal AUTO routing mode;
|
||
// it does NOT grant access to specific model ids.
|
||
func (g *Gateway) hasScopeModel(list []config.ModelScope, s string) bool {
|
||
if r := g.core.Registry(); r != nil {
|
||
s = r.EffectiveModel(s)
|
||
}
|
||
for _, x := range list {
|
||
if x.Model == s {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
// writeScopeReject answers a model-scope or quota rejection. A spent quota is
|
||
// 429 (rate_limit_exceeded) with Retry-After, so a client waits and resumes
|
||
// after the reset; a model the key may not use stays 403 (model_not_allowed),
|
||
// because retrying cannot help.
|
||
func (g *Gateway) writeReject(w http.ResponseWriter, q *quotaRejection) {
|
||
if q.retry > 0 {
|
||
w.Header().Set("Retry-After", strconv.FormatInt(q.retry, 10))
|
||
writeError(w, http.StatusTooManyRequests, "rate_limit_exceeded", q.msg)
|
||
return
|
||
}
|
||
// no window to wait for: the cap is either permanent or key-wide with no
|
||
// period. "rate_limit_exceeded" still says "come back after the operator
|
||
// raises the cap", which 403 would not.
|
||
if strings.Contains(q.msg, "quota exceeded") {
|
||
writeError(w, http.StatusTooManyRequests, "rate_limit_exceeded", q.msg)
|
||
return
|
||
}
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", q.msg)
|
||
}
|
||
|
||
func (g *Gateway) resolveByModel(model string) ([]*provider.Provider, string) {
|
||
if isAuto(model) {
|
||
return g.core.Registry().Resolve("AUTO"), ""
|
||
}
|
||
return g.core.Registry().Resolve(model), g.core.Registry().EffectiveModel(model)
|
||
}
|
||
|
||
// chatOnly keeps providers that expose at least one chat-capable model, so a
|
||
// chat/AUTO request never lands on an image-only source (or borrows its image
|
||
// model id). Explicit image-kind requests stay on the imageOnly path.
|
||
func chatOnly(cands []*provider.Provider) []*provider.Provider {
|
||
out := make([]*provider.Provider, 0, len(cands))
|
||
for _, p := range cands {
|
||
for _, id := range p.Models() {
|
||
if m := p.ModelByID(id); m == nil || m.Kind != "image" {
|
||
out = append(out, p)
|
||
break
|
||
}
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
|
||
// toolRequest reports whether the request participates in a tool-call round.
|
||
func toolRequest(req *chatRequest) bool {
|
||
if len(req.Tools) > 0 || req.ToolChoice != nil {
|
||
return true
|
||
}
|
||
for _, m := range req.Messages {
|
||
if m.Role == "tool" || len(m.ToolCalls) > 0 {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) {
|
||
if r.Method != http.MethodPost {
|
||
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "use POST")
|
||
return
|
||
}
|
||
var req chatRequest
|
||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
||
return
|
||
}
|
||
// 客户端自带的会话标识(若有)比推导出来的更准,见 clientSessionFromRequest。
|
||
clientSession := clientSessionFromRequest(r, req.PromptCacheKey)
|
||
if len(req.Messages) == 0 {
|
||
writeError(w, http.StatusBadRequest, "invalid_request", "messages is required")
|
||
return
|
||
}
|
||
model := req.Model
|
||
if model == "" {
|
||
model = g.core.DefaultModel()
|
||
}
|
||
// request_start fires for EVERY chat request, on both the AUTO and the
|
||
// direct path, and it fires BEFORE the quota / model-scope gates on
|
||
// purpose: a plugin that counts volume or audits traffic must also see the
|
||
// requests the gateway rejected, otherwise "requests accepted" would be all
|
||
// it could ever report. It sits after authentication (so the key and role in
|
||
// the payload are real) and after the messages check (a body with no
|
||
// messages is not a chat request at all).
|
||
//
|
||
// Calling it here rather than inside each branch is what keeps the two paths
|
||
// honest: an earlier version called it only from the AUTO branch, so every
|
||
// direct (model-pinned) request silently skipped it. That was caught by
|
||
// TestHooksFireOnRealDirectChat, not by reading the code.
|
||
g.fireStart(r.Context(), &req, "chat", model, len(req.Messages), len(req.Tools))
|
||
if isAuto(model) {
|
||
chain := g.core.AutoChain()
|
||
if chain == nil || len(chain.Tiers) == 0 {
|
||
writeError(w, http.StatusServiceUnavailable, "no_provider", "no auto slot configured")
|
||
return
|
||
}
|
||
if q := g.checkQuota(r.Context(), "AUTO"); q != nil {
|
||
g.writeReject(w, q)
|
||
return
|
||
}
|
||
ctx := r.Context()
|
||
done := g.stats.Begin()
|
||
defer done()
|
||
inner := &types.ChatRequest{
|
||
Messages: req.Messages,
|
||
Temperature: req.Temperature,
|
||
MaxTokens: req.MaxTokens,
|
||
Stream: req.Stream,
|
||
Tools: req.Tools,
|
||
ToolChoice: req.ToolChoice,
|
||
DisableThinking: req.DisableThinking,
|
||
ExtraBody: req.ExtraBody,
|
||
ClientSession: clientSession,
|
||
}
|
||
rec := &Req{
|
||
Key: keyID(reqKey(ctx)),
|
||
Type: "chat",
|
||
OK: false,
|
||
}
|
||
// quotaExhausted reports a slot whose token window has been used up;
|
||
// exhausted slots are dropped from scheduling without penalty.
|
||
quotaExhausted := func(sl *scheduler.Slot) bool {
|
||
if sl.Quota <= 0 {
|
||
return false
|
||
}
|
||
win := AutoPeriodSeconds(sl.Period, sl.Hours)
|
||
return g.stats.WindowTokens(sl.Model, sl.Source, win) >= sl.Quota
|
||
}
|
||
if req.Stream {
|
||
rec.Type = "stream"
|
||
g.streamChatAuto(w, ctx, chain, inner, rec, quotaExhausted)
|
||
return
|
||
}
|
||
g.singleChatAuto(w, ctx, chain, inner, rec, quotaExhausted)
|
||
return
|
||
}
|
||
if !isAuto(model) {
|
||
if allow := g.allowedModels(r.Context()); allow != nil && !g.hasScopeModel(allow, model) {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", fmt.Sprintf("model %q is not allowed for this key", model))
|
||
return
|
||
}
|
||
}
|
||
cands, effective := g.resolveCands(r.Context(), &req)
|
||
if len(cands) == 0 {
|
||
writeError(w, http.StatusNotFound, "model_not_found", fmt.Sprintf("model %q is not configured", model))
|
||
return
|
||
}
|
||
if effective == "" {
|
||
effective = firstModel(cands[0])
|
||
}
|
||
if q := g.checkQuota(r.Context(), effective); q != nil {
|
||
g.writeReject(w, q)
|
||
return
|
||
}
|
||
ctx := r.Context()
|
||
done := g.stats.Begin()
|
||
defer done()
|
||
|
||
inner := &types.ChatRequest{
|
||
Model: effective,
|
||
Messages: req.Messages,
|
||
Temperature: req.Temperature,
|
||
MaxTokens: req.MaxTokens,
|
||
Stream: req.Stream,
|
||
Tools: req.Tools,
|
||
ToolChoice: req.ToolChoice,
|
||
DisableThinking: req.DisableThinking,
|
||
ExtraBody: req.ExtraBody,
|
||
ClientSession: clientSession,
|
||
}
|
||
rec := &Req{
|
||
Key: keyID(reqKey(ctx)),
|
||
Type: "chat",
|
||
Model: effective,
|
||
Source: firstSource(cands),
|
||
OK: false,
|
||
}
|
||
if req.Stream {
|
||
rec.Type = "stream"
|
||
g.streamChat(w, ctx, cands, inner, effective, rec)
|
||
return
|
||
}
|
||
g.singleChat(w, ctx, cands, inner, effective, rec)
|
||
}
|
||
|
||
func firstModel(p *provider.Provider) string {
|
||
ms := p.Models()
|
||
if len(ms) > 0 {
|
||
return ms[0]
|
||
}
|
||
return "auto"
|
||
}
|
||
|
||
// imageOnly keeps providers exposing at least one image-kind model.
|
||
func imageOnly(cands []*provider.Provider) []*provider.Provider {
|
||
var out []*provider.Provider
|
||
for _, p := range cands {
|
||
for _, id := range p.Models() {
|
||
if m := p.ModelByID(id); m != nil && m.Kind == "image" {
|
||
out = append(out, p)
|
||
break
|
||
}
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
|
||
// toolCallsWire converts unified tool calls to the OpenAI wire format:
|
||
// tool_calls:[{id,type,function:{name,arguments:StringJSON}}]. Clients expect
|
||
// arguments to be a JSON string, not an object.
|
||
func toolCallsWire(tcs []types.ToolCall) json.RawMessage {
|
||
wire := make([]map[string]interface{}, 0, len(tcs))
|
||
for _, tc := range tcs {
|
||
args := "{}"
|
||
if tc.Arguments != nil {
|
||
if b, err := json.Marshal(tc.Arguments); err == nil {
|
||
args = string(b)
|
||
}
|
||
}
|
||
wire = append(wire, map[string]interface{}{
|
||
"id": tc.ID,
|
||
"type": tc.Type,
|
||
"function": map[string]interface{}{
|
||
"name": tc.Name,
|
||
"arguments": args,
|
||
},
|
||
})
|
||
}
|
||
b, _ := json.Marshal(wire)
|
||
return b
|
||
}
|
||
|
||
func firstSource(cands []*provider.Provider) string {
|
||
if len(cands) > 0 {
|
||
return cands[0].Name()
|
||
}
|
||
return ""
|
||
}
|
||
|
||
// effectiveImageModel picks the image model id that will actually be used so
|
||
// quota checks can target it (AUTO resolves to the first image candidate).
|
||
func effectiveImageModel(model string, cands []*provider.Provider) string {
|
||
if !isAuto(model) && model != "" {
|
||
return model
|
||
}
|
||
if len(cands) > 0 {
|
||
for _, id := range cands[0].Models() {
|
||
if m := cands[0].ModelByID(id); m != nil && m.Kind == "image" {
|
||
return id
|
||
}
|
||
}
|
||
}
|
||
return model
|
||
}
|
||
|
||
func estimatePromptTokens(req *types.ChatRequest) int64 {
|
||
if req == nil {
|
||
return 0
|
||
}
|
||
b, _ := json.Marshal(struct {
|
||
Messages []types.ChatMessage `json:"messages"`
|
||
Tools []interface{} `json:"tools,omitempty"`
|
||
}{Messages: req.Messages, Tools: req.Tools})
|
||
if len(b) == 0 {
|
||
return 0
|
||
}
|
||
return int64(len(b)/3 + 1)
|
||
}
|
||
|
||
func estimateTextTokens(parts ...interface{}) int64 {
|
||
var n int
|
||
for _, p := range parts {
|
||
switch v := p.(type) {
|
||
case string:
|
||
n += len(v)
|
||
case json.RawMessage:
|
||
n += len(v)
|
||
case []types.ToolCall:
|
||
b, _ := json.Marshal(v)
|
||
n += len(b)
|
||
}
|
||
}
|
||
if n == 0 {
|
||
return 0
|
||
}
|
||
return int64(n/3 + 1)
|
||
}
|
||
|
||
// toScheduler adapts concrete providers to the scheduler.Provider interface.
|
||
// It lives here (not in the scheduler package) so scheduler tests do not pull
|
||
// in the provider package and with it the Lua runtime's link requirements.
|
||
func toScheduler(cands []*provider.Provider) []scheduler.Provider {
|
||
out := make([]scheduler.Provider, len(cands))
|
||
for i, p := range cands {
|
||
out[i] = p
|
||
}
|
||
return out
|
||
}
|
||
|
||
// upstreamErrStatus maps a scheduling error to its HTTP status: a failed
|
||
// AUTO chain answers 503 with its per-tier summary, a busy source (every
|
||
// concurrency slot in use) is a transient capacity condition answered with
|
||
// 429 so clients fail fast, while other upstream failures stay 502.
|
||
func upstreamErrStatus(err error) int {
|
||
var ce *scheduler.ChainErr
|
||
if errors.As(err, &ce) {
|
||
return http.StatusServiceUnavailable
|
||
}
|
||
if errors.Is(err, provider.ErrBusy) {
|
||
return http.StatusTooManyRequests
|
||
}
|
||
return http.StatusBadGateway
|
||
}
|
||
|
||
// overflowMarkers 匹配上游「上下文超窗」类措辞。
|
||
//
|
||
// 上游写法五花八门,且**不在 pi 客户端的识别列表里**。pi 靠
|
||
// @earendil-works/pi-ai 的 OVERFLOW_PATTERNS 判断超窗并据此触发压缩重试,
|
||
// 而 justworker 返回的是「请精简对话历史…(Context window is full…)」——
|
||
// 与那 25 条正则一条都不匹配,于是 pi 既不压缩也不重试,只把它当成一条
|
||
// 普通上游错误。这里把可识别的超窗措辞归一化成 pi 一定认得的标记。
|
||
var overflowMarkers = []*regexp.Regexp{
|
||
regexp.MustCompile(`(?i)context[ _-]?window is full`),
|
||
regexp.MustCompile(`(?i)context[_ ]length[_ ]exceeded`),
|
||
regexp.MustCompile(`(?i)exceeds? the context window`),
|
||
regexp.MustCompile(`(?i)maximum context length`),
|
||
regexp.MustCompile(`(?i)reduce the length of the messages`),
|
||
regexp.MustCompile(`(?i)too many tokens`),
|
||
regexp.MustCompile(`(?i)token limit exceeded`),
|
||
regexp.MustCompile(`请精简对话历史`),
|
||
regexp.MustCompile(`上下文(长度)?超(出|限)`),
|
||
regexp.MustCompile(`对话历史过长`),
|
||
}
|
||
|
||
// overflowCanonical 命中 pi 的 /context[_ ]length[_ ]exceeded/i。
|
||
//
|
||
// 注意:这个前缀只是必要条件,不是充分条件——pi 还会先用
|
||
// NON_OVERFLOW_PATTERNS 排除整条消息,见 overflowClientMessage。
|
||
const overflowCanonical = "context_length_exceeded"
|
||
|
||
// looksLikeOverflow 判断错误文本是否属于上下文超窗。
|
||
func looksLikeOverflow(s string) bool {
|
||
for _, re := range overflowMarkers {
|
||
if re.MatchString(s) {
|
||
return true
|
||
}
|
||
}
|
||
return false
|
||
}
|
||
|
||
// clientUpstreamErr collapses upstream failure details into a short message
|
||
// for the client: per-tier bodies (WAF HTML pages, quota payloads, ...) stay
|
||
// in rec.Err / the stats API and the server log instead of the response.
|
||
// Per-tier one-line reasons are kept (quota/cooling skips carry no error body
|
||
// and are the actionable part); each is capped so HTML dumps can't leak.
|
||
//
|
||
// 超窗会被归一化成「干净」的 overflowCanonical 消息,见 overflowClientMessage。
|
||
func clientUpstreamErr(err error) string {
|
||
log.Printf("[gateway] upstream failure surfaced to client: %v", err)
|
||
if looksLikeOverflow(err.Error()) {
|
||
return overflowClientMessage(err)
|
||
}
|
||
return upstreamErrSummary(err)
|
||
}
|
||
|
||
// overflowHint 是给客户端的超窗说明,刻意只用 pi 认得的措辞。
|
||
const overflowHint = "context window is full; reduce the length of the messages"
|
||
|
||
// clientSessionFromRequest 取客户端自带的会话标识,拿不到时返回 ""。
|
||
//
|
||
// 通用 OpenAI 客户端默认不带会话 id,但 pi 支持:只要 provider 的 compat 里打开
|
||
// sendSessionAffinityHeaders,pi 就会把**平台真实的会话 id**(uuidv7,整个会话
|
||
// 恒定)放到 x-session-affinity / x-client-request-id 上(sessionAffinityFormat
|
||
// 为 openrouter 时是 x-session-id,为 openai 时是 session_id)。
|
||
// 有了它就不必再从首条 user 消息推导会话指纹(那只在客户端不发会话 id 时才作为
|
||
// 退路,且历史压缩后会漂移)。prompt_cache_key 是 OpenAI 原生的缓存键,
|
||
// 客户端若带也一并当会话用。
|
||
//
|
||
// 刻意不采纳 x-client-request-id:名字含 request,部分客户端每请求都换,
|
||
// 拿它当会话会让上游缓存永不命中。pi 总会同时发 x-session-affinity,够用。
|
||
func clientSessionFromRequest(r *http.Request, bodyKey string) string {
|
||
for _, h := range []string{"X-Session-Affinity", "X-Session-Id", "Session-Id"} {
|
||
if v := strings.TrimSpace(r.Header.Get(h)); v != "" {
|
||
return v
|
||
}
|
||
}
|
||
return strings.TrimSpace(bodyKey)
|
||
}
|
||
|
||
// overflowClientMessage 为超窗失败生成「干净」的客户端消息。
|
||
//
|
||
// 为什么不能沿用 upstreamErrSummary:pi 的 isContextOverflow 先查
|
||
// NON_OVERFLOW_PATTERNS(/rate limit/、/too many requests/、Bedrock 前缀),
|
||
// 一旦命中就直接判为「非超窗」——**哪怕消息里已经有 context_length_exceeded**,
|
||
// pi 也不会压缩重试。而 AUTO 链的失败消息天生是多 tier 原因的拼接:
|
||
// 超窗 tier(gozen 400 maximum context length)常与配额/限流 tier
|
||
// (429 token plan exhausted、cooling、no free slot)同时出现。
|
||
// 把明细原样带出去,等于让一条限流 tier 的措辞反过来封杀超窗识别。
|
||
// 所以这里只保留「超窗」措辞 + 是哪个源超的窗,其余一律不带。
|
||
func overflowClientMessage(err error) string {
|
||
var ce *scheduler.ChainErr
|
||
if errors.As(err, &ce) {
|
||
for _, t := range ce.Tiers {
|
||
if looksLikeOverflow(t.Err.Error()) {
|
||
return fmt.Sprintf("%s: %s (%s/%s)", overflowCanonical, overflowHint, t.Source, t.Model)
|
||
}
|
||
}
|
||
}
|
||
return fmt.Sprintf("%s: %s", overflowCanonical, overflowHint)
|
||
}
|
||
|
||
// upstreamErrSummary 把上游失败压成一行短消息。
|
||
func upstreamErrSummary(err error) string {
|
||
var ce *scheduler.ChainErr
|
||
if errors.As(err, &ce) {
|
||
parts := make([]string, 0, len(ce.Tiers)+len(ce.Skipped))
|
||
for _, t := range ce.Tiers {
|
||
// 160 而不是 80:短诊断词("Context window is full")常落在尾部,
|
||
// 80 字节按字节截断正好会把它切掉,超窗就再也认不出来。
|
||
parts = append(parts, types.OneLine(fmt.Sprintf("%s/%s: %v", t.Source, t.Model, t.Err), 160))
|
||
}
|
||
for _, sk := range ce.Skipped {
|
||
parts = append(parts, types.OneLine(sk, 160))
|
||
}
|
||
msg := strings.Join(parts, "; ")
|
||
if len(msg) > 300 {
|
||
msg = msg[:300] + "..."
|
||
}
|
||
return fmt.Sprintf("all %d auto providers failed: %s", len(parts), msg)
|
||
}
|
||
return types.OneLine(err.Error(), 160)
|
||
}
|
||
|
||
// failChat maps a scheduling failure onto the audit record and answers the
|
||
// client. A failed AUTO chain additionally pins its first failed tier onto
|
||
// the record; direct-path errors never match *scheduler.ChainErr, so the
|
||
// extraction is safely shared by all four entry points. Callers own the
|
||
// writeRec call (inline for non-stream paths, deferred for stream paths).
|
||
func (g *Gateway) failChat(w http.ResponseWriter, rec *Req, err error) {
|
||
rec.OK = false
|
||
rec.Status = upstreamErrStatus(err)
|
||
rec.Err = err.Error()
|
||
var ce *scheduler.ChainErr
|
||
if errors.As(err, &ce) && len(ce.Tiers) > 0 {
|
||
rec.Source = ce.Tiers[0].Source
|
||
rec.Model = ce.Tiers[0].Model
|
||
}
|
||
writeError(w, rec.Status, "upstream_error", clientUpstreamErr(err))
|
||
}
|
||
|
||
// recordChatUsage fills token accounting for a finished non-streaming
|
||
// request: exact upstream numbers win, byte estimates fill the gaps.
|
||
func recordChatUsage(rec *Req, req *types.ChatRequest, resp *types.UnifiedResponse) {
|
||
rec.Prompt = int64(resp.TokenUsage.Prompt)
|
||
if rec.Prompt == 0 {
|
||
rec.Prompt = estimatePromptTokens(req)
|
||
}
|
||
rec.Compl = int64(resp.TokenUsage.Completion)
|
||
if rec.Compl == 0 {
|
||
rec.Compl = estimateTextTokens(resp.Content, resp.ReasoningContent, resp.ToolCalls)
|
||
}
|
||
// Cache accounting: whichever source format the adapter normalized into
|
||
// (prompt_tokens_details.cached_tokens or legacy prompt_cache_hit_tokens),
|
||
// read it back so the request record carries the hit/miss split.
|
||
if d := resp.TokenUsage.PromptTokensDetails; d != nil && d.CachedTokens > 0 {
|
||
rec.CacheHit = int64(d.CachedTokens)
|
||
}
|
||
if resp.TokenUsage.PromptTokensDetails != nil {
|
||
// Upstream reported cache details (even a 0 hit) — tag the row so
|
||
// the UI can show 0% rather than “—”.
|
||
rec.CacheReported = true
|
||
} else if resp.TokenUsage.PromptCacheHit > 0 {
|
||
rec.CacheHit = int64(resp.TokenUsage.PromptCacheHit)
|
||
rec.CacheReported = true
|
||
}
|
||
if resp.TokenUsage.PromptCacheMiss > 0 {
|
||
rec.CacheMiss = int64(resp.TokenUsage.PromptCacheMiss)
|
||
}
|
||
}
|
||
|
||
// writeChatCompletion renders a unified response as an OpenAI
|
||
// chat.completion object. modelName is the id clients see as the serving
|
||
// model: the requested id for direct routes, the exact slot model for AUTO.
|
||
func writeChatCompletion(w http.ResponseWriter, resp *types.UnifiedResponse, modelName string) {
|
||
msg := RespMessage{Role: "assistant", Content: resp.Content}
|
||
if resp.ReasoningContent != "" {
|
||
msg.ReasoningContent = resp.ReasoningContent
|
||
}
|
||
if len(resp.ToolCalls) > 0 {
|
||
msg.ToolCalls = toolCallsWire(resp.ToolCalls)
|
||
}
|
||
out := ChatCompletion{
|
||
ID: newID(),
|
||
Object: "chat.completion",
|
||
Created: time.Now().Unix(),
|
||
Model: modelName,
|
||
Choices: []ChatChoice{{Index: 0, Message: msg, FinishReason: resp.FinishReason}},
|
||
Cost: resp.Cost,
|
||
}
|
||
if resp.TokenUsage.Total > 0 || resp.TokenUsage.Prompt > 0 || resp.TokenUsage.Completion > 0 {
|
||
out.Usage = &resp.TokenUsage
|
||
}
|
||
writeJSON(w, http.StatusOK, out)
|
||
}
|
||
|
||
// singleChat runs a direct (model-pinned) non-streaming request across the
|
||
// candidate list, falling back on failure.
|
||
func (g *Gateway) singleChat(w http.ResponseWriter, ctx context.Context, cands []*provider.Provider, req *types.ChatRequest, effective string, rec *Req) {
|
||
rec.LatMs = 0
|
||
t0 := time.Now()
|
||
resp, usedSrc, usedModel, err := g.core.Scheduler().Chat(ctx, toScheduler(cands), req)
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
if err != nil {
|
||
g.failChat(w, rec, err)
|
||
g.writeRec(rec)
|
||
return
|
||
}
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
recordChatUsage(rec, req, resp)
|
||
rec.Source = usedSrc
|
||
rec.Model = usedModel
|
||
g.fireRouted(ctx, "chat", usedSrc, usedModel, -1, false)
|
||
// Non-streaming: the whole response arrives at once, so TTFB equals
|
||
// the total latency.
|
||
rec.FirstByteMs = rec.LatMs
|
||
g.writeRec(rec)
|
||
writeChatCompletion(w, resp, effective)
|
||
}
|
||
|
||
// writeRec records a finished request (audit + aggregates) and fires the
|
||
// plugin request_end stage.
|
||
//
|
||
// This is the ONE place every request passes through on its way out, which is
|
||
// what makes it the right hook point: the four entry points (single/stream ×
|
||
// direct/auto) all funnel here, so a plugin sees each request exactly once with
|
||
// its final accounting. Firing earlier would miss the streamed ones (their
|
||
// numbers are only known once the stream finishes), and firing in each entry
|
||
// point would mean four call sites to keep in sync.
|
||
//
|
||
// Hooks run AFTER the record is written: a plugin must not be able to delay or
|
||
// lose the audit trail, and a plugin that throws is contained by Fire.
|
||
func (g *Gateway) writeRec(rec *Req) {
|
||
if rec == nil {
|
||
return
|
||
}
|
||
if rec.Time == 0 {
|
||
rec.Time = time.Now().UnixMilli()
|
||
}
|
||
g.stats.Record(*rec)
|
||
g.fireEnd(rec)
|
||
}
|
||
|
||
// fireStart dispatches the plugin request_start stage: the request has been
|
||
// parsed and authorized but no upstream slot has been chosen yet, so `source`
|
||
// is empty. A plugin that only wants volume/acceptance counts can subscribe
|
||
// here and stay out of the per-request hot path entirely.
|
||
func (g *Gateway) fireStart(ctx context.Context, req *chatRequest, kind, model string, msgs, tools int) {
|
||
ps := g.core.Plugins()
|
||
if ps == nil || ps.Count() == 0 {
|
||
return
|
||
}
|
||
ps.Fire(lua.StageRequestStart, map[string]interface{}{
|
||
"stage": string(lua.StageRequestStart),
|
||
"type": kind,
|
||
"model": model,
|
||
"key": keyID(reqKey(ctx)),
|
||
"role": reqRole(ctx),
|
||
"source": "",
|
||
"stream": req.Stream,
|
||
"messages_count": msgs,
|
||
"tools_count": tools,
|
||
"ts": time.Now().Unix(),
|
||
})
|
||
}
|
||
|
||
// fireImageStart dispatches request_start for /v1/images/generations.
|
||
//
|
||
// It is a separate function rather than a call to fireStart with a nil
|
||
// chatRequest because the image body has no messages and no tools: passing
|
||
// zeroes through a struct built for chat would invite someone to read a field
|
||
// that simply does not exist on this path.
|
||
func (g *Gateway) fireImageStart(ctx context.Context, model string) {
|
||
ps := g.core.Plugins()
|
||
if ps == nil || ps.Count() == 0 {
|
||
return
|
||
}
|
||
ps.Fire(lua.StageRequestStart, map[string]interface{}{
|
||
"stage": string(lua.StageRequestStart),
|
||
"type": "image",
|
||
"model": model,
|
||
"key": keyID(reqKey(ctx)),
|
||
"role": reqRole(ctx),
|
||
"source": "",
|
||
"stream": false,
|
||
"messages_count": 0,
|
||
"tools_count": 0,
|
||
"ts": time.Now().Unix(),
|
||
})
|
||
}
|
||
|
||
// fireRouted dispatches the plugin routed stage once a (source, model) slot has
|
||
// been selected. tier is the AUTO tier index, or -1 on the direct path, so a
|
||
// plugin can tell "this came from tier 1" from "this bypassed the chain".
|
||
func (g *Gateway) fireRouted(ctx context.Context, kind, source, model string, tier int, stream bool) {
|
||
ps := g.core.Plugins()
|
||
if ps == nil || ps.Count() == 0 {
|
||
return
|
||
}
|
||
ps.Fire(lua.StageRouted, map[string]interface{}{
|
||
"stage": string(lua.StageRouted),
|
||
"type": kind,
|
||
"source": source,
|
||
"model": model,
|
||
"key": keyID(reqKey(ctx)),
|
||
"tier": tier,
|
||
"stream": stream,
|
||
"ts": time.Now().Unix(),
|
||
})
|
||
}
|
||
|
||
// chainTraceSink adapts a scheduler TraceSink into the plugin chain_step stage.
|
||
//
|
||
// It returns nil when no plugin is loaded, so the scheduler's emit() does a
|
||
// single nil check per event and the AUTO hot path pays nothing on a gateway
|
||
// with no plugins.
|
||
//
|
||
// The events are also accumulated into walk so request_end can carry a compact
|
||
// summary: a plugin that only listens to request_end still learns that a
|
||
// degradation happened, which is the common case for a dashboard that does not
|
||
// want to subscribe to a high-frequency stage.
|
||
func (g *Gateway) chainTraceSink(ctx context.Context, kind string, walk *[]map[string]interface{}) scheduler.TraceSink {
|
||
ps := g.core.Plugins()
|
||
if ps == nil || ps.Count() == 0 {
|
||
return nil
|
||
}
|
||
key := keyID(reqKey(ctx))
|
||
return func(ev scheduler.TraceEvent) {
|
||
payload := map[string]interface{}{
|
||
"stage": string(lua.StageChainStep),
|
||
"kind": string(ev.Kind),
|
||
"type": kind,
|
||
"key": key,
|
||
"tier": ev.Tier,
|
||
"attempt": ev.Attempt,
|
||
}
|
||
if ev.Source != "" {
|
||
payload["source"] = ev.Source
|
||
}
|
||
if ev.Model != "" {
|
||
payload["model"] = ev.Model
|
||
}
|
||
if ev.Reason != "" {
|
||
payload["reason"] = ev.Reason
|
||
}
|
||
if ev.Err != "" {
|
||
payload["error"] = ev.Err
|
||
}
|
||
if walk != nil {
|
||
// Keep the summary bounded: a pathological chain could emit many
|
||
// steps, and request_end's payload is written to the audit trail.
|
||
if len(*walk) < maxWalkSummary {
|
||
*walk = append(*walk, map[string]interface{}{
|
||
"kind": string(ev.Kind), "tier": ev.Tier,
|
||
"source": ev.Source, "model": ev.Model, "reason": ev.Reason,
|
||
})
|
||
}
|
||
}
|
||
ps.Fire(lua.StageChainStep, payload)
|
||
}
|
||
}
|
||
|
||
// maxWalkSummary caps how many chain steps request_end carries, so a long
|
||
// degradation cannot inflate every audit record.
|
||
const maxWalkSummary = 12
|
||
|
||
// tierServed returns the AUTO tier that actually served the request, or -1 when
|
||
// the walk is empty (a direct request) or ended without a selection (total
|
||
// failure). It is the single most useful number for "why did my expensive tier
|
||
// not get used".
|
||
func tierServed(walk []map[string]interface{}) int {
|
||
for i := len(walk) - 1; i >= 0; i-- {
|
||
if k, _ := walk[i]["kind"].(string); k == string(scheduler.TraceSelected) {
|
||
if t, ok := walk[i]["tier"].(int); ok {
|
||
return t
|
||
}
|
||
}
|
||
}
|
||
return -1
|
||
}
|
||
|
||
// fireEnd dispatches the plugin request_end stage for one finished request.
|
||
func (g *Gateway) fireEnd(rec *Req) {
|
||
ps := g.core.Plugins()
|
||
if ps == nil || ps.Count() == 0 {
|
||
return
|
||
}
|
||
payload := map[string]interface{}{
|
||
"stage": string(lua.StageRequestEnd),
|
||
"type": rec.Type,
|
||
"model": rec.Model,
|
||
"source": rec.Source,
|
||
"key": rec.Key,
|
||
"ok": rec.OK,
|
||
"status": rec.Status,
|
||
"latency_ms": rec.LatMs,
|
||
"first_byte_ms": rec.FirstByteMs,
|
||
"prompt_tokens": rec.Prompt,
|
||
"completion_tokens": rec.Compl,
|
||
"cache_hit_tokens": rec.CacheHit,
|
||
"cache_miss_tokens": rec.CacheMiss,
|
||
"image_count": rec.ImageCount,
|
||
"error": rec.Err,
|
||
"time": rec.Time,
|
||
// chain_walk: the AUTO tier-by-tier trace, when the request went
|
||
// through the chain. Empty for a direct request and for a gateway with
|
||
// no plugins loaded. Absent rather than empty so a plugin can tell
|
||
// "no chain" from "chain with no degradation".
|
||
"degraded": len(rec.Walk) > 1,
|
||
"chain_walk": rec.Walk,
|
||
"tier_served": tierServed(rec.Walk),
|
||
}
|
||
// The merged result is intentionally discarded: request_end is the last
|
||
// stage, so there is nobody downstream to read a plugin's additions. Plugins
|
||
// that need to publish derived numbers (the billing plugin) do it in their
|
||
// OWN state and expose them through the /api/plugins/<name>/state endpoint.
|
||
ps.Fire(lua.StageRequestEnd, payload)
|
||
}
|
||
|
||
// mergeUsage combines token usage across stream chunks additively. Some
|
||
// providers split usage across chunks (e.g. Anthropic reports prompt tokens
|
||
// in message_start and the final completion tokens in message_delta); a plain
|
||
// "last non-nil wins" would discard the prompt half. Non-zero fields from cur
|
||
// override prev; total is recomputed from the merged parts so a partial later
|
||
// chunk can't shrink it. For the common single-chunk case (OpenAI's terminal
|
||
// empty-choices+usage chunk) upstream totals are preserved exactly.
|
||
func mergeUsage(prev, cur *types.TokenUsage) *types.TokenUsage {
|
||
if prev == nil {
|
||
u := *cur
|
||
if u.Total == 0 && (u.Prompt > 0 || u.Completion > 0) {
|
||
u.Total = u.Prompt + u.Completion
|
||
}
|
||
return &u
|
||
}
|
||
out := *prev
|
||
if cur.Prompt > 0 {
|
||
out.Prompt = cur.Prompt
|
||
}
|
||
if cur.Completion > 0 {
|
||
out.Completion = cur.Completion
|
||
}
|
||
out.Total = out.Prompt + out.Completion
|
||
if cur.PromptTokensDetails != nil {
|
||
// Keep the details object even when CachedTokens is 0: a reported
|
||
// zero-hit is meaningful ("cache missed") and must stay
|
||
// distinguishable from "upstream never reported cache info".
|
||
// Dropping it here made streaming rows show “—” instead of 0%.
|
||
out.PromptTokensDetails = cur.PromptTokensDetails
|
||
}
|
||
if cur.PromptCacheHit > 0 {
|
||
out.PromptCacheHit = cur.PromptCacheHit
|
||
}
|
||
if cur.PromptCacheMiss > 0 {
|
||
out.PromptCacheMiss = cur.PromptCacheMiss
|
||
}
|
||
return &out
|
||
}
|
||
|
||
// pumpStream writes the full SSE sequence for a started stream: role
|
||
// preamble, one chunk per unified delta, the terminating finish_reason, the
|
||
// OpenAI-standard final usage chunk (empty choices) and [DONE]. modelName
|
||
// follows writeChatCompletion's rule (requested id for direct routes, exact
|
||
// slot model for AUTO).
|
||
func (g *Gateway) pumpStream(w http.ResponseWriter, rec *Req, chunks <-chan types.UnifiedChunk, modelName string, t0 time.Time) {
|
||
w.Header().Set("Content-Type", "text/event-stream")
|
||
w.Header().Set("Cache-Control", "no-cache")
|
||
w.Header().Set("Connection", "keep-alive")
|
||
w.WriteHeader(http.StatusOK)
|
||
|
||
flusher, _ := w.(http.Flusher)
|
||
id := newID()
|
||
created := time.Now().Unix()
|
||
send := func(obj interface{}) bool {
|
||
b, err := json.Marshal(obj)
|
||
if err != nil {
|
||
return false
|
||
}
|
||
if _, err := fmt.Fprintf(w, "data: %s\n\n", b); err != nil {
|
||
return false
|
||
}
|
||
if flusher != nil {
|
||
flusher.Flush()
|
||
}
|
||
return true
|
||
}
|
||
|
||
if !send(ChatChunk{
|
||
ID: id, Object: "chat.completion.chunk", Created: created, Model: modelName,
|
||
Choices: []ChunkChoice{{Index: 0, Delta: RespMessage{Role: "assistant"}}},
|
||
}) {
|
||
return
|
||
}
|
||
// First SSE byte sent to the client: record time-to-first-byte for the
|
||
// source's status-page latency average.
|
||
rec.FirstByteMs = time.Since(t0).Milliseconds()
|
||
var lastUsage *types.TokenUsage
|
||
var lastCost string
|
||
lastFinish := ""
|
||
for ck := range chunks {
|
||
if ck.Usage != nil {
|
||
lastUsage = mergeUsage(lastUsage, ck.Usage)
|
||
}
|
||
if ck.Cost != "" {
|
||
lastCost = ck.Cost
|
||
}
|
||
chunk := ChatChunk{
|
||
ID: id, Object: "chat.completion.chunk", Created: created, Model: modelName,
|
||
}
|
||
delta := RespMessage{Role: "assistant", Content: ck.Content}
|
||
if ck.ReasoningContent != "" {
|
||
delta.ReasoningContent = ck.ReasoningContent
|
||
}
|
||
if len(ck.ToolCalls) > 0 {
|
||
delta.ToolCalls = ck.ToolCalls
|
||
}
|
||
choice := ChunkChoice{Index: 0, Delta: delta}
|
||
if ck.Done {
|
||
fin := ck.FinishReason
|
||
if fin == "" {
|
||
fin = "stop"
|
||
}
|
||
lastFinish = fin
|
||
choice.FinishReason = &fin
|
||
}
|
||
chunk.Choices = []ChunkChoice{choice}
|
||
rec.Compl += int64(len(ck.Content)+len(ck.ReasoningContent)+len(ck.ToolCalls)) / 3
|
||
if !send(chunk) {
|
||
return
|
||
}
|
||
}
|
||
finalFinish := lastFinish
|
||
if finalFinish == "" {
|
||
finalFinish = "stop"
|
||
}
|
||
send(ChatChunk{
|
||
ID: id, Object: "chat.completion.chunk", Created: created, Model: modelName,
|
||
Choices: []ChunkChoice{{Index: 0, Delta: RespMessage{}, FinishReason: &finalFinish}},
|
||
})
|
||
// Final usage chunk (OpenAI standard: empty choices + usage before [DONE]).
|
||
// Prefer the upstream's exact usage if the stream carried it; fall back to
|
||
// the gateway's estimate otherwise.
|
||
// Write the upstream's exact numbers back onto the audit record. Without
|
||
// this the streamed path kept only the per-chunk byte estimate, so the same
|
||
// request recorded ~1.4x its real completion tokens (measured: upstream 100,
|
||
// audit 145) while the non-streaming path recorded 100. Two paths, two
|
||
// different numbers for one request is a reporting bug, not a rounding one.
|
||
if lastUsage != nil {
|
||
if lastUsage.Prompt > 0 {
|
||
rec.Prompt = int64(lastUsage.Prompt)
|
||
}
|
||
if lastUsage.Completion > 0 {
|
||
rec.Compl = int64(lastUsage.Completion)
|
||
}
|
||
}
|
||
var tut *types.TokenUsage
|
||
if lastUsage != nil {
|
||
tut = lastUsage
|
||
// Write the upstream's cache accounting back onto the request record
|
||
// so the audit trail carries the hit/miss split for streaming too.
|
||
if d := tut.PromptTokensDetails; d != nil {
|
||
rec.CacheReported = true
|
||
if d.CachedTokens > 0 {
|
||
rec.CacheHit = int64(d.CachedTokens)
|
||
}
|
||
} else if tut.PromptCacheHit > 0 {
|
||
rec.CacheHit = int64(tut.PromptCacheHit)
|
||
rec.CacheReported = true
|
||
}
|
||
if tut.PromptCacheMiss > 0 {
|
||
rec.CacheMiss = int64(tut.PromptCacheMiss)
|
||
}
|
||
} else if rec.Prompt+rec.Compl > 0 {
|
||
u := types.TokenUsage{
|
||
Prompt: int(rec.Prompt),
|
||
Completion: int(rec.Compl),
|
||
Total: int(rec.Prompt + rec.Compl),
|
||
}
|
||
tut = &u
|
||
}
|
||
if tut != nil {
|
||
send(ChatChunk{
|
||
ID: id, Object: "chat.completion.chunk", Created: created, Model: modelName,
|
||
Choices: []ChunkChoice{},
|
||
Usage: tut,
|
||
Cost: lastCost,
|
||
})
|
||
} else if lastCost != "" {
|
||
// Cost arrived without any usage (OpenCode sends it on its own frame).
|
||
send(ChatChunk{
|
||
ID: id, Object: "chat.completion.chunk", Created: created, Model: modelName,
|
||
Choices: []ChunkChoice{},
|
||
Cost: lastCost,
|
||
})
|
||
}
|
||
fmt.Fprintf(w, "data: [DONE]\n\n")
|
||
if flusher != nil {
|
||
flusher.Flush()
|
||
}
|
||
}
|
||
|
||
// streamChat runs a direct (model-pinned) streaming request across the
|
||
// candidate list, falling back early on connect errors.
|
||
func (g *Gateway) streamChat(w http.ResponseWriter, ctx context.Context, cands []*provider.Provider, req *types.ChatRequest, effective string, rec *Req) {
|
||
rec.LatMs = 0
|
||
t0 := time.Now()
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
defer func() {
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
g.writeRec(rec)
|
||
}()
|
||
chunks, usedSrc, usedModel, err := g.core.Scheduler().ChatStream(ctx, toScheduler(cands), req)
|
||
if err != nil {
|
||
g.failChat(w, rec, err)
|
||
return
|
||
}
|
||
if usedModel != "" {
|
||
rec.Model = usedModel
|
||
}
|
||
// Audit accuracy: pin the source that actually served the stream (after a
|
||
// failover it differs from the first candidate). Direct streams previously
|
||
// discarded it.
|
||
rec.Source = usedSrc
|
||
g.fireRouted(ctx, "stream", usedSrc, rec.Model, -1, true)
|
||
rec.Prompt = estimatePromptTokens(req)
|
||
g.pumpStream(w, rec, chunks, effective, t0)
|
||
}
|
||
|
||
// singleChatAuto runs a non-streaming AUTO request down the chain (see
|
||
// scheduler.ChainChat): tiers ascending (tier 1 = highest priority first),
|
||
// per-tier round-robin ordered by preference, cooldown as the only hard skip,
|
||
// busy slots skipped without penalty and a bounded busy wait. When every
|
||
// tier fails, the response is a 503 carrying the per-tier error summary
|
||
// (which source/model failed why).
|
||
func (g *Gateway) singleChatAuto(w http.ResponseWriter, ctx context.Context, chain *scheduler.Chain, req *types.ChatRequest, rec *Req, quotaExhausted func(*scheduler.Slot) bool) {
|
||
rec.LatMs = 0
|
||
t0 := time.Now()
|
||
var walk []map[string]interface{}
|
||
resp, usedSrc, usedModel, err := g.core.Scheduler().ChainChat(ctx, chain, req, quotaExhausted,
|
||
g.chainTraceSink(ctx, "chat", &walk))
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
rec.Walk = walk
|
||
if err != nil {
|
||
g.failChat(w, rec, err)
|
||
g.writeRec(rec)
|
||
return
|
||
}
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
recordChatUsage(rec, req, resp)
|
||
rec.Source = usedSrc
|
||
rec.Model = usedModel
|
||
// AUTO has no single tier to report: the chain may have walked several
|
||
// before this slot served the request, so -2 means "resolved by the chain"
|
||
// and a plugin can tell that apart from the direct path's -1.
|
||
g.fireRouted(ctx, "chat", usedSrc, usedModel, -2, false)
|
||
rec.FirstByteMs = rec.LatMs
|
||
g.writeRec(rec)
|
||
writeChatCompletion(w, resp, usedModel)
|
||
}
|
||
|
||
// streamChatAuto streams an AUTO request down the chain. A slot is abandoned
|
||
// only before its first chunk (connect error / non-200 / busy); once a stream
|
||
// starts it stays pinned. Total failure writes a JSON 503 (with the per-tier
|
||
// summary) before any SSE byte is sent.
|
||
func (g *Gateway) streamChatAuto(w http.ResponseWriter, ctx context.Context, chain *scheduler.Chain, req *types.ChatRequest, rec *Req, quotaExhausted func(*scheduler.Slot) bool) {
|
||
rec.LatMs = 0
|
||
t0 := time.Now()
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
defer func() {
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
g.writeRec(rec)
|
||
}()
|
||
var walk []map[string]interface{}
|
||
chunks, usedSrc, usedModel, err := g.core.Scheduler().ChainChatStream(ctx, chain, req, quotaExhausted,
|
||
g.chainTraceSink(ctx, "stream", &walk))
|
||
rec.Walk = walk
|
||
if err != nil {
|
||
g.failChat(w, rec, err)
|
||
return
|
||
}
|
||
if usedModel != "" {
|
||
rec.Model = usedModel
|
||
}
|
||
rec.Source = usedSrc
|
||
g.fireRouted(ctx, "stream", usedSrc, rec.Model, -2, true)
|
||
rec.Prompt = estimatePromptTokens(req)
|
||
g.pumpStream(w, rec, chunks, usedModel, t0)
|
||
}
|
||
|
||
func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) {
|
||
if r.Method != http.MethodPost {
|
||
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "use POST")
|
||
return
|
||
}
|
||
var req types.ImageGenRequest
|
||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||
writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error())
|
||
return
|
||
}
|
||
if req.Prompt == "" {
|
||
writeError(w, http.StatusBadRequest, "invalid_request", "prompt is required")
|
||
return
|
||
}
|
||
model := req.Model
|
||
if model == "" {
|
||
model = g.core.DefaultModel()
|
||
}
|
||
// Same rule as the chat path, and for the same reason: an image request is
|
||
// billable traffic, so a cost plugin must see it. It fires before the
|
||
// quota/scope gates so rejected image requests are visible too.
|
||
// messages_count/tools_count are 0: the image request has neither.
|
||
g.fireImageStart(r.Context(), model)
|
||
if isAuto(model) {
|
||
if chain := g.core.AutoImageChain(); chain != nil && len(chain.Tiers) > 0 {
|
||
if q := g.checkQuota(r.Context(), "AUTO"); q != nil {
|
||
g.writeReject(w, q)
|
||
return
|
||
}
|
||
done := g.stats.Begin()
|
||
defer done()
|
||
rec := &Req{Key: keyID(reqKey(r.Context())), Type: "image", Model: model, OK: false}
|
||
t0 := time.Now()
|
||
resp, usedSrc, usedModel, err := g.core.Scheduler().ChainImage(r.Context(), chain, &req)
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
if err != nil {
|
||
rec.Status = upstreamErrStatus(err)
|
||
rec.Err = err.Error()
|
||
g.writeRec(rec)
|
||
writeError(w, rec.Status, "upstream_error", clientUpstreamErr(err))
|
||
return
|
||
}
|
||
rec.Source = usedSrc
|
||
if usedModel != "" {
|
||
rec.Model = usedModel // actual image model served, not "AUTO"
|
||
}
|
||
g.fireRouted(r.Context(), "image", usedSrc, rec.Model, -2, false)
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
// Image generation has no token concept. Recording len(ImageData)
|
||
// (the image COUNT) in completion_tokens mislabels image count as
|
||
// tokens and feeds it into the token totals; leave it 0.
|
||
rec.ImageCount = len(resp.ImageData)
|
||
g.writeRec(rec)
|
||
writeJSON(w, http.StatusOK, types.ImageGenResponse{
|
||
Created: time.Now().Unix(),
|
||
Data: resp.ImageData,
|
||
})
|
||
return
|
||
}
|
||
// no image chain configured: fall through to legacy discovery (all
|
||
// sources exposing an image model, tried in registry order)
|
||
}
|
||
if allow := g.allowedModels(r.Context()); allow != nil && !g.hasScopeModel(allow, model) {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", fmt.Sprintf("model %q is not allowed for this key", model))
|
||
return
|
||
}
|
||
cands, _ := g.resolveByModel(model)
|
||
cands = imageOnly(cands)
|
||
if allow := g.allowedModels(r.Context()); allow != nil {
|
||
cands = filterCandsByModels(cands, allow)
|
||
}
|
||
if len(cands) == 0 {
|
||
writeError(w, http.StatusServiceUnavailable, "no_provider", "no image source configured")
|
||
return
|
||
}
|
||
if q := g.checkQuota(r.Context(), effectiveImageModel(model, cands)); q != nil {
|
||
g.writeReject(w, q)
|
||
return
|
||
}
|
||
done := g.stats.Begin()
|
||
defer done()
|
||
rec := &Req{Key: keyID(reqKey(r.Context())), Type: "image", Model: model, Source: firstSource(cands), OK: false}
|
||
t0 := time.Now()
|
||
resp, usedSrc, err := g.core.Scheduler().Image(r.Context(), toScheduler(cands), &req)
|
||
rec.LatMs = time.Since(t0).Milliseconds()
|
||
if err != nil {
|
||
rec.Status = upstreamErrStatus(err)
|
||
rec.Err = err.Error()
|
||
g.writeRec(rec)
|
||
writeError(w, rec.Status, "upstream_error", clientUpstreamErr(err))
|
||
return
|
||
}
|
||
if usedSrc != "" {
|
||
rec.Source = usedSrc
|
||
}
|
||
if resp.Model != "" {
|
||
rec.Model = resp.Model // record the actual model served, not the raw request id
|
||
}
|
||
g.fireRouted(r.Context(), "image", rec.Source, rec.Model, -1, false)
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
// Image generation has no token concept — see the AUTO path above.
|
||
rec.ImageCount = len(resp.ImageData)
|
||
g.writeRec(rec)
|
||
writeJSON(w, http.StatusOK, types.ImageGenResponse{
|
||
Created: time.Now().Unix(),
|
||
Data: resp.ImageData,
|
||
})
|
||
}
|