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