mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 09:35:02 -04:00
fix(localai-proxy): keep rerank, stream errors and chat intact through the proxy
Rerank no longer sends top_n 0, which the upstream rejects. A mid-stream upstream error frame now fails the call instead of ending it as a short success. Temperature 0 is forwarded. An upstream 429 becomes ResourceExhausted, which failover skips like Unimplemented. A localai-proxy config sends its own name upstream when upstream_model is unset, and a chat proxy defaults to the tokenizer template so chat reaches the upstream as messages. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
6d5c9600d7
commit
6df1767133
12 files changed
+242
-10
No files matched your search
@@ -189,7 +189,7 @@ func transportError(path string, err error) error {
|
||||
// statusError maps a non-2xx upstream reply to a gRPC status. 5xx means the
|
||||
// upstream is unhealthy (Unavailable, so failover retries elsewhere); 4xx
|
||||
// means the request itself is wrong (InvalidArgument, so failover does not
|
||||
// trip a healthy target over a client error). 501 is the upstream saying it
|
||||
// trip a healthy target over a client error), except 429. 501 is the upstream saying it
|
||||
// cannot serve this kind of request, which failover treats as a capability
|
||||
// gap, like our own Unimplemented methods.
|
||||
func statusError(path string, resp *http.Response) error {
|
||||
@@ -208,6 +208,10 @@ func statusError(path string, resp *http.Response) error {
|
||||
code = codes.Unimplemented
|
||||
case resp.StatusCode >= 500:
|
||||
code = codes.Unavailable
|
||||
case resp.StatusCode == http.StatusTooManyRequests:
|
||||
// Rate limited: the upstream is healthy but out of capacity, so
|
||||
// failover skips to the next target without tripping this one.
|
||||
code = codes.ResourceExhausted
|
||||
case resp.StatusCode >= 400:
|
||||
code = codes.InvalidArgument
|
||||
default:
|
||||
|
||||
@@ -16,13 +16,15 @@ import (
|
||||
// textRequest is the body for /v1/chat/completions (Messages) and
|
||||
// /v1/completions (Prompt). Zero sampling values are omitted so the upstream
|
||||
// model's own config defaults apply, as they would for a direct caller.
|
||||
// Temperature is the exception: 0 is a real choice (greedy decoding), and
|
||||
// core always fills it from the model config, so it is always sent.
|
||||
type textRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []chatMessage `json:"messages,omitempty"`
|
||||
Prompt string `json:"prompt,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
MaxTokens int32 `json:"max_tokens,omitempty"`
|
||||
Temperature float32 `json:"temperature,omitempty"`
|
||||
Temperature float32 `json:"temperature"`
|
||||
TopP float32 `json:"top_p,omitempty"`
|
||||
TopK int32 `json:"top_k,omitempty"`
|
||||
Seed int32 `json:"seed,omitempty"`
|
||||
@@ -67,6 +69,11 @@ type choiceDelta struct {
|
||||
}
|
||||
|
||||
type textResponse struct {
|
||||
// Error is set on the frame LocalAI sends when generation fails after
|
||||
// the stream has started (followed by [DONE]).
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
Choices []textChoice `json:"choices"`
|
||||
Usage *struct {
|
||||
PromptTokens int32 `json:"prompt_tokens"`
|
||||
@@ -79,6 +86,9 @@ type textResponse struct {
|
||||
// too); otherwise core already rendered the prompt and completions takes it
|
||||
// verbatim.
|
||||
func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string, textRequest) {
|
||||
if dropped := unforwardedFields(opts); len(dropped) > 0 {
|
||||
xlog.Warn("localai-proxy: request fields are not forwarded upstream", "fields", dropped)
|
||||
}
|
||||
req := textRequest{
|
||||
Model: p.model(""),
|
||||
Stream: stream,
|
||||
@@ -113,6 +123,27 @@ func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string
|
||||
return "/v1/chat/completions", req
|
||||
}
|
||||
|
||||
// unforwardedFields names the request inputs the REST text endpoints cannot
|
||||
// carry from here: a grammar core compiled locally, and media that core hands
|
||||
// over as local paths or base64 outside the messages. They are dropped, so
|
||||
// say so instead of letting the answer silently ignore them.
|
||||
func unforwardedFields(opts *pb.PredictOptions) []string {
|
||||
var out []string
|
||||
if opts.GetGrammar() != "" {
|
||||
out = append(out, "grammar")
|
||||
}
|
||||
if len(opts.GetImages()) > 0 {
|
||||
out = append(out, "images")
|
||||
}
|
||||
if len(opts.GetAudios()) > 0 {
|
||||
out = append(out, "audios")
|
||||
}
|
||||
if len(opts.GetVideos()) > 0 {
|
||||
out = append(out, "videos")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// rawJSON passes a JSON string through untouched, or omits it when it is
|
||||
// empty or invalid rather than failing the whole request.
|
||||
func rawJSON(s string) json.RawMessage {
|
||||
@@ -207,6 +238,13 @@ func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<-
|
||||
xlog.Debug("localai-proxy: skip malformed SSE frame", "path", path, "error", err)
|
||||
continue
|
||||
}
|
||||
if chunk.Error != nil {
|
||||
// The upstream failed mid-generation. Returning nil would turn a
|
||||
// cut-off answer into a success; Unavailable lets failover and
|
||||
// the client see the failure.
|
||||
xlog.Warn("localai-proxy: upstream stream error", "path", path, "error", chunk.Error.Message)
|
||||
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", path, chunk.Error.Message)
|
||||
}
|
||||
if chunk.Usage != nil && len(chunk.Choices) == 0 {
|
||||
results <- &pb.Reply{PromptTokens: chunk.Usage.PromptTokens, Tokens: chunk.Usage.CompletionTokens}
|
||||
continue
|
||||
@@ -273,7 +311,11 @@ func (p *LocalAIProxy) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.Re
|
||||
"model": p.model(""),
|
||||
"query": in.GetQuery(),
|
||||
"documents": in.GetDocuments(),
|
||||
"top_n": in.GetTopN(),
|
||||
}
|
||||
// TopN 0 means "score every document" (the router's reranker sends it);
|
||||
// upstream rejects top_n < 1, and an absent top_n means the same thing.
|
||||
if n := in.GetTopN(); n > 0 {
|
||||
body["top_n"] = n
|
||||
}
|
||||
var resp struct {
|
||||
Usage struct {
|
||||
|
||||
@@ -176,6 +176,27 @@ var _ = Describe("localai-proxy", func() {
|
||||
Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700))
|
||||
})
|
||||
|
||||
It("maps a 429 upstream to ResourceExhausted so failover skips without tripping", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/completions", scriptedResponse{Status: http.StatusTooManyRequests, Body: "slow down"})
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(codeOf(err)).To(Equal(codes.ResourceExhausted))
|
||||
Expect(err.Error()).To(ContainSubstring("slow down"))
|
||||
})
|
||||
|
||||
It("always sends temperature, even 0, so greedy decoding survives", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "t"}}})
|
||||
|
||||
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
req := up.last()
|
||||
Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0)))
|
||||
Expect(req.JSON).NotTo(HaveKey("top_p"))
|
||||
Expect(req.JSON).NotTo(HaveKey("top_k"))
|
||||
})
|
||||
|
||||
It("maps an unreachable upstream to Unavailable", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.Close()
|
||||
@@ -229,6 +250,21 @@ var _ = Describe("localai-proxy", func() {
|
||||
Expect(string((<-results).GetMessage())).To(Equal("b"))
|
||||
})
|
||||
|
||||
It("returns a mid-stream upstream error frame as Unavailable", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "partial"}}}}),
|
||||
sseJSON(map[string]any{"error": map[string]any{"message": "backend crashed", "type": "server_error", "code": "server_error"}}),
|
||||
"[DONE]",
|
||||
}})
|
||||
|
||||
results := make(chan *pb.Reply, 10)
|
||||
err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Expect(err.Error()).To(ContainSubstring("backend crashed"))
|
||||
Expect(results).To(HaveLen(1))
|
||||
})
|
||||
|
||||
It("maps a failing upstream to a gRPC code", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/completions", scriptedResponse{Status: http.StatusBadGateway, Body: "gateway"})
|
||||
@@ -260,6 +296,13 @@ var _ = Describe("localai-proxy", func() {
|
||||
})
|
||||
})
|
||||
|
||||
It("names the request fields it cannot forward", func() {
|
||||
Expect(unforwardedFields(&pb.PredictOptions{Prompt: "x"})).To(BeEmpty())
|
||||
Expect(unforwardedFields(&pb.PredictOptions{
|
||||
Grammar: "root ::= x", Images: []string{"i"}, Audios: []string{"a"}, Videos: []string{"v"},
|
||||
})).To(Equal([]string{"grammar", "images", "audios", "videos"}))
|
||||
})
|
||||
|
||||
Describe("Embeddings", func() {
|
||||
It("posts the input to /v1/embeddings and returns the first vector", func() {
|
||||
p := loadProxy(up, nil)
|
||||
@@ -314,6 +357,15 @@ var _ = Describe("localai-proxy", func() {
|
||||
})
|
||||
})
|
||||
|
||||
It("Rerank omits top_n when it is 0, which means score every document", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/rerank", map[string]any{"results": []any{}})
|
||||
|
||||
_, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(up.last().JSON).NotTo(HaveKey("top_n"))
|
||||
})
|
||||
|
||||
Describe("TokenizeString and Detokenize", func() {
|
||||
It("posts the prompt to /v1/tokenize", func() {
|
||||
p := loadProxy(up, nil)
|
||||
|
||||
@@ -553,6 +553,14 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions {
|
||||
RequestTimeoutSeconds: int32(c.Proxy.RequestTimeoutSeconds),
|
||||
CachePrompt: c.Proxy.CachePrompt,
|
||||
}
|
||||
// localai-proxy calls a LocalAI that knows the model by name, so an
|
||||
// unset upstream_model means this config's name, the same derivation
|
||||
// failover.UpstreamModel uses. Not for cloud-proxy: its translate mode
|
||||
// falls back to parameters.model and passthrough keeps the client's
|
||||
// model when upstream_model is empty.
|
||||
if c.Backend == "localai-proxy" && opts.Proxy.UpstreamModel == "" {
|
||||
opts.Proxy.UpstreamModel = c.Name
|
||||
}
|
||||
}
|
||||
|
||||
if c.MMProj != "" {
|
||||
|
||||
@@ -62,6 +62,33 @@ var _ = Describe("grpcModelOpts Proxy options", func() {
|
||||
Expect(opts.Proxy.Mode).To(Equal(config.ProxyModePassthrough))
|
||||
})
|
||||
|
||||
It("sends the config name as the localai-proxy upstream model when none is set", func() {
|
||||
threads := 1
|
||||
cfg := config.ModelConfig{
|
||||
Name: "argus-whisper",
|
||||
Threads: &threads,
|
||||
Backend: "localai-proxy",
|
||||
Proxy: config.ProxyConfig{UpstreamURL: "http://127.0.0.1:8081"},
|
||||
}
|
||||
cfg.Model = "some-file"
|
||||
|
||||
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("argus-whisper"))
|
||||
|
||||
cfg.Proxy.UpstreamModel = "whisper-large"
|
||||
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("whisper-large"))
|
||||
})
|
||||
|
||||
It("leaves the cloud-proxy upstream model unset so translate mode keeps its fallback", func() {
|
||||
threads := 1
|
||||
cfg := config.ModelConfig{
|
||||
Name: "claude-strict",
|
||||
Threads: &threads,
|
||||
Backend: "cloud-proxy",
|
||||
Proxy: config.ProxyConfig{UpstreamURL: "https://api.example.com", Mode: config.ProxyModeTranslate},
|
||||
}
|
||||
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("leaves Proxy nil for a backend that is not a proxy", func() {
|
||||
threads := 1
|
||||
opts := grpcModelOpts(config.ModelConfig{Threads: &threads, Backend: "llama-cpp"}, "/tmp/models")
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
package config
|
||||
|
||||
func init() {
|
||||
RegisterBackendHook("localai-proxy", localAIProxyDefaults)
|
||||
}
|
||||
|
||||
// localAIProxyDefaults makes chat requests reach the upstream as structured
|
||||
// messages. Without the tokenizer template core renders the prompt itself
|
||||
// with no template, and the proxy can only send that text to
|
||||
// /v1/completions, bypassing the upstream model's chat template, tool
|
||||
// handling and reasoning parsing.
|
||||
//
|
||||
// Only configs that declare the chat usecase get it: usecase guessing reads
|
||||
// the tokenizer template as "this model chats", so setting it on a
|
||||
// transcription or embedding proxy would offer that model to chat pickers
|
||||
// and default-model selection. A config that brings its own templates keeps
|
||||
// them: the operator chose local templating.
|
||||
func localAIProxyDefaults(cfg *ModelConfig, _ string) {
|
||||
t := cfg.TemplateConfig
|
||||
if t.UseTokenizerTemplate || t.Chat != "" || t.ChatMessage != "" || t.Completion != "" || t.Edit != "" {
|
||||
return
|
||||
}
|
||||
declared := GetUsecasesFromYAML(cfg.KnownUsecaseStrings)
|
||||
if declared == nil || *declared&FLAG_CHAT != FLAG_CHAT {
|
||||
return
|
||||
}
|
||||
cfg.TemplateConfig.UseTokenizerTemplate = true
|
||||
}
|
||||
@@ -127,6 +127,38 @@ var _ = Describe("Backend hooks and parser defaults", func() {
|
||||
})
|
||||
})
|
||||
|
||||
Context("localai-proxy hook", func() {
|
||||
It("defaults a chat proxy to the tokenizer template so chat sends messages upstream", func() {
|
||||
cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}}
|
||||
cfg.SetDefaults()
|
||||
Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeTrue())
|
||||
})
|
||||
|
||||
It("leaves non-chat and undeclared proxies alone so they are not guessed as chat", func() {
|
||||
cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"transcript"}}
|
||||
cfg.SetDefaults()
|
||||
Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse())
|
||||
Expect(cfg.HasUsecases(FLAG_CHAT)).To(BeFalse())
|
||||
|
||||
bare := &ModelConfig{Backend: "localai-proxy"}
|
||||
bare.SetDefaults()
|
||||
Expect(bare.TemplateConfig.UseTokenizerTemplate).To(BeFalse())
|
||||
})
|
||||
|
||||
It("keeps a config that brings its own templates", func() {
|
||||
cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}}
|
||||
cfg.TemplateConfig.Chat = "{{.Input}}"
|
||||
cfg.SetDefaults()
|
||||
Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse())
|
||||
})
|
||||
|
||||
It("does not touch other backends", func() {
|
||||
cfg := &ModelConfig{Backend: "cloud-proxy", KnownUsecaseStrings: []string{"chat"}}
|
||||
cfg.SetDefaults()
|
||||
Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
Context("vllmDefaults hook", func() {
|
||||
It("auto-sets parsers for known model families on vllm backend", func() {
|
||||
cfg := &ModelConfig{
|
||||
|
||||
@@ -293,6 +293,18 @@ var _ = Describe("failover chains in the request pipeline", func() {
|
||||
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
|
||||
})
|
||||
|
||||
It("spills a rate-limited target (gRPC ResourceExhausted) to the next target without tripping it", func() {
|
||||
behavior["a"] = func(echo.Context) error {
|
||||
return grpcstatus.Error(codes.ResourceExhausted, "localai-proxy: upstream /v1/chat/completions returned 429: slow down")
|
||||
}
|
||||
rec := chat("chain")
|
||||
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
|
||||
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
|
||||
Expect(calls).To(Equal([]string{"a", "b"}))
|
||||
st, _ := fm.ChainStatus("chain")
|
||||
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
|
||||
})
|
||||
|
||||
It("skips a disabled target without tripping it", func() {
|
||||
rec := chat("chain-off")
|
||||
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
|
||||
|
||||
@@ -61,16 +61,25 @@ func IsRetryable(err error, status int) bool {
|
||||
return !isRequestError(msg)
|
||||
}
|
||||
|
||||
// IsCapabilityGap reports a target that cannot serve this kind of request at
|
||||
// all (gRPC Unimplemented, anywhere in the error chain). The next target may
|
||||
// serve it, and this target is not broken: the failure carries no signal
|
||||
// about its health, so callers must skip it without tripping.
|
||||
// IsCapabilityGap reports a target that cannot serve this request right now
|
||||
// for a reason that says nothing about its health: it cannot serve this kind
|
||||
// of request at all (gRPC Unimplemented), or it is out of capacity, such as a
|
||||
// rate-limited upstream (gRPC ResourceExhausted, what localai-proxy returns
|
||||
// for an upstream 429). Matched anywhere in the error chain. The next target
|
||||
// may serve it, so callers must skip this one without tripping it.
|
||||
func IsCapabilityGap(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
return ok && st.Code() == codes.Unimplemented
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
switch st.Code() {
|
||||
case codes.Unimplemented, codes.ResourceExhausted:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func retryableStatus(code int) bool {
|
||||
|
||||
@@ -45,6 +45,8 @@ var _ = DescribeTable("IsCapabilityGap",
|
||||
Entry("nil error", nil, false),
|
||||
Entry("grpc unimplemented", grpcstatus.Error(codes.Unimplemented, "x"), true),
|
||||
Entry("wrapped grpc unimplemented", fmt.Errorf("call: %w", grpcstatus.Error(codes.Unimplemented, "x")), true),
|
||||
Entry("grpc resource exhausted", grpcstatus.Error(codes.ResourceExhausted, "x"), true),
|
||||
Entry("wrapped grpc resource exhausted", fmt.Errorf("call: %w", grpcstatus.Error(codes.ResourceExhausted, "x")), true),
|
||||
Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), false),
|
||||
Entry("plain error", errors.New("boom"), false),
|
||||
)
|
||||
@@ -728,8 +728,9 @@ func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Cont
|
||||
att.Succeed()
|
||||
return nil
|
||||
case IsCapabilityGap(err) && !committed.Load():
|
||||
// This target cannot serve this kind of request at all; it is
|
||||
// not broken, so move on without counting a failure.
|
||||
// This target cannot serve this kind of request, or is out of
|
||||
// capacity; it is not broken, so move on without counting a
|
||||
// failure.
|
||||
if !att.Skip() {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -305,5 +305,20 @@ var _ = Describe("Manager", func() {
|
||||
st, _ := m.ChainStatus("chain")
|
||||
Expect(st.Targets[0].State).To(Equal(StateHealthy))
|
||||
})
|
||||
|
||||
It("skips a rate-limited (ResourceExhausted) target without tripping it", func() {
|
||||
var tried []string
|
||||
err := m.Do(context.Background(), "chain", func(_ context.Context, target string, _ func()) error {
|
||||
tried = append(tried, target)
|
||||
if target == "a" {
|
||||
return grpcstatus.Error(codes.ResourceExhausted, "localai-proxy: upstream /v1/rerank returned 429: slow down")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(tried).To(Equal([]string{"a", "b"}))
|
||||
st, _ := m.ChainStatus("chain")
|
||||
Expect(st.Targets[0].State).To(Equal(StateHealthy))
|
||||
})
|
||||
})
|
||||
})
|
||||
Reference in new issue
Block a user