diff --git a/backend/go/parakeet-cpp/diarize.go b/backend/go/parakeet-cpp/diarize.go index 128c42007..c70afcd99 100644 --- a/backend/go/parakeet-cpp/diarize.go +++ b/backend/go/parakeet-cpp/diarize.go @@ -109,15 +109,20 @@ func unsupportedDiarizeFields(req *pb.DiarizeRequest) []string { // turn); otherwise, or when no ASR companion is loaded, segments carry no // text (parakeet_capi_diarize_pcm) and no error is raised. func (p *ParakeetCpp) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) { + // Check the diarization model first. A backend whose models were freed + // holds no diarization context either, and it must answer + // FailedPrecondition, which LocalAI reads as a stale replica and reloads. + // Unimplemented is final: it would hide the empty backend behind a 501 + // that every later request repeats. + if p.diarCtx == 0 { + return pb.DiarizeResponse{}, status.Error(codes.FailedPrecondition, + "parakeet-cpp: model is not a diarization model"+p.roleHint(componentDiar, "diar_component")) + } if req.GetIncludeSpeakerProfiles() { if CppDiarizeProfilesPCMJSON == nil || CppSpeakerIdentity == nil || CppSpeakerDim == nil || p.spkCtx == 0 { return pb.DiarizeResponse{}, status.Error(codes.Unimplemented, "parakeet-cpp: speaker profiles require a loaded speaker encoder and profile-capable library") } } - if p.diarCtx == 0 { - return pb.DiarizeResponse{}, status.Error(codes.FailedPrecondition, - "parakeet-cpp: model is not a diarization model"+p.roleHint(componentDiar, "diar_component")) - } if CppDiarizePCM == nil { return pb.DiarizeResponse{}, status.Error(codes.Unimplemented, "parakeet-cpp: loaded libparakeet.so has no diarization support (parakeet_capi_diarize_pcm missing)") diff --git a/backend/go/parakeet-cpp/profiles_test.go b/backend/go/parakeet-cpp/profiles_test.go index a44310361..33b184020 100644 --- a/backend/go/parakeet-cpp/profiles_test.go +++ b/backend/go/parakeet-cpp/profiles_test.go @@ -18,6 +18,16 @@ var _ = Describe("profile capability", func() { _, err := p.Diarize(&pb.DiarizeRequest{IncludeSpeakerProfiles: true}) Expect(status.Code(err)).To(Equal(codes.Unimplemented)) }) + + It("answers FailedPrecondition, not Unimplemented, when every model was freed", func() { + // A freed backend holds no contexts. The answer must let LocalAI + // drop the stale replica and reload it; Unimplemented never heals. + restore := diarizeStubs() + defer restore() + p := &ParakeetCpp{} + _, err := p.Diarize(&pb.DiarizeRequest{Dst: diarizeWav(1), IncludeSpeakerProfiles: true}) + Expect(status.Code(err)).To(Equal(codes.FailedPrecondition)) + }) }) var _ = Describe("profile export transport", func() { It("exports with no registry and copies trusted metadata without freeing borrowed identity", func() { diff --git a/core/services/worker/lifecycle.go b/core/services/worker/lifecycle.go index 86579c2fc..e6561b9da 100644 --- a/core/services/worker/lifecycle.go +++ b/core/services/worker/lifecycle.go @@ -367,25 +367,38 @@ func (s *backendSupervisor) backendList(_ context.Context, _ workerctl.BackendLi return workerctl.BackendListReply{Backends: infos} } +// unloadTargets returns the gRPC addresses a model.unload request must free. +// +// The address in the request wins when set. Otherwise the request names a +// model, and only that model's processes (every replica) are returned. A +// request that names no running model frees nothing: freeing some other +// model's process would empty its loaded weights while the control plane +// still counts it as loaded, and its next request would fail. +func (s *backendSupervisor) unloadTargets(req workerctl.ModelUnloadRequest) []string { + if req.Address != "" { + return []string{req.Address} + } + if req.ModelName == "" { + return nil + } + keys := s.resolveProcessKeys(req.ModelName) + s.mu.Lock() + defer s.mu.Unlock() + var addrs []string + for _, k := range keys { + if bp, ok := s.processes[k]; ok && bp.addr != "" { + addrs = append(addrs, bp.addr) + } + } + return addrs +} + // unloadModel answers model.unload: call gRPC Free() to release GPU memory // without killing the backend process. func (s *backendSupervisor) unloadModel(ctx context.Context, req workerctl.ModelUnloadRequest) workerctl.ModelUnloadReply { - xlog.Info("Received NATS model.unload event") + xlog.Info("Received NATS model.unload event", "model", req.ModelName) - // Find the backend address for this model's backend type - // The request includes an Address field if the router knows which process to target - targetAddr := req.Address - if targetAddr == "" { - // Fallback: try all running backends - s.mu.Lock() - for _, bp := range s.processes { - targetAddr = bp.addr - break - } - s.mu.Unlock() - } - - if targetAddr != "" { + for _, targetAddr := range s.unloadTargets(req) { // Best-effort bounded gRPC Free(). A model.unload request must not // occupy the NATS reply handler forever when a backend is wedged. client := grpc.NewClientWithToken(targetAddr, false, nil, false, s.cfg.RegistrationToken) diff --git a/core/services/worker/unload_targets_test.go b/core/services/worker/unload_targets_test.go new file mode 100644 index 000000000..d912b6c5e --- /dev/null +++ b/core/services/worker/unload_targets_test.go @@ -0,0 +1,47 @@ +package worker + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// model.unload frees GPU memory of one model. Freeing another model's +// process empties it while the control plane still counts it as loaded, so +// the target must come from the model name and never from map order. +var _ = Describe("backendSupervisor unload targets", func() { + newSupervisor := func() *backendSupervisor { + return &backendSupervisor{ + processes: map[string]*backendProcess{ + "chat#0": {addr: "127.0.0.1:50051"}, + "chat#1": {addr: "127.0.0.1:50052"}, + "parakeet#0": {addr: "127.0.0.1:50053"}, + }, + } + } + + It("targets every replica of the named model and nothing else", func() { + s := newSupervisor() + for range 50 { // map iteration order is random; the answer must not be + Expect(s.unloadTargets(workerctl.ModelUnloadRequest{ModelName: "chat"})). + To(ConsistOf("127.0.0.1:50051", "127.0.0.1:50052")) + } + }) + + It("targets the process named by an exact key", func() { + Expect(newSupervisor().unloadTargets(workerctl.ModelUnloadRequest{ModelName: "parakeet#0"})). + To(ConsistOf("127.0.0.1:50053")) + }) + + It("prefers the address in the request", func() { + Expect(newSupervisor().unloadTargets(workerctl.ModelUnloadRequest{ModelName: "chat", Address: "10.0.0.1:1"})). + To(ConsistOf("10.0.0.1:1")) + }) + + It("frees nothing for a model that is not running", func() { + s := newSupervisor() + Expect(s.unloadTargets(workerctl.ModelUnloadRequest{ModelName: "missing"})).To(BeEmpty()) + Expect(s.unloadTargets(workerctl.ModelUnloadRequest{})).To(BeEmpty()) + }) +})