From 10f7b1ca6055df046b36f85a0e4f911130978818 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 20:16:02 +0000 Subject: [PATCH] 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 Assisted-by: Claude:claude-opus-5-5 [Claude Code] --- backend/go/localai-proxy/main.go | 4 ++ backend/go/localai-proxy/text.go | 17 +++++- backend/go/localai-proxy/text_test.go | 38 ++++++++++++++ pkg/grpc/interface.go | 10 ++++ pkg/grpc/rich_context_test.go | 75 +++++++++++++++++++++++++++ pkg/grpc/server.go | 10 +++- 6 files changed, 151 insertions(+), 3 deletions(-) create mode 100644 pkg/grpc/rich_context_test.go diff --git a/backend/go/localai-proxy/main.go b/backend/go/localai-proxy/main.go index 7c3c3f0e8..93e5534fa 100644 --- a/backend/go/localai-proxy/main.go +++ b/backend/go/localai-proxy/main.go @@ -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) diff --git a/backend/go/localai-proxy/text.go b/backend/go/localai-proxy/text.go index 34bc98411..531718267 100644 --- a/backend/go/localai-proxy/text.go +++ b/backend/go/localai-proxy/text.go @@ -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 } diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go index e6e2f787a..c8e8fad0d 100644 --- a/backend/go/localai-proxy/text_test.go +++ b/backend/go/localai-proxy/text_test.go @@ -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{ diff --git a/pkg/grpc/interface.go b/pkg/grpc/interface.go index 5287e18a4..c4d6d4967 100644 --- a/pkg/grpc/interface.go +++ b/pkg/grpc/interface.go @@ -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 diff --git a/pkg/grpc/rich_context_test.go b/pkg/grpc/rich_context_test.go new file mode 100644 index 000000000..b7dee0bdf --- /dev/null +++ b/pkg/grpc/rich_context_test.go @@ -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) + }) +}) diff --git a/pkg/grpc/server.go b/pkg/grpc/server.go index ec3021ebb..4608ae748 100644 --- a/pkg/grpc/server.go +++ b/pkg/grpc/server.go @@ -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