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) }