From d3724eef509acdc63ccb882e65b612b1083a9479 Mon Sep 17 00:00:00 2001 From: mudler-agent Date: Sun, 4 Oct 2026 13:14:00 +0200 Subject: [PATCH] fix(nodes): stage dedicated diarization audio (#12474) Dedicated diarization forwards frontend audio paths to remote workers, which cannot read those temporary files. Stage the input before the RPC and release it afterward, following the transcription lifecycle. Clone the request so staging does not change caller-owned data. Cover input bytes, request fields, cleanup, and error propagation in tests. Document distributed diarization staging on the existing feature page. Assisted-by: nib:gpt-6-astra Signed-off-by: Ettore Di Giacinto Co-authored-by: Ettore Di Giacinto --- core/services/nodes/file_staging_client.go | 17 +++ .../nodes/file_staging_diarization_test.go | 113 ++++++++++++++++++ docs/content/features/audio-diarization.md | 4 + 3 files changed, 134 insertions(+) create mode 100644 core/services/nodes/file_staging_diarization_test.go diff --git a/core/services/nodes/file_staging_client.go b/core/services/nodes/file_staging_client.go index 7fa91d716..e03671a15 100644 --- a/core/services/nodes/file_staging_client.go +++ b/core/services/nodes/file_staging_client.go @@ -506,6 +506,23 @@ func (f *FileStagingClient) SoundDetection(ctx context.Context, in *pb.SoundDete return f.Backend.SoundDetection(ctx, in, opts...) } +func (f *FileStagingClient) Diarize(ctx context.Context, in *pb.DiarizeRequest, opts ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + lifecycle := f.newStagedInputLifecycle() + defer lifecycle.release() + in = proto.Clone(in).(*pb.DiarizeRequest) + + // Stage input audio file + if in.Dst != "" && isFilePath(in.Dst) { + backendPath, err := f.stageInputFile(ctx, lifecycle, in.Dst, "inputs") + if err != nil { + return nil, fmt.Errorf("staging audio for diarization: %w", err) + } + in.Dst = backendPath + } + + return f.Backend.Diarize(ctx, in, opts...) +} + func (f *FileStagingClient) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest, opts ...ggrpc.CallOption) (*pb.TranscriptResult, error) { lifecycle := f.newStagedInputLifecycle() defer lifecycle.release() diff --git a/core/services/nodes/file_staging_diarization_test.go b/core/services/nodes/file_staging_diarization_test.go new file mode 100644 index 000000000..6cc775ad6 --- /dev/null +++ b/core/services/nodes/file_staging_diarization_test.go @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: MIT + +package nodes + +import ( + "context" + "errors" + "os" + "path/filepath" + + grpc "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" + "google.golang.org/protobuf/proto" +) + +type diarizationStager struct { + lifecycleStager + root string +} + +func (s *diarizationStager) EnsureRemote(ctx context.Context, nodeID, localPath, key string) (string, error) { + if _, err := s.lifecycleStager.EnsureRemote(ctx, nodeID, localPath, key); err != nil { + return "", err + } + remotePath := filepath.Join(s.root, key) + return remotePath, copyFile(localPath, remotePath) +} + +type diarizationStagingBackend struct { + grpc.Backend + call func(context.Context, *pb.DiarizeRequest, ...ggrpc.CallOption) (*pb.DiarizeResponse, error) +} + +func (b *diarizationStagingBackend) Diarize(ctx context.Context, in *pb.DiarizeRequest, opts ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + return b.call(ctx, in, opts...) +} + +var _ = Describe("FileStagingClient diarization", func() { + DescribeTable("stages audio and releases it after the backend returns", func(backendErr error) { + ctx := context.WithValue(context.Background(), struct{}{}, "request-context") + root := GinkgoT().TempDir() + source := filepath.Join(root, "frontend", "clip.wav") + audio := []byte("RIFF\x00\x01\x02\xffWAVEtest audio") + Expect(os.MkdirAll(filepath.Dir(source), 0750)).To(Succeed()) + Expect(os.WriteFile(source, audio, 0600)).To(Succeed()) + stager := &diarizationStager{root: filepath.Join(root, "worker")} + request := &pb.DiarizeRequest{ + Dst: source, Threads: 4, Language: "en", NumSpeakers: 2, + MinSpeakers: 1, MaxSpeakers: 3, ClusteringThreshold: 0.7, + MinDurationOn: 0.2, MinDurationOff: 0.4, IncludeText: true, + ModelIdentity: "diarization-model", IncludeSpeakerProfiles: true, + KnownVoices: []*pb.KnownVoice{{Id: "voice-1", Name: "speaker", Model: "encoder", Embedding: []float32{0.1, 0.2}}}, + } + original := proto.Clone(request) + response := &pb.DiarizeResponse{NumSpeakers: 2} + option := ggrpc.WaitForReady(true) + calls := 0 + backend := &diarizationStagingBackend{call: func(gotCtx context.Context, in *pb.DiarizeRequest, opts ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + calls++ + Expect(gotCtx).To(Equal(ctx)) + Expect(opts).To(Equal([]ggrpc.CallOption{option})) + Expect(in).NotTo(BeIdenticalTo(request)) + Expect(stager.ensureCalls).To(HaveLen(1)) + staged := stager.ensureCalls[0] + Expect(staged.nodeID).To(Equal("worker-1")) + Expect(staged.localPath).To(Equal(source)) + Expect(in.Dst).To(Equal(filepath.Join(stager.root, staged.key))) + Expect(in.Dst).NotTo(Equal(source)) + contents, err := os.ReadFile(in.Dst) + Expect(err).NotTo(HaveOccurred()) + Expect(contents).To(Equal(audio)) + expected := proto.Clone(request).(*pb.DiarizeRequest) + expected.Dst = in.Dst + Expect(proto.Equal(in, expected)).To(BeTrue()) + Expect(stager.releasedKeys).To(BeEmpty()) + in.KnownVoices[0].Embedding[0] = 99 + return response, backendErr + }} + result, err := NewFileStagingClient(backend, stager, "worker-1").Diarize(ctx, request, option) + if backendErr != nil { + Expect(err).To(BeIdenticalTo(backendErr)) + } else { + Expect(err).NotTo(HaveOccurred()) + } + Expect(result).To(BeIdenticalTo(response)) + Expect(calls).To(Equal(1)) + Expect(proto.Equal(request, original)).To(BeTrue()) + Expect(stager.releasedKeys).To(Equal([]string{stager.ensureCalls[0].key})) + }, Entry("success", nil), Entry("backend error", errors.New("diarization failed"))) + + It("releases the attempted stage and does not forward when staging fails", func(ctx SpecContext) { + failure := errors.New("upload failed") + stager := &lifecycleStager{ensureErr: failure} + called := false + backend := &diarizationStagingBackend{call: func(context.Context, *pb.DiarizeRequest, ...ggrpc.CallOption) (*pb.DiarizeResponse, error) { + called = true + return &pb.DiarizeResponse{}, nil + }} + request := &pb.DiarizeRequest{Dst: filepath.Join(GinkgoT().TempDir(), "clip.wav")} + original := proto.Clone(request) + result, err := NewFileStagingClient(backend, stager, "worker-1").Diarize(ctx, request) + Expect(err).To(MatchError(ContainSubstring("staging audio for diarization"))) + Expect(errors.Is(err, failure)).To(BeTrue()) + Expect(result).To(BeNil()) + Expect(called).To(BeFalse()) + Expect(proto.Equal(request, original)).To(BeTrue()) + Expect(stager.ensureCalls).To(HaveLen(1)) + Expect(stager.releasedKeys).To(Equal([]string{stager.ensureCalls[0].key})) + }) +}) diff --git a/docs/content/features/audio-diarization.md b/docs/content/features/audio-diarization.md index 53e5e487f..72bc0cdf1 100644 --- a/docs/content/features/audio-diarization.md +++ b/docs/content/features/audio-diarization.md @@ -19,6 +19,10 @@ LocalAI exposes this through the `/v1/audio/diarization` endpoint, modelled afte Because diarization is exposed as a regular OpenAI-compatible endpoint, any HTTP client works. There is no Python dependency on pyannote or NeMo on the consumer side. +In distributed mode, LocalAI stages uploaded audio on the remote worker before +running dedicated diarization and releases the staged input after the request. +The worker does not need access to the frontend’s temporary upload directory. + ## Endpoint ```