fix(localai-proxy): cancel the upstream request when the caller goes away

The proxy sent every chat and completion request with
context.Background, and the gRPC server gave rich backends no context
at all. When a client disconnected, or failover gave up on the target,
the upstream kept generating to the end, which costs tokens on a paid
or shared upstream. A silent upstream held the backend forever.

Add the optional AIModelRichContext interface. The gRPC server prefers
it and passes the call's context, like the Score and Rerank
extensions. The proxy implements it, so the upstream request ends with
the gRPC call.

Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
Assisted-by: Claude:claude-opus-5-5 [Claude Code]
This commit is contained in:
Ettore Di Giacinto committed 2026-09-27 20:16:02 +00:00
1 parent 19c1717bdb
commit 10f7b1ca60
6 files changed
+151 -3

No files matched your search

+4
View File
@@ -30,3 +30,7 @@ func main() {
panic(err)
}
}
// The gRPC server only passes the call's context to backends that implement
// this, so keep the proxy on it.
var _ grpc.AIModelRichContext = (*LocalAIProxy)(nil)
+15 -2
View File
@@ -203,9 +203,15 @@ func replyFromChoice(c textChoice, streaming bool) *pb.Reply {
}
func (p *LocalAIProxy) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) {
return p.PredictRichContext(context.Background(), opts)
}
// PredictRichContext is PredictRich bound to the gRPC call: when the caller
// goes away, the upstream request is cancelled too.
func (p *LocalAIProxy) PredictRichContext(ctx context.Context, opts *pb.PredictOptions) (*pb.Reply, error) {
path, body := p.textRequest(opts, false)
var resp textResponse
if err := p.postJSON(context.Background(), path, body, &resp); err != nil {
if err := p.postJSON(ctx, path, body, &resp); err != nil {
return nil, err
}
if len(resp.Choices) == 0 {
@@ -222,8 +228,15 @@ func (p *LocalAIProxy) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) {
// PredictStreamRich sends one Reply per upstream SSE delta. It does not close
// results: the gRPC server does, after this returns.
func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<- *pb.Reply) error {
return p.PredictStreamRichContext(context.Background(), opts, results)
}
// PredictStreamRichContext is PredictStreamRich bound to the gRPC stream: a
// client that disconnects, or a failover that abandons this target, stops the
// upstream generation instead of letting it run (and bill) to the end.
func (p *LocalAIProxy) PredictStreamRichContext(ctx context.Context, opts *pb.PredictOptions, results chan<- *pb.Reply) error {
path, body := p.textRequest(opts, true)
resp, err := p.postStream(context.Background(), path, body)
resp, err := p.postStream(ctx, path, body)
if err != nil {
return err
}
+38
View File
@@ -285,6 +285,44 @@ var _ = Describe("localai-proxy", func() {
Expect(string((<-results).GetMessage())).To(Equal("b"))
})
It("stops the upstream request when the caller cancels", func() {
upstreamGone := make(chan struct{})
slow := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = w.Write([]byte("data: " + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "a"}}}}) + "\n\n"))
w.(http.Flusher).Flush()
// A generation that outlasts the spec unless the proxy hangs up.
select {
case <-r.Context().Done():
close(upstreamGone)
case <-time.After(8 * time.Second):
}
})
DeferCleanup(slow.Close)
p := loadProxy(slow, nil)
addr := "test://localai-proxy-cancel"
grpc.Provide(addr, p)
client := grpc.NewClient(addr, true, nil, false)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
errCh := make(chan error, 1)
first := make(chan struct{}, 1)
go func() {
errCh <- client.PredictStream(ctx, &pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, func(*pb.Reply) {
select {
case first <- struct{}{}:
default:
}
})
}()
Eventually(first, 5*time.Second).Should(Receive())
cancel()
Eventually(upstreamGone, 5*time.Second).Should(BeClosed(), "the upstream generation must stop with the caller")
Eventually(errCh, 5*time.Second).Should(Receive(HaveOccurred()))
})
It("returns a mid-stream upstream error frame as Unavailable", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
+10
View File
@@ -104,6 +104,16 @@ type AIModelRich interface {
PredictStreamRich(*pb.PredictOptions, chan<- *pb.Reply) error
}
// AIModelRichContext is an optional extension to AIModelRich for backends
// whose work outlives a plain function call, such as a proxy waiting on a
// remote server. The gRPC server prefers it and passes the call's context, so
// a caller that disconnects or gives up stops the work instead of letting it
// run to the end. The channel contract is the same as PredictStreamRich.
type AIModelRichContext interface {
PredictRichContext(context.Context, *pb.PredictOptions) (*pb.Reply, error)
PredictStreamRichContext(context.Context, *pb.PredictOptions, chan<- *pb.Reply) error
}
// ClassifyModel is an optional extension to AIModel for backends that
// implement the TokenClassify RPC (zero-shot NER). The gRPC server
// type-asserts to this interface; backends that do not implement it
+75
View File
@@ -0,0 +1,75 @@
package grpc
import (
"context"
"time"
"github.com/mudler/LocalAI/pkg/grpc/base"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
// ctxBackend reports the context each rich call receives, so a spec can check
// that cancelling the caller reaches the backend.
type ctxBackend struct {
base.SingleThread
seen chan context.Context
}
func (b *ctxBackend) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) {
return b.PredictRichContext(context.Background(), opts)
}
func (b *ctxBackend) PredictStreamRich(opts *pb.PredictOptions, out chan<- *pb.Reply) error {
return b.PredictStreamRichContext(context.Background(), opts, out)
}
func (b *ctxBackend) PredictRichContext(ctx context.Context, _ *pb.PredictOptions) (*pb.Reply, error) {
b.seen <- ctx
<-ctx.Done()
return nil, ctx.Err()
}
func (b *ctxBackend) PredictStreamRichContext(ctx context.Context, _ *pb.PredictOptions, _ chan<- *pb.Reply) error {
b.seen <- ctx
<-ctx.Done()
return ctx.Err()
}
var _ AIModelRichContext = (*ctxBackend)(nil)
var _ = Describe("AIModelRichContext dispatch", func() {
// A backend that only saw context.Background would block forever here,
// so bound each call and check the backend returned because of the cancel.
expectCancelReaches := func(call func(ctx context.Context) error, seen chan context.Context) {
ctx, cancel := context.WithCancel(context.Background())
errCh := make(chan error, 1)
go func() { errCh <- call(ctx) }()
var got context.Context
Eventually(seen).Should(Receive(&got))
cancel()
Eventually(got.Done(), 2*time.Second).Should(BeClosed())
Eventually(errCh, 2*time.Second).Should(Receive(HaveOccurred()))
}
It("passes the caller's context to PredictStreamRichContext", func() {
b := &ctxBackend{seen: make(chan context.Context, 1)}
Provide("test://rich-ctx-stream", b)
c := NewClient("test://rich-ctx-stream", true, nil, false)
expectCancelReaches(func(ctx context.Context) error {
return c.PredictStream(ctx, &pb.PredictOptions{}, func(*pb.Reply) {})
}, b.seen)
})
It("passes the caller's context to PredictRichContext", func() {
b := &ctxBackend{seen: make(chan context.Context, 1)}
Provide("test://rich-ctx-predict", b)
c := NewClient("test://rich-ctx-predict", true, nil, false)
expectCancelReaches(func(ctx context.Context) error {
_, err := c.Predict(ctx, &pb.PredictOptions{})
return err
}, b.seen)
})
})
+9 -1
View File
@@ -176,6 +176,9 @@ func (s *server) Predict(ctx context.Context, in *pb.PredictOptions) (*pb.Reply,
s.llm.Lock()
defer s.llm.Unlock()
}
if rich, ok := s.llm.(AIModelRichContext); ok {
return rich.PredictRichContext(ctx, in)
}
if rich, ok := s.llm.(AIModelRich); ok {
return rich.PredictRich(in)
}
@@ -580,7 +583,12 @@ func (s *server) PredictStream(in *pb.PredictOptions, stream pb.Backend_PredictS
// Server-side close: PredictStreamRich implementations send into
// the channel and return when finished; closing is the host's
// concern so impls don't have to remember `defer close(...)`.
err := rich.PredictStreamRich(in, replyChan)
var err error
if withCtx, ok := s.llm.(AIModelRichContext); ok {
err = withCtx.PredictStreamRichContext(stream.Context(), in, replyChan)
} else {
err = rich.PredictStreamRich(in, replyChan)
}
close(replyChan)
<-done
return err