package model import ( "context" "sync" grpc "github.com/mudler/LocalAI/pkg/grpc" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "github.com/mudler/xlog" ggrpc "google.golang.org/grpc" ) // ConnectionEvictingClient wraps a grpc.Backend. When any inference method // fails with a connection error (server unreachable), it calls the evict // callback to remove the model from the ModelLoader's cache. The error is // still returned to the caller — the NEXT request will trigger rescheduling // via SmartRouter. type ConnectionEvictingClient struct { grpc.WrappedBackend modelID string evict func() once sync.Once } var _ grpc.BackendUnwrapper = (*ConnectionEvictingClient)(nil) func newConnectionEvictingClient(inner grpc.Backend, modelID string, evict func()) grpc.Backend { return &ConnectionEvictingClient{ WrappedBackend: grpc.WrappedBackend{Backend: inner}, modelID: modelID, evict: evict, } } func (c *ConnectionEvictingClient) checkErr(err error) { if err == nil || !isConnectionError(err) { return } // The fifth site of the same shape, and the one reached during INFERENCE // rather than a health check. evict() runs ShutdownModel, which for a remote // model sends a backend.stop control RPC over the tunnel of every node // holding it and deletes every replica row. In distributed mode the client underneath reaches the // backend over the worker's tunnel, and a failure of THAT transport arrives // as the same codes.Unavailable a dead backend produces; evicting on it // stops a model that is loaded and serving, on a worker that is // heartbeating. A locally spawned backend has no custom transport, so this // reports nil and the behaviour there is exactly what it always was. // transportFailure and not LastDialErrorOf: a refusal the WORKER wrote is // the worker answering that it could not reach the process, which is what a // crashed backend produces now that a worker listens on nothing. Treating // that as a transport failure kept a genuinely dead model loaded and // failing every request, which is the mirror image of the mistake this // guard exists to prevent. if dialErr := transportFailure(c.Backend); dialErr != nil { xlog.Warn("Inference failed because the worker could not be reached; keeping the model", "model", c.modelID, "error", dialErr) return } c.once.Do(func() { xlog.Warn("Connection error during inference, evicting model from cache", "model", c.modelID, "error", err) c.evict() }) } // --- Intercepted inference methods --- func (c *ConnectionEvictingClient) Predict(ctx context.Context, in *pb.PredictOptions, opts ...ggrpc.CallOption) (*pb.Reply, error) { reply, err := c.Backend.Predict(ctx, in, opts...) c.checkErr(err) return reply, err } func (c *ConnectionEvictingClient) PredictStream(ctx context.Context, in *pb.PredictOptions, f func(reply *pb.Reply), opts ...ggrpc.CallOption) error { err := c.Backend.PredictStream(ctx, in, f, opts...) c.checkErr(err) return err } func (c *ConnectionEvictingClient) Embeddings(ctx context.Context, in *pb.PredictOptions, opts ...ggrpc.CallOption) (*pb.EmbeddingResult, error) { result, err := c.Backend.Embeddings(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.GenerateImage(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.GenerateVideo(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) Animate3D(ctx context.Context, in *pb.Animate3DRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.Animate3D(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) Generate3D(ctx context.Context, in *pb.Generate3DRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.Generate3D(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) TTS(ctx context.Context, in *pb.TTSRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.TTS(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) TTSStream(ctx context.Context, in *pb.TTSRequest, f func(reply *pb.Reply), opts ...ggrpc.CallOption) error { err := c.Backend.TTSStream(ctx, in, f, opts...) c.checkErr(err) return err } func (c *ConnectionEvictingClient) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequest, opts ...ggrpc.CallOption) (*pb.Result, error) { result, err := c.Backend.SoundGeneration(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest, opts ...ggrpc.CallOption) (*pb.TranscriptResult, error) { result, err := c.Backend.AudioTranscription(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) AudioTranscriptionStream(ctx context.Context, in *pb.TranscriptRequest, f func(chunk *pb.TranscriptStreamResponse), opts ...ggrpc.CallOption) error { err := c.Backend.AudioTranscriptionStream(ctx, in, f, opts...) c.checkErr(err) return err } func (c *ConnectionEvictingClient) Detect(ctx context.Context, in *pb.DetectOptions, opts ...ggrpc.CallOption) (*pb.DetectResponse, error) { result, err := c.Backend.Detect(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) Depth(ctx context.Context, in *pb.DepthRequest, opts ...ggrpc.CallOption) (*pb.DepthResponse, error) { result, err := c.Backend.Depth(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) Rerank(ctx context.Context, in *pb.RerankRequest, opts ...ggrpc.CallOption) (*pb.RerankResult, error) { result, err := c.Backend.Rerank(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) TokenClassify(ctx context.Context, in *pb.TokenClassifyRequest, opts ...ggrpc.CallOption) (*pb.TokenClassifyResponse, error) { result, err := c.Backend.TokenClassify(ctx, in, opts...) c.checkErr(err) return result, err } func (c *ConnectionEvictingClient) Score(ctx context.Context, in *pb.ScoreRequest, opts ...ggrpc.CallOption) (*pb.ScoreResponse, error) { result, err := c.Backend.Score(ctx, in, opts...) c.checkErr(err) return result, err }