mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-30 01:54:31 -04:00
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:
1 parent
19c1717bdb
commit
10f7b1ca60
6 files changed
+151
-3
No files matched your search
@@ -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)
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in new issue
Block a user