mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-30 10:04:32 -04:00
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]
76 lines
2.4 KiB
Go
76 lines
2.4 KiB
Go
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)
|
|
})
|
|
})
|