mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
feat(localai-proxy): serve speech, transcription and audio APIs remotely
The proxy now forwards TTS, streaming TTS, sound generation, transcription (plain and streaming), diarization, VAD, sound detection and audio transforms to the upstream LocalAI. Streaming TTS passes the upstream WAV bytes through unchanged. A streaming transcription that stops before its final frame, or sends an error frame, fails with Unavailable instead of ending as a short success. Transcription always sends diarize, because the upstream treats a missing field as true. Audio transforms also download the separation stems the upstream names and write them beside Dst. Sound generation from a source clip returns Unimplemented, because the REST endpoint has no field for the clip. The multipart helper now takes repeated fields and several files. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
882f51fd62
commit
87fd7da531
4 files changed
+1128
-33
No files matched your search
@@ -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)
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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
|
||||
}
|
||||
@@ -27,10 +27,11 @@ type recordedRequest struct {
|
||||
}
|
||||
|
||||
// scriptedResponse is the reply for one path. SSE, when set, is written as
|
||||
// "data: <frame>" events and wins over Body.
|
||||
// "data: <frame>" 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)
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user