diff --git a/backend/go/localai-proxy/audio.go b/backend/go/localai-proxy/audio.go
new file mode 100644
index 000000000..07e011209
--- /dev/null
+++ b/backend/go/localai-proxy/audio.go
@@ -0,0 +1,453 @@
+package main
+
+import (
+ "bufio"
+ "context"
+ "encoding/json"
+ "errors"
+ "io"
+ "net/url"
+ "path/filepath"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/mudler/xlog"
+ "google.golang.org/grpc/codes"
+ "google.golang.org/grpc/status"
+
+ pb "github.com/mudler/LocalAI/pkg/grpc/proto"
+)
+
+const transcriptionsPath = "/v1/audio/transcriptions"
+
+// ttsRequest is the body of LocalAI's /tts (schema.TTSRequest).
+type ttsRequest struct {
+ Model string `json:"model"`
+ Input string `json:"input"`
+ Voice string `json:"voice,omitempty"`
+ Language string `json:"language,omitempty"`
+ Instructions string `json:"instructions,omitempty"`
+ Params map[string]string `json:"params,omitempty"`
+ Stream bool `json:"stream,omitempty"`
+}
+
+func (p *LocalAIProxy) ttsRequest(req *pb.TTSRequest, stream bool) ttsRequest {
+ return ttsRequest{
+ Model: p.model(""),
+ Input: req.GetText(),
+ Voice: req.GetVoice(),
+ Language: req.GetLanguage(),
+ Instructions: req.GetInstructions(),
+ Params: req.GetParams(),
+ Stream: stream,
+ }
+}
+
+func (p *LocalAIProxy) TTS(req *pb.TTSRequest) error {
+ return p.postJSONToFile(context.Background(), "/tts", p.ttsRequest(req, false), req.GetDst())
+}
+
+// TTSStream forwards the upstream's chunked audio unchanged. That body is
+// already what core expects from a streaming backend: a WAV header followed
+// by PCM. out is closed on every path because the gRPC server drains it until
+// closed and would otherwise hang.
+func (p *LocalAIProxy) TTSStream(req *pb.TTSRequest, out chan []byte) error {
+ defer close(out)
+ resp, err := p.postStream(context.Background(), "/tts", p.ttsRequest(req, true))
+ if err != nil {
+ return err
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ buf := make([]byte, 32*1024)
+ for {
+ n, err := resp.Body.Read(buf)
+ if n > 0 {
+ // The reader reuses buf, so each chunk needs its own copy.
+ out <- append([]byte(nil), buf[:n]...)
+ }
+ if errors.Is(err, io.EOF) {
+ return nil
+ }
+ if err != nil {
+ // A cut-off stream is a failed synthesis, not a short one.
+ return transportError("/tts", err)
+ }
+ }
+}
+
+// soundGenerationRequest is the body of /v1/sound-generation
+// (schema.ElevenLabsSoundGenerationRequest). Pointers keep "unset" distinct
+// from zero so the upstream model's defaults apply.
+type soundGenerationRequest struct {
+ ModelID string `json:"model_id"`
+ Text string `json:"text"`
+ Duration *float32 `json:"duration_seconds,omitempty"`
+ Temperature *float32 `json:"prompt_influence,omitempty"`
+ DoSample *bool `json:"do_sample,omitempty"`
+ Think *bool `json:"think,omitempty"`
+ Caption string `json:"caption,omitempty"`
+ Lyrics string `json:"lyrics,omitempty"`
+ BPM *int32 `json:"bpm,omitempty"`
+ Keyscale string `json:"keyscale,omitempty"`
+ Language string `json:"language,omitempty"`
+ Timesignature string `json:"timesignature,omitempty"`
+ Instrumental *bool `json:"instrumental,omitempty"`
+}
+
+// SoundGeneration refuses audio-conditioned requests: the REST endpoint has no
+// field for a source clip, and dropping it would return unconditioned audio as
+// if it were the answer.
+func (p *LocalAIProxy) SoundGeneration(req *pb.SoundGenerationRequest) error {
+ if req.GetSrc() != "" {
+ return unimplemented("SoundGeneration with src")
+ }
+ body := soundGenerationRequest{
+ ModelID: p.model(""),
+ Text: req.GetText(),
+ Duration: req.Duration,
+ Temperature: req.Temperature,
+ DoSample: req.Sample,
+ Think: req.Think,
+ Caption: req.GetCaption(),
+ Lyrics: req.GetLyrics(),
+ BPM: req.Bpm,
+ Keyscale: req.GetKeyscale(),
+ Language: req.GetLanguage(),
+ Timesignature: req.GetTimesignature(),
+ Instrumental: req.Instrumental,
+ }
+ return p.postJSONToFile(context.Background(), "/v1/sound-generation", body, req.GetDst())
+}
+
+// transcriptionForm builds the /v1/audio/transcriptions upload. Dst is the
+// input audio (core names it that way). Language and translate are only sent
+// when set so the upstream model config's defaults still apply; diarize is
+// always sent because the upstream treats a missing field as true.
+func (p *LocalAIProxy) transcriptionForm(req *pb.TranscriptRequest, stream bool) multipartForm {
+ f := url.Values{}
+ f.Set("model", p.model(""))
+ f.Set("diarize", strconv.FormatBool(req.GetDiarize()))
+ // verbose_json keeps segments and words; the default drops nothing today,
+ // but naming it keeps the reply shape pinned.
+ f.Set("response_format", "verbose_json")
+ if v := req.GetLanguage(); v != "" {
+ f.Set("language", v)
+ }
+ if req.GetTranslate() {
+ f.Set("translate", "true")
+ }
+ if v := req.GetPrompt(); v != "" {
+ f.Set("prompt", v)
+ }
+ if v := req.GetTemperature(); v != 0 {
+ f.Set("temperature", formatFloat(v))
+ }
+ for _, g := range req.GetTimestampGranularities() {
+ f.Add("timestamp_granularities[]", g)
+ }
+ if stream {
+ f.Set("stream", "true")
+ }
+ return multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}
+}
+
+// transcriptionResult is TranscriptionResultSeconds: the REST API reports
+// times in seconds, pb in nanoseconds (core reads them as time.Duration).
+type transcriptionResult struct {
+ Text string `json:"text"`
+ Language string `json:"language"`
+ Duration float64 `json:"duration"`
+ Segments []transcriptionSegment `json:"segments"`
+}
+
+type transcriptionSegment struct {
+ ID int32 `json:"id"`
+ Start float64 `json:"start"`
+ End float64 `json:"end"`
+ Text string `json:"text"`
+ Tokens []int32 `json:"tokens"`
+ Speaker string `json:"speaker"`
+ Words []struct {
+ Start float64 `json:"start"`
+ End float64 `json:"end"`
+ Text string `json:"text"`
+ } `json:"words"`
+}
+
+// toProto drops the top-level words: core rebuilds them from segment words.
+func (r transcriptionResult) toProto() *pb.TranscriptResult {
+ out := &pb.TranscriptResult{Text: r.Text, Language: r.Language, Duration: float32(r.Duration)}
+ for _, s := range r.Segments {
+ seg := &pb.TranscriptSegment{
+ Id: s.ID, Start: nanos(s.Start), End: nanos(s.End), Text: s.Text,
+ Tokens: s.Tokens, Speaker: s.Speaker,
+ }
+ for _, w := range s.Words {
+ seg.Words = append(seg.Words, &pb.TranscriptWord{Start: nanos(w.Start), End: nanos(w.End), Text: w.Text})
+ }
+ out.Segments = append(out.Segments, seg)
+ }
+ return out
+}
+
+func nanos(seconds float64) int64 {
+ return int64(seconds * float64(time.Second))
+}
+
+func (p *LocalAIProxy) AudioTranscription(ctx context.Context, req *pb.TranscriptRequest) (pb.TranscriptResult, error) {
+ var resp transcriptionResult
+ if err := p.postForm(ctx, transcriptionsPath, p.transcriptionForm(req, false), &resp); err != nil {
+ return pb.TranscriptResult{}, err
+ }
+ return *resp.toProto(), nil
+}
+
+// transcriptEvent covers every frame of the upstream transcription SSE stream.
+type transcriptEvent struct {
+ Type string `json:"type"`
+ Delta string `json:"delta"`
+ Error *struct {
+ Message string `json:"message"`
+ } `json:"error"`
+ transcriptionResult
+}
+
+// AudioTranscriptionStream maps upstream SSE frames to stream responses. out
+// is closed on every path: the gRPC server drains it until closed.
+func (p *LocalAIProxy) AudioTranscriptionStream(ctx context.Context, req *pb.TranscriptRequest, out chan *pb.TranscriptStreamResponse) error {
+ defer close(out)
+ resp, err := p.postMultipartStream(ctx, transcriptionsPath, p.transcriptionForm(req, true))
+ if err != nil {
+ return err
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ scanner := bufio.NewScanner(resp.Body)
+ // The done frame carries every segment of the recording.
+ scanner.Buffer(make([]byte, 0, 64*1024), 4<<20)
+ for scanner.Scan() {
+ payload, ok := strings.CutPrefix(scanner.Text(), "data:")
+ if !ok {
+ continue
+ }
+ payload = strings.TrimSpace(payload)
+ if payload == "" || payload == "[DONE]" {
+ continue
+ }
+ var ev transcriptEvent
+ if err := json.Unmarshal([]byte(payload), &ev); err != nil {
+ xlog.Debug("localai-proxy: skip malformed SSE frame", "path", transcriptionsPath, "error", err)
+ continue
+ }
+ switch ev.Type {
+ case "transcript.text.delta":
+ out <- &pb.TranscriptStreamResponse{Delta: ev.Delta}
+ case "transcript.text.done":
+ out <- &pb.TranscriptStreamResponse{FinalResult: ev.toProto()}
+ return nil
+ case "error":
+ msg := "unknown error"
+ if ev.Error != nil {
+ msg = ev.Error.Message
+ }
+ xlog.Warn("localai-proxy: upstream stream error", "path", transcriptionsPath, "error", msg)
+ return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", transcriptionsPath, msg)
+ }
+ }
+ if err := scanner.Err(); err != nil {
+ return transportError(transcriptionsPath, err)
+ }
+ // The upstream always ends with a done or error frame, so a stream that
+ // stops without one was cut off; returning nil would pass the partial
+ // deltas off as the whole transcript.
+ return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream ended before the final transcript", transcriptionsPath)
+}
+
+func (p *LocalAIProxy) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) {
+ f := url.Values{}
+ f.Set("model", p.model(""))
+ // verbose_json keeps per-segment text; core strips it again for callers
+ // that did not ask for it.
+ f.Set("response_format", "verbose_json")
+ if v := req.GetLanguage(); v != "" {
+ f.Set("language", v)
+ }
+ for name, v := range map[string]int32{
+ "num_speakers": req.GetNumSpeakers(), "min_speakers": req.GetMinSpeakers(), "max_speakers": req.GetMaxSpeakers(),
+ } {
+ if v != 0 {
+ f.Set(name, strconv.Itoa(int(v)))
+ }
+ }
+ for name, v := range map[string]float32{
+ "clustering_threshold": req.GetClusteringThreshold(),
+ "min_duration_on": req.GetMinDurationOn(),
+ "min_duration_off": req.GetMinDurationOff(),
+ } {
+ if v != 0 {
+ f.Set(name, formatFloat(v))
+ }
+ }
+ if req.GetIncludeText() {
+ f.Set("include_text", "true")
+ }
+
+ var resp struct {
+ Duration float64 `json:"duration"`
+ Language string `json:"language"`
+ NumSpeakers int32 `json:"num_speakers"`
+ Segments []struct {
+ ID int32 `json:"id"`
+ Speaker string `json:"speaker"`
+ Label string `json:"label"`
+ Start float32 `json:"start"`
+ End float32 `json:"end"`
+ Text string `json:"text"`
+ } `json:"segments"`
+ }
+ if err := p.postForm(context.Background(), "/v1/audio/diarization",
+ multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}, &resp); err != nil {
+ return pb.DiarizeResponse{}, err
+ }
+ var segments []*pb.DiarizeSegment
+ for _, s := range resp.Segments {
+ // The upstream renames speakers to SPEAKER_NN and keeps the backend's
+ // own id in label. Core renames again, so hand it the raw id: the
+ // caller then sees the same speakers and labels a direct call gives.
+ speaker := s.Label
+ if speaker == "" {
+ speaker = s.Speaker
+ }
+ segments = append(segments, &pb.DiarizeSegment{
+ Id: s.ID, Start: s.Start, End: s.End, Speaker: speaker, Text: s.Text,
+ })
+ }
+ return pb.DiarizeResponse{
+ Segments: segments, NumSpeakers: resp.NumSpeakers, Duration: float32(resp.Duration), Language: resp.Language,
+ }, nil
+}
+
+func (p *LocalAIProxy) VAD(req *pb.VADRequest) (pb.VADResponse, error) {
+ var resp struct {
+ Segments []struct {
+ Start float32 `json:"start"`
+ End float32 `json:"end"`
+ } `json:"segments"`
+ }
+ body := map[string]any{"model": p.model(""), "audio": req.GetAudio()}
+ if err := p.postJSON(context.Background(), "/v1/vad", body, &resp); err != nil {
+ return pb.VADResponse{}, err
+ }
+ var segments []*pb.VADSegment
+ for _, s := range resp.Segments {
+ segments = append(segments, &pb.VADSegment{Start: s.Start, End: s.End})
+ }
+ return pb.VADResponse{Segments: segments}, nil
+}
+
+func (p *LocalAIProxy) SoundDetection(ctx context.Context, req *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) {
+ f := url.Values{}
+ f.Set("model", p.model(""))
+ if v := req.GetTopK(); v != 0 {
+ f.Set("top_k", strconv.Itoa(int(v)))
+ }
+ if v := req.GetThreshold(); v != 0 {
+ f.Set("threshold", formatFloat(v))
+ }
+ var resp struct {
+ Detections []struct {
+ Index int32 `json:"index"`
+ Label string `json:"label"`
+ Score float32 `json:"score"`
+ } `json:"detections"`
+ }
+ if err := p.postForm(ctx, "/v1/audio/classification",
+ multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetSrc()}}}, &resp); err != nil {
+ return nil, err
+ }
+ out := &pb.SoundDetectionResponse{}
+ for _, d := range resp.Detections {
+ out.Detections = append(out.Detections, &pb.SoundClass{Index: d.Index, Label: d.Label, Score: d.Score})
+ }
+ return out, nil
+}
+
+// stemsHeader names the other outputs of a separation run: the body carries
+// one file, and the rest are served under /generated-audio/.
+const stemsHeader = "X-Audio-Stems"
+
+// AudioTransform uploads the input (and reference) to /audio/transformations
+// and writes the returned audio to Dst. SampleRate and Samples stay 0: the
+// REST reply does not report them, and core only uses them for tracing.
+func (p *LocalAIProxy) AudioTransform(req *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) {
+ f := url.Values{}
+ f.Set("model", p.model(""))
+ for k, v := range req.GetParams() {
+ f.Set("params["+k+"]", v)
+ }
+ form := multipartForm{fields: f, files: []formFile{{field: "audio", path: req.GetAudioPath()}}}
+ if ref := req.GetReferencePath(); ref != "" {
+ form.files = append(form.files, formFile{field: "reference", path: ref})
+ }
+
+ ctx := context.Background()
+ header, err := p.postMultipartToFile(ctx, "/audio/transformations", form, req.GetDst())
+ if err != nil {
+ return nil, err
+ }
+ return &pb.AudioTransformResult{
+ Dst: req.GetDst(),
+ ReferenceProvided: req.GetReferencePath() != "",
+ Stems: p.fetchStems(ctx, header.Get(stemsHeader), req.GetDst()),
+ }, nil
+}
+
+// fetchStems downloads each stem the upstream names and writes it beside
+// dst, where core looks for them. A stem that cannot be fetched is dropped
+// with a warning rather than failing the call: the caller still gets the
+// audio it asked for.
+func (p *LocalAIProxy) fetchStems(ctx context.Context, header, dst string) []*pb.AudioTransformStem {
+ if header == "" {
+ return nil
+ }
+ var entries []struct {
+ Name string `json:"name"`
+ URL string `json:"url"`
+ }
+ if err := json.Unmarshal([]byte(header), &entries); err != nil {
+ xlog.Warn("localai-proxy: ignore malformed stems header", "error", err)
+ return nil
+ }
+ dir := filepath.Dir(dst)
+ prefix := strings.TrimSuffix(filepath.Base(dst), filepath.Ext(dst))
+ var stems []*pb.AudioTransformStem
+ for _, e := range entries {
+ // Only files under /generated-audio/ are stems. Anything else is not
+ // a path this API hands out, so it is not fetched.
+ escaped, ok := strings.CutPrefix(e.URL, "/generated-audio/")
+ if !ok || e.Name == "" {
+ xlog.Warn("localai-proxy: skip stem with unexpected url", "name", e.Name, "url", e.URL)
+ continue
+ }
+ name, err := url.PathUnescape(escaped)
+ if err != nil || name == "" || name != filepath.Base(name) || name == "." || name == ".." {
+ xlog.Warn("localai-proxy: skip stem with unsafe file name", "name", e.Name, "url", e.URL)
+ continue
+ }
+ local := filepath.Join(dir, prefix+"-"+name)
+ if err := p.getToFile(ctx, e.URL, local); err != nil {
+ xlog.Warn("localai-proxy: stem download failed", "name", e.Name, "error", err)
+ continue
+ }
+ stems = append(stems, &pb.AudioTransformStem{Name: e.Name, Dst: local})
+ }
+ return stems
+}
+
+// formatFloat prints a float32 form value without float64 noise
+// (0.2, not 0.20000000298023224).
+func formatFloat(v float32) string {
+ return strconv.FormatFloat(float64(v), 'g', -1, 32)
+}
diff --git a/backend/go/localai-proxy/audio_test.go b/backend/go/localai-proxy/audio_test.go
new file mode 100644
index 000000000..8f5013c7b
--- /dev/null
+++ b/backend/go/localai-proxy/audio_test.go
@@ -0,0 +1,489 @@
+package main
+
+import (
+ "context"
+ "fmt"
+ "net/http"
+ "os"
+ "path/filepath"
+ "time"
+
+ . "github.com/onsi/ginkgo/v2"
+ . "github.com/onsi/gomega"
+ "google.golang.org/grpc/codes"
+
+ pb "github.com/mudler/LocalAI/pkg/grpc/proto"
+)
+
+// writeInput writes content to a fresh temp file and returns its path, so a
+// spec can check the upstream received exactly those bytes.
+func writeInput(name, content string) string {
+ path := filepath.Join(GinkgoT().TempDir(), name)
+ Expect(os.WriteFile(path, []byte(content), 0o600)).To(Succeed())
+ return path
+}
+
+// drainBytes collects every chunk sent on ch until it is closed.
+func drainBytes(ch chan []byte) <-chan [][]byte {
+ done := make(chan [][]byte, 1)
+ go func() {
+ var got [][]byte
+ for c := range ch {
+ got = append(got, c)
+ }
+ done <- got
+ }()
+ return done
+}
+
+func drainTranscript(ch chan *pb.TranscriptStreamResponse) <-chan []*pb.TranscriptStreamResponse {
+ done := make(chan []*pb.TranscriptStreamResponse, 1)
+ go func() {
+ var got []*pb.TranscriptStreamResponse
+ for c := range ch {
+ got = append(got, c)
+ }
+ done <- got
+ }()
+ return done
+}
+
+// cutMidStream answers 200, flushes body, then drops the connection without
+// finishing the chunked encoding, as a crashed upstream would.
+func cutMidStream(contentType, body string) http.HandlerFunc {
+ return func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", contentType)
+ w.WriteHeader(http.StatusOK)
+ _, _ = w.Write([]byte(body))
+ w.(http.Flusher).Flush()
+ conn, _, err := w.(http.Hijacker).Hijack()
+ if err == nil {
+ _ = conn.Close()
+ }
+ }
+}
+
+var _ = Describe("audio methods", func() {
+ var up *fakeUpstream
+
+ BeforeEach(func() {
+ up = newFakeUpstream()
+ DeferCleanup(up.Close)
+ })
+
+ Describe("TTS", func() {
+ It("posts to /tts and writes the upstream audio to Dst", func() {
+ p := loadProxy(up, nil)
+ up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-audio-bytes"})
+ dst := filepath.Join(GinkgoT().TempDir(), "out.wav")
+ lang := "it"
+ instr := "cheerful"
+
+ Expect(p.TTS(&pb.TTSRequest{
+ Text: "ciao", Model: "local-path", Dst: dst, Voice: "v1", Language: &lang,
+ Instructions: &instr, Params: map[string]string{"speed": "1.2"},
+ })).To(Succeed())
+
+ got, err := os.ReadFile(dst)
+ Expect(err).NotTo(HaveOccurred())
+ Expect(string(got)).To(Equal("RIFF-audio-bytes"))
+ req := up.last()
+ Expect(req.Path).To(Equal("/tts"))
+ Expect(req.JSON).To(Equal(map[string]any{
+ "model": "remote-model", "input": "ciao", "voice": "v1", "language": "it",
+ "instructions": "cheerful", "params": map[string]any{"speed": "1.2"},
+ }))
+ })
+
+ It("maps an upstream failure and leaves no partial file", func() {
+ p := loadProxy(up, nil)
+ up.script("/tts", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"})
+ dst := filepath.Join(GinkgoT().TempDir(), "out.wav")
+ err := p.TTS(&pb.TTSRequest{Text: "x", Dst: dst})
+ Expect(codeOf(err)).To(Equal(codes.Unavailable))
+ Expect(dst).NotTo(BeAnExistingFile())
+ })
+ })
+
+ Describe("TTSStream", func() {
+ It("forwards the chunked WAV body in order and closes the channel", func() {
+ header := "RIFF" + string(make([]byte, 40))
+ up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "audio/wav")
+ w.WriteHeader(http.StatusOK)
+ for _, c := range []string{header, "pcm-1", "pcm-2"} {
+ _, _ = w.Write([]byte(c))
+ w.(http.Flusher).Flush()
+ time.Sleep(10 * time.Millisecond)
+ }
+ })
+ DeferCleanup(up.Close)
+ p := loadProxy(up, nil)
+
+ out := make(chan []byte)
+ done := drainBytes(out)
+ Expect(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)).To(Succeed())
+
+ var chunks [][]byte
+ Eventually(done).Should(Receive(&chunks))
+ Expect(chunks).NotTo(BeEmpty())
+ Expect(string(chunks[0][:4])).To(Equal("RIFF"))
+ var all []byte
+ for _, c := range chunks {
+ all = append(all, c...)
+ }
+ Expect(string(all)).To(Equal(header + "pcm-1pcm-2"))
+ })
+
+ It("asks the upstream to stream", func() {
+ p := loadProxy(up, nil)
+ up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF"})
+ out := make(chan []byte)
+ done := drainBytes(out)
+ Expect(p.TTSStream(&pb.TTSRequest{Text: "hi", Voice: "v"}, out)).To(Succeed())
+ Eventually(done).Should(Receive())
+ Expect(up.last().JSON).To(HaveKeyWithValue("stream", true))
+ Expect(up.last().JSON).To(HaveKeyWithValue("input", "hi"))
+ })
+
+ It("reports a mid-stream disconnect as Unavailable and still closes the channel", func() {
+ cut := newFakeUpstreamWithHandler(cutMidStream("audio/wav", "RIFFpartial"))
+ DeferCleanup(cut.Close)
+ p := loadProxy(cut, nil)
+ out := make(chan []byte)
+ done := drainBytes(out)
+ err := p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)
+ Expect(codeOf(err)).To(Equal(codes.Unavailable))
+ Eventually(done).Should(Receive())
+ })
+
+ It("closes the channel when the upstream refuses the request", func() {
+ p := loadProxy(up, nil)
+ up.script("/tts", scriptedResponse{Status: http.StatusBadRequest, Body: "bad"})
+ out := make(chan []byte)
+ done := drainBytes(out)
+ Expect(codeOf(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out))).To(Equal(codes.InvalidArgument))
+ Eventually(done).Should(Receive())
+ })
+ })
+
+ Describe("SoundGeneration", func() {
+ It("posts the ElevenLabs body to /v1/sound-generation and writes Dst", func() {
+ p := loadProxy(up, nil)
+ up.script("/v1/sound-generation", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-sfx"})
+ dst := filepath.Join(GinkgoT().TempDir(), "sfx.wav")
+ dur, temp, bpm := float32(4.5), float32(0.3), int32(120)
+ sample, instrumental := true, false
+
+ Expect(p.SoundGeneration(&pb.SoundGenerationRequest{
+ Text: "rain", Dst: dst, Duration: &dur, Temperature: &temp, Sample: &sample,
+ Bpm: &bpm, Caption: ptr("cap"), Lyrics: ptr("la"), Keyscale: ptr("C major"),
+ Language: ptr("en"), Timesignature: ptr("4/4"), Instrumental: &instrumental,
+ })).To(Succeed())
+
+ got, err := os.ReadFile(dst)
+ Expect(err).NotTo(HaveOccurred())
+ Expect(string(got)).To(Equal("RIFF-sfx"))
+ Expect(up.last().JSON).To(Equal(map[string]any{
+ "model_id": "remote-model", "text": "rain", "duration_seconds": 4.5,
+ "prompt_influence": 0.3, "do_sample": true, "bpm": float64(120),
+ "caption": "cap", "lyrics": "la", "keyscale": "C major", "language": "en",
+ "timesignature": "4/4", "instrumental": false,
+ }))
+ })
+
+ It("refuses audio-conditioned generation it cannot upload", func() {
+ p := loadProxy(up, nil)
+ err := p.SoundGeneration(&pb.SoundGenerationRequest{Text: "x", Src: ptr("/tmp/in.wav")})
+ Expect(codeOf(err)).To(Equal(codes.Unimplemented))
+ Expect(up.recorded()).To(BeEmpty())
+ })
+ })
+
+ Describe("AudioTranscription", func() {
+ It("uploads Dst with the request fields and maps the seconds-based result", func() {
+ p := loadProxy(up, nil)
+ up.replyJSON("/v1/audio/transcriptions", map[string]any{
+ "text": "hello world", "language": "en", "duration": 2.5,
+ "segments": []any{map[string]any{
+ "id": 0, "start": 0.5, "end": 1.25, "text": "hello world", "tokens": []int{1, 2},
+ "speaker": "SPEAKER_00",
+ "words": []any{map[string]any{"start": 0.5, "end": 0.75, "text": "hello"}},
+ }},
+ })
+ in := writeInput("speech.wav", "RIFF-speech")
+
+ res, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{
+ Dst: in, Language: "en", Translate: true, Diarize: true, Prompt: "names: Ada",
+ Temperature: 0.2, TimestampGranularities: []string{"word"},
+ })
+ Expect(err).NotTo(HaveOccurred())
+
+ req := up.last()
+ Expect(req.Path).To(Equal("/v1/audio/transcriptions"))
+ Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-speech"}))
+ Expect(req.Fields).To(Equal(map[string]string{
+ "model": "remote-model", "language": "en", "translate": "true", "diarize": "true",
+ "prompt": "names: Ada", "temperature": "0.2", "timestamp_granularities[]": "word",
+ "response_format": "verbose_json",
+ }))
+
+ Expect(res.Text).To(Equal("hello world"))
+ Expect(res.Language).To(Equal("en"))
+ Expect(res.Duration).To(BeNumerically("==", 2.5))
+ Expect(res.Segments).To(HaveLen(1))
+ seg := res.Segments[0]
+ // pb carries nanoseconds (core reads them as time.Duration).
+ Expect(seg.Start).To(Equal(int64(500 * time.Millisecond)))
+ Expect(seg.End).To(Equal(int64(1250 * time.Millisecond)))
+ Expect(seg.Text).To(Equal("hello world"))
+ Expect(seg.Tokens).To(Equal([]int32{1, 2}))
+ Expect(seg.Speaker).To(Equal("SPEAKER_00"))
+ Expect(seg.Words).To(HaveLen(1))
+ Expect(seg.Words[0].Start).To(Equal(int64(500 * time.Millisecond)))
+ Expect(seg.Words[0].Text).To(Equal("hello"))
+ })
+
+ It("sends diarize=false explicitly because the upstream defaults it on", func() {
+ p := loadProxy(up, nil)
+ up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "x"})
+ _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")})
+ Expect(err).NotTo(HaveOccurred())
+ f := up.last().Fields
+ Expect(f).To(HaveKeyWithValue("diarize", "false"))
+ // Unset language and translate are left to the upstream model config.
+ Expect(f).NotTo(HaveKey("language"))
+ Expect(f).NotTo(HaveKey("translate"))
+ })
+
+ It("maps a 4xx to InvalidArgument", func() {
+ p := loadProxy(up, nil)
+ up.script("/v1/audio/transcriptions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad audio"})
+ _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")})
+ Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
+ })
+ })
+
+ Describe("AudioTranscriptionStream", func() {
+ It("emits deltas, then the final result, and closes the channel", func() {
+ p := loadProxy(up, nil)
+ up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{
+ sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}),
+ sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "lo"}),
+ sseJSON(map[string]any{
+ "type": "transcript.text.done", "text": "hello", "language": "en", "duration": 1.5,
+ "segments": []any{map[string]any{"id": 0, "start": 0.0, "end": 1.5, "text": "hello"}},
+ }),
+ "[DONE]",
+ }})
+ out := make(chan *pb.TranscriptStreamResponse)
+ done := drainTranscript(out)
+
+ Expect(p.AudioTranscriptionStream(context.Background(),
+ &pb.TranscriptRequest{Dst: writeInput("a.wav", "RIFF-a"), Stream: true}, out)).To(Succeed())
+
+ var got []*pb.TranscriptStreamResponse
+ Eventually(done).Should(Receive(&got))
+ Expect(got).To(HaveLen(3))
+ Expect(got[0].Delta).To(Equal("hel"))
+ Expect(got[1].Delta).To(Equal("lo"))
+ final := got[2].FinalResult
+ Expect(final).NotTo(BeNil())
+ Expect(final.Text).To(Equal("hello"))
+ Expect(final.Language).To(Equal("en"))
+ Expect(final.Duration).To(BeNumerically("==", 1.5))
+ Expect(final.Segments).To(HaveLen(1))
+ Expect(final.Segments[0].End).To(Equal(int64(1500 * time.Millisecond)))
+
+ req := up.last()
+ Expect(req.Fields).To(HaveKeyWithValue("stream", "true"))
+ Expect(req.Files).To(HaveKeyWithValue("file", "RIFF-a"))
+ })
+
+ It("returns an upstream error event as Unavailable", func() {
+ p := loadProxy(up, nil)
+ up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{
+ sseJSON(map[string]any{"type": "error", "error": map[string]any{"message": "decoder died"}}),
+ "[DONE]",
+ }})
+ out := make(chan *pb.TranscriptStreamResponse)
+ done := drainTranscript(out)
+ err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out)
+ Expect(codeOf(err)).To(Equal(codes.Unavailable))
+ Expect(err).To(MatchError(ContainSubstring("decoder died")))
+ Eventually(done).Should(Receive())
+ })
+
+ It("reports an upstream disconnect before the final result as Unavailable", func() {
+ frame := "data: " + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}) + "\n\n"
+ cut := newFakeUpstreamWithHandler(cutMidStream("text/event-stream", frame))
+ DeferCleanup(cut.Close)
+ p := loadProxy(cut, nil)
+ out := make(chan *pb.TranscriptStreamResponse)
+ done := drainTranscript(out)
+ err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out)
+ Expect(codeOf(err)).To(Equal(codes.Unavailable))
+ var got []*pb.TranscriptStreamResponse
+ Eventually(done).Should(Receive(&got))
+ Expect(got).To(HaveLen(1))
+ })
+
+ It("closes the channel when the local file is missing", func() {
+ p := loadProxy(up, nil)
+ out := make(chan *pb.TranscriptStreamResponse)
+ done := drainTranscript(out)
+ err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: "/nonexistent.wav"}, out)
+ Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
+ Eventually(done).Should(Receive())
+ Expect(up.recorded()).To(BeEmpty())
+ })
+ })
+
+ Describe("Diarize", func() {
+ It("uploads Dst with the tuning fields and maps the segments", func() {
+ p := loadProxy(up, nil)
+ up.replyJSON("/v1/audio/diarization", map[string]any{
+ "task": "diarize", "duration": 3.0, "language": "en", "num_speakers": 2,
+ "segments": []any{
+ map[string]any{"id": 0, "speaker": "SPEAKER_00", "label": "spk_a", "start": 0.0, "end": 1.5, "text": "hi"},
+ map[string]any{"id": 1, "speaker": "SPEAKER_01", "start": 1.5, "end": 3.0},
+ },
+ })
+ res, err := p.Diarize(&pb.DiarizeRequest{
+ Dst: writeInput("talk.wav", "RIFF-talk"), Language: "en", NumSpeakers: 2, MinSpeakers: 1,
+ MaxSpeakers: 3, ClusteringThreshold: 0.5, MinDurationOn: 0.1, MinDurationOff: 0.2, IncludeText: true,
+ })
+ Expect(err).NotTo(HaveOccurred())
+
+ req := up.last()
+ Expect(req.Path).To(Equal("/v1/audio/diarization"))
+ Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-talk"}))
+ Expect(req.Fields).To(Equal(map[string]string{
+ "model": "remote-model", "language": "en", "num_speakers": "2", "min_speakers": "1",
+ "max_speakers": "3", "clustering_threshold": "0.5", "min_duration_on": "0.1",
+ "min_duration_off": "0.2", "include_text": "true", "response_format": "verbose_json",
+ }))
+
+ Expect(res.NumSpeakers).To(Equal(int32(2)))
+ Expect(res.Duration).To(BeNumerically("==", 3.0))
+ Expect(res.Language).To(Equal("en"))
+ Expect(res.Segments).To(HaveLen(2))
+ // The raw upstream label survives, so core's own normalisation
+ // yields the same speakers a direct call would.
+ Expect(res.Segments[0].Speaker).To(Equal("spk_a"))
+ Expect(res.Segments[0].Text).To(Equal("hi"))
+ Expect(res.Segments[1].Speaker).To(Equal("SPEAKER_01"))
+ Expect(res.Segments[1].Start).To(BeNumerically("==", 1.5))
+ Expect(res.Segments[1].Id).To(Equal(int32(1)))
+ })
+ })
+
+ Describe("VAD", func() {
+ It("posts the samples to /v1/vad and maps the segments", func() {
+ p := loadProxy(up, nil)
+ up.replyJSON("/v1/vad", map[string]any{"segments": []any{map[string]any{"start": 0.25, "end": 1.5}}})
+ res, err := p.VAD(&pb.VADRequest{Audio: []float32{0.5, -0.25}})
+ Expect(err).NotTo(HaveOccurred())
+ Expect(up.last().JSON).To(Equal(map[string]any{"model": "remote-model", "audio": []any{0.5, -0.25}}))
+ Expect(res.Segments).To(HaveLen(1))
+ Expect(res.Segments[0].Start).To(BeNumerically("==", 0.25))
+ Expect(res.Segments[0].End).To(BeNumerically("==", 1.5))
+ })
+ })
+
+ Describe("SoundDetection", func() {
+ It("uploads Src to /v1/audio/classification and maps the detections", func() {
+ p := loadProxy(up, nil)
+ up.replyJSON("/v1/audio/classification", map[string]any{"model": "remote-model", "detections": []any{
+ map[string]any{"index": 74, "label": "Dog", "score": 0.9},
+ map[string]any{"index": 0, "label": "Speech", "score": 0.4},
+ }})
+ res, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{
+ Src: writeInput("bark.wav", "RIFF-bark"), TopK: 5, Threshold: 0.25,
+ })
+ Expect(err).NotTo(HaveOccurred())
+ req := up.last()
+ Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-bark"}))
+ Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "top_k": "5", "threshold": "0.25"}))
+ Expect(res.Detections).To(HaveLen(2))
+ Expect(res.Detections[0].Label).To(Equal("Dog"))
+ Expect(res.Detections[0].Index).To(Equal(int32(74)))
+ Expect(res.Detections[0].Score).To(BeNumerically("~", 0.9, 1e-6))
+ })
+ })
+
+ Describe("AudioTransform", func() {
+ It("uploads audio and reference with params and writes the result to Dst", func() {
+ p := loadProxy(up, nil)
+ up.script("/audio/transformations", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-clean"})
+ dst := filepath.Join(GinkgoT().TempDir(), "transform.wav")
+
+ res, err := p.AudioTransform(&pb.AudioTransformRequest{
+ AudioPath: writeInput("mic.wav", "RIFF-mic"), ReferencePath: writeInput("ref.wav", "RIFF-ref"),
+ Dst: dst, Params: map[string]string{"noise_gate": "true"},
+ })
+ Expect(err).NotTo(HaveOccurred())
+ got, err := os.ReadFile(dst)
+ Expect(err).NotTo(HaveOccurred())
+ Expect(string(got)).To(Equal("RIFF-clean"))
+ Expect(res.Dst).To(Equal(dst))
+ Expect(res.ReferenceProvided).To(BeTrue())
+ Expect(res.Stems).To(BeEmpty())
+
+ req := up.last()
+ Expect(req.Path).To(Equal("/audio/transformations"))
+ Expect(req.Files).To(Equal(map[string]string{"audio": "RIFF-mic", "reference": "RIFF-ref"}))
+ Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "params[noise_gate]": "true"}))
+ })
+
+ It("fetches the stems the upstream names and writes them beside Dst", func() {
+ p := loadProxy(up, nil)
+ stems := `[{"name":"vocals","url":"/generated-audio/sep%20vocals.wav"},{"name":"evil","url":"/etc/passwd"}]`
+ up.script("/audio/transformations", scriptedResponse{
+ Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-drums",
+ Header: map[string]string{"X-Audio-Stems": stems},
+ })
+ up.script("/generated-audio/sep vocals.wav", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-vocals"})
+ dir := GinkgoT().TempDir()
+ dst := filepath.Join(dir, "transform.wav")
+
+ res, err := p.AudioTransform(&pb.AudioTransformRequest{AudioPath: writeInput("song.wav", "RIFF-song"), Dst: dst})
+ Expect(err).NotTo(HaveOccurred())
+ Expect(res.ReferenceProvided).To(BeFalse())
+ Expect(res.Stems).To(HaveLen(1))
+ Expect(res.Stems[0].Name).To(Equal("vocals"))
+ // Core keeps only stems that are direct children of Dst's directory.
+ Expect(filepath.Dir(res.Stems[0].Dst)).To(Equal(dir))
+ got, err := os.ReadFile(res.Stems[0].Dst)
+ Expect(err).NotTo(HaveOccurred())
+ Expect(string(got)).To(Equal("RIFF-vocals"))
+
+ var paths []string
+ for _, r := range up.recorded() {
+ paths = append(paths, r.Method+" "+r.Path)
+ }
+ Expect(paths).To(Equal([]string{"POST /audio/transformations", "GET /generated-audio/sep vocals.wav"}))
+ })
+ })
+
+ It("names the upload file after the local file", func() {
+ // The upstream saves the upload under its base name and some handlers
+ // pick the decoder by extension, so the name must reach it intact.
+ var name string
+ named := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
+ _, fh, err := r.FormFile("file")
+ if err == nil {
+ name = fh.Filename
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = fmt.Fprint(w, `{"detections":[]}`)
+ })
+ DeferCleanup(named.Close)
+ p := loadProxy(named, nil)
+ _, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{Src: writeInput("clip.mp3", "ID3")})
+ Expect(err).NotTo(HaveOccurred())
+ Expect(name).To(Equal("clip.mp3"))
+ })
+})
+
+func ptr[T any](v T) *T { return &v }
diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go
index 100a931dd..166fa1ab8 100644
--- a/backend/go/localai-proxy/client.go
+++ b/backend/go/localai-proxy/client.go
@@ -9,6 +9,7 @@ import (
"io"
"mime/multipart"
"net/http"
+ "net/url"
"os"
"path/filepath"
"strings"
@@ -37,7 +38,7 @@ func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any)
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
- req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload))
+ req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
if err != nil {
return err
}
@@ -51,48 +52,92 @@ func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any)
// paths, and LocalAI's upload endpoints take them as multipart files. The
// request_timeout_seconds limit applies.
func (p *LocalAIProxy) postMultipart(ctx context.Context, path string, fields map[string]string, fileField, filePath string, out any) error {
+ form := multipartForm{fields: url.Values{}}
+ for k, v := range fields {
+ form.fields.Set(k, v)
+ }
+ if fileField != "" {
+ form.files = []formFile{{field: fileField, path: filePath}}
+ }
+ return p.postForm(ctx, path, form, out)
+}
+
+// postForm uploads form to path and decodes the 2xx JSON reply into out. The
+// request_timeout_seconds limit applies.
+func (p *LocalAIProxy) postForm(ctx context.Context, path string, form multipartForm, out any) error {
cfg, err := p.config()
if err != nil {
return err
}
- var file *os.File
- if fileField != "" {
- // Open before contacting the upstream so a bad path is reported as a
- // request error, not as a failure of the remote host.
- if file, err = os.Open(filePath); err != nil {
- return status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", filePath, err)
- }
- defer func() { _ = file.Close() }()
- }
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
-
- // Stream the form through a pipe so large audio files are not buffered
- // in memory. The transport closes pr when the request ends, which
- // unblocks the writer on every error path.
- pr, pw := io.Pipe()
- mw := multipart.NewWriter(pw)
- go func() {
- pw.CloseWithError(writeMultipart(mw, fields, fileField, file))
- }()
-
- req, err := p.newRequest(ctx, cfg, path, pr)
+ req, err := p.newMultipartRequest(ctx, cfg, path, form)
if err != nil {
- _ = pr.CloseWithError(err)
return err
}
- req.Header.Set("Content-Type", mw.FormDataContentType())
return p.do(req, path, out)
}
-func writeMultipart(mw *multipart.Writer, fields map[string]string, fileField string, file *os.File) error {
- for k, v := range fields {
- if err := mw.WriteField(k, v); err != nil {
- return err
+// multipartForm is an upload: repeated fields (timestamp_granularities[]) need
+// url.Values, and audio transforms send two files.
+type multipartForm struct {
+ fields url.Values
+ files []formFile
+}
+
+type formFile struct {
+ field string
+ path string
+}
+
+// newMultipartRequest builds a POST whose body streams form through a pipe,
+// so large audio files are not buffered in memory. The writer goroutine owns
+// the opened files and closes them when it finishes; the transport closes the
+// pipe when the request ends, which unblocks the writer on every error path.
+func (p *LocalAIProxy) newMultipartRequest(ctx context.Context, cfg *proxyConfig, path string, form multipartForm) (*http.Request, error) {
+ // Open before contacting the upstream so a bad path is reported as a
+ // request error, not as a failure of the remote host.
+ files := make([]*os.File, 0, len(form.files))
+ closeAll := func() {
+ for _, f := range files {
+ _ = f.Close()
}
}
- if file != nil {
- part, err := mw.CreateFormFile(fileField, filepath.Base(file.Name()))
+ for _, ff := range form.files {
+ f, err := os.Open(ff.path)
+ if err != nil {
+ closeAll()
+ return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", ff.path, err)
+ }
+ files = append(files, f)
+ }
+
+ pr, pw := io.Pipe()
+ mw := multipart.NewWriter(pw)
+ go func() {
+ defer closeAll()
+ pw.CloseWithError(writeMultipart(mw, form, files))
+ }()
+
+ req, err := p.newRequest(ctx, cfg, http.MethodPost, path, pr)
+ if err != nil {
+ _ = pr.CloseWithError(err)
+ return nil, err
+ }
+ req.Header.Set("Content-Type", mw.FormDataContentType())
+ return req, nil
+}
+
+func writeMultipart(mw *multipart.Writer, form multipartForm, files []*os.File) error {
+ for k, vs := range form.fields {
+ for _, v := range vs {
+ if err := mw.WriteField(k, v); err != nil {
+ return err
+ }
+ }
+ }
+ for i, file := range files {
+ part, err := mw.CreateFormFile(form.files[i].field, filepath.Base(file.Name()))
if err != nil {
return err
}
@@ -115,11 +160,30 @@ func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (*
if err != nil {
return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
}
- req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload))
+ req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
+ return p.doStream(req, path)
+}
+
+// postMultipartStream uploads form and returns the open response of a 2xx
+// reply; the caller must close its body. Like postStream, only ctx bounds it.
+func (p *LocalAIProxy) postMultipartStream(ctx context.Context, path string, form multipartForm) (*http.Response, error) {
+ cfg, err := p.config()
+ if err != nil {
+ return nil, err
+ }
+ req, err := p.newMultipartRequest(ctx, cfg, path, form)
+ if err != nil {
+ return nil, err
+ }
+ return p.doStream(req, path)
+}
+
+// doStream runs req and returns the open response of a 2xx reply.
+func (p *LocalAIProxy) doStream(req *http.Request, path string) (*http.Response, error) {
resp, err := p.client.Do(req)
if err != nil {
return nil, transportError(path, err)
@@ -131,8 +195,8 @@ func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (*
return resp, nil
}
-func (p *LocalAIProxy) newRequest(ctx context.Context, cfg *proxyConfig, path string, body io.Reader) (*http.Request, error) {
- req, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.base+path, body)
+func (p *LocalAIProxy) newRequest(ctx context.Context, cfg *proxyConfig, method, path string, body io.Reader) (*http.Request, error) {
+ req, err := http.NewRequestWithContext(ctx, method, cfg.base+path, body)
if err != nil {
return nil, status.Errorf(codes.Internal, "localai-proxy: build %s request: %v", path, err)
}
@@ -224,3 +288,88 @@ func statusError(path string, resp *http.Response) error {
xlog.Warn("localai-proxy: upstream error", "path", path, "status", resp.StatusCode)
return status.Error(code, fmt.Sprintf("localai-proxy: upstream %s returned %d: %s", path, resp.StatusCode, msg))
}
+
+// postJSONToFile sends body as JSON to path and writes a 2xx reply's body,
+// which is audio rather than JSON, to dst. The request_timeout_seconds limit
+// applies.
+func (p *LocalAIProxy) postJSONToFile(ctx context.Context, path string, body any, dst string) error {
+ cfg, err := p.config()
+ if err != nil {
+ return err
+ }
+ payload, err := json.Marshal(body)
+ if err != nil {
+ return status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
+ }
+ ctx, cancel := withTimeout(ctx, cfg)
+ defer cancel()
+ req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
+ if err != nil {
+ return err
+ }
+ req.Header.Set("Content-Type", "application/json")
+ _, err = p.doToFile(req, path, dst)
+ return err
+}
+
+// postMultipartToFile uploads form to path and writes a 2xx reply's body to
+// dst, returning the reply headers for endpoints that describe extra outputs
+// there. The request_timeout_seconds limit applies.
+func (p *LocalAIProxy) postMultipartToFile(ctx context.Context, path string, form multipartForm, dst string) (http.Header, error) {
+ cfg, err := p.config()
+ if err != nil {
+ return nil, err
+ }
+ ctx, cancel := withTimeout(ctx, cfg)
+ defer cancel()
+ req, err := p.newMultipartRequest(ctx, cfg, path, form)
+ if err != nil {
+ return nil, err
+ }
+ return p.doToFile(req, path, dst)
+}
+
+// getToFile downloads path to dst. The request_timeout_seconds limit applies.
+func (p *LocalAIProxy) getToFile(ctx context.Context, path, dst string) error {
+ cfg, err := p.config()
+ if err != nil {
+ return err
+ }
+ ctx, cancel := withTimeout(ctx, cfg)
+ defer cancel()
+ req, err := p.newRequest(ctx, cfg, http.MethodGet, path, nil)
+ if err != nil {
+ return err
+ }
+ _, err = p.doToFile(req, path, dst)
+ return err
+}
+
+// doToFile runs req and writes a 2xx body to dst. A failed copy removes dst:
+// core serves whatever file it finds there, and a truncated recording must
+// not pass for a finished one.
+func (p *LocalAIProxy) doToFile(req *http.Request, path, dst string) (http.Header, error) {
+ resp, err := p.doStream(req, path)
+ if err != nil {
+ return nil, err
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ f, err := os.Create(dst)
+ if err != nil {
+ return nil, status.Errorf(codes.Internal, "localai-proxy: create %s: %v", dst, err)
+ }
+ _, copyErr := io.Copy(f, resp.Body)
+ closeErr := f.Close()
+ if copyErr != nil || closeErr != nil {
+ _ = os.Remove(dst)
+ if copyErr != nil {
+ if ctxErr := req.Context().Err(); ctxErr != nil {
+ return nil, transportError(path, ctxErr)
+ }
+ return nil, transportError(path, copyErr)
+ }
+ return nil, status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, closeErr)
+ }
+ return resp.Header, nil
+}
diff --git a/backend/go/localai-proxy/fake_upstream_test.go b/backend/go/localai-proxy/fake_upstream_test.go
index 6ced223a8..933e012be 100644
--- a/backend/go/localai-proxy/fake_upstream_test.go
+++ b/backend/go/localai-proxy/fake_upstream_test.go
@@ -27,10 +27,11 @@ type recordedRequest struct {
}
// scriptedResponse is the reply for one path. SSE, when set, is written as
-// "data: " events and wins over Body.
+// "data: " events and wins over Body. Header adds response headers.
type scriptedResponse struct {
Status int
ContentType string
+ Header map[string]string
Body string
SSE []string
}
@@ -133,6 +134,9 @@ func (f *fakeUpstream) serve(w http.ResponseWriter, r *http.Request) {
if resp.ContentType != "" {
w.Header().Set("Content-Type", resp.ContentType)
}
+ for k, v := range resp.Header {
+ w.Header().Set(k, v)
+ }
w.WriteHeader(resp.Status)
_, _ = io.WriteString(w, resp.Body)
}