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 <mudler@localai.io>
Co-authored-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
mudler-agentandEttore Di Giacinto authored and GitHub committed 2026-10-04 13:14:00 +02:00
1 parent ed4a3975be
commit d3724eef50
3 files changed
+134

No files matched your search

@@ -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()
@@ -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}))
})
})
@@ -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
```