mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-09-20 17:07:59 +00:00
问题:上游返回上下文超窗时,客户端(pi)既不压缩也不重试,只看到一条 普通上游错误。链路两处叠加: 1. 措辞不在客户端识别列表里。pi 靠 @earendil-works/pi-ai 的 OVERFLOW_PATTERNS(25 条正则)判断超窗,而 justworker 返回的是 「请精简对话历史…(Context window is full…)」——与那 25 条一条都不匹配 (最接近的 /context window exceeds limit/i 也不命中,因为措辞是 "is full" 而非 "exceeds limit")。 2. AUTO 链路的每条 tier 错误按 80 字节 OneLine 截断,而 "Context window is full" 这类短诊断词常出现在尾部,正好被按字节切掉(直连路径是 160 才保住)。 修法: - 新增 overflowMarkers 识别超窗措辞(含中文写法),命中时把客户端可见 错误归一化为 context_length_exceeded 前缀——它命中 pi 的 /context[_ ]length[_ ]exceeded/i,超窗因此可被发现并触发压缩重试。 - AUTO 链路每条 tier 错误宽度 80→160,短诊断词不再被截断。 归一化只加前缀,原始诊断信息保留,便于定位是哪一层超窗。 测试:internal/gateway/overflow_err_test.go(5 例,含「标记必须命中 pi 正则」、非超窗不得误标、80 vs 160 宽度的回归对比)。
1027 lines
32 KiB
Go
1027 lines
32 KiB
Go
package gateway
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"log"
|
||
"net/http"
|
||
"regexp"
|
||
"strings"
|
||
"sync/atomic"
|
||
"time"
|
||
|
||
"llmsproxy/internal/config"
|
||
"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"`
|
||
}
|
||
|
||
// 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"`
|
||
}
|
||
|
||
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"`
|
||
}
|
||
|
||
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.
|
||
func (g *Gateway) checkModelScope(ctx context.Context, model string) string {
|
||
allow := g.allowedModels(ctx)
|
||
if allow == nil {
|
||
return ""
|
||
}
|
||
for _, sc := range allow {
|
||
if sc.Model != model {
|
||
continue
|
||
}
|
||
if sc.TokenQuota > 0 {
|
||
used := g.scopeTokens(ctx, sc)
|
||
if used >= sc.TokenQuota {
|
||
return fmt.Sprintf("token quota exceeded for %q (%d/%d)", model, used, sc.TokenQuota)
|
||
}
|
||
}
|
||
return ""
|
||
}
|
||
return fmt.Sprintf("model %q is not allowed for this key", model)
|
||
}
|
||
|
||
// scopeTokens returns the tokens a scope entry has consumed within its reset
|
||
// window (total for AUTO / per model otherwise).
|
||
func (g *Gateway) scopeTokens(ctx context.Context, sc config.ModelScope) int64 {
|
||
k := keyID(reqKey(ctx))
|
||
if sc.Model != "" && strings.EqualFold(sc.Model, "AUTO") {
|
||
return g.stats.KeyTokens(k)
|
||
}
|
||
win := AutoPeriodSeconds(sc.Period, sc.Hours)
|
||
return g.stats.WindowTokens(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
|
||
}
|
||
|
||
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
|
||
}
|
||
if len(req.Messages) == 0 {
|
||
writeError(w, http.StatusBadRequest, "invalid_request", "messages is required")
|
||
return
|
||
}
|
||
model := req.Model
|
||
if model == "" {
|
||
model = g.core.DefaultModel()
|
||
}
|
||
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 msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
||
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,
|
||
}
|
||
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 msg := g.checkModelScope(r.Context(), effective); msg != "" {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
||
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,
|
||
}
|
||
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。
|
||
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 前缀,见 overflowMarkers 的说明。
|
||
func clientUpstreamErr(err error) string {
|
||
log.Printf("[gateway] upstream failure surfaced to client: %v", err)
|
||
msg := upstreamErrSummary(err)
|
||
if looksLikeOverflow(err.Error()) || looksLikeOverflow(msg) {
|
||
return fmt.Sprintf("%s: %s", overflowCanonical, msg)
|
||
}
|
||
return msg
|
||
}
|
||
|
||
// 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}},
|
||
}
|
||
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
|
||
// 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).
|
||
func (g *Gateway) writeRec(rec *Req) {
|
||
if rec == nil {
|
||
return
|
||
}
|
||
if rec.Time == 0 {
|
||
rec.Time = time.Now().UnixMilli()
|
||
}
|
||
g.stats.Record(*rec)
|
||
}
|
||
|
||
// 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
|
||
lastFinish := ""
|
||
for ck := range chunks {
|
||
if ck.Usage != nil {
|
||
lastUsage = mergeUsage(lastUsage, ck.Usage)
|
||
}
|
||
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.
|
||
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,
|
||
})
|
||
}
|
||
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
|
||
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()
|
||
resp, usedSrc, usedModel, err := g.core.Scheduler().ChainChat(ctx, chain, req, quotaExhausted)
|
||
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
|
||
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)
|
||
}()
|
||
chunks, usedSrc, usedModel, err := g.core.Scheduler().ChainChatStream(ctx, chain, req, quotaExhausted)
|
||
if err != nil {
|
||
g.failChat(w, rec, err)
|
||
return
|
||
}
|
||
if usedModel != "" {
|
||
rec.Model = usedModel
|
||
}
|
||
rec.Source = usedSrc
|
||
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()
|
||
}
|
||
if isAuto(model) {
|
||
if chain := g.core.AutoImageChain(); chain != nil && len(chain.Tiers) > 0 {
|
||
if msg := g.checkModelScope(r.Context(), "AUTO"); msg != "" {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
||
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"
|
||
}
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
rec.Compl = int64(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 msg := g.checkModelScope(r.Context(), effectiveImageModel(model, cands)); msg != "" {
|
||
writeError(w, http.StatusForbidden, "model_not_allowed", msg)
|
||
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
|
||
}
|
||
rec.OK = true
|
||
rec.Status = http.StatusOK
|
||
rec.Compl = int64(len(resp.ImageData))
|
||
g.writeRec(rec)
|
||
writeJSON(w, http.StatusOK, types.ImageGenResponse{
|
||
Created: time.Now().Unix(),
|
||
Data: resp.ImageData,
|
||
})
|
||
}
|