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 ```