diff --git a/AGENTS.md b/AGENTS.md index 96c043e..35dca3e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -29,6 +29,7 @@ go test -tags e2e -run 'TestE2E' -timeout 15m -v . # LIVE provider e2e (see be | `openai.go` / `gemini.go` / `anthropic.go` | Per-format request builders, response/stream mappers, model listing | | `responses.go` | OpenAI Responses API (`/v1/responses`) for GPT-5.6+ tools+reasoning | | `tts.go` | Text-to-speech: `Speak`/`SpeakRequest`/`SpeakResult` — OpenAI-compat `/audio/speech` + Gemini AUDIO modality, no transcoding | +| `stt.go` | Speech-to-text: `Transcribe`/`TranscribeRequest`/`TranscribeResult` — OpenAI-compat `/audio/transcriptions` (multipart), 25MB input cap | | `sse.go` | SSE parser (abort-safe via `done` channel) + idle-watchdog pump | | `retry.go` | Backoff/jitter/`Retry-After`/`retrySleep` (8 attempts, cap 30s) | | `provider.go` | Built-in registry, quirks flags, config validation | diff --git a/README.md b/README.md index 2f3c1ee..1ff5b97 100644 --- a/README.md +++ b/README.md @@ -256,6 +256,25 @@ res, err := sdk.Speak(ctx, "openai", "tts-1", llm.SpeakRequest{ The SDK never transcodes: it returns the provider's bytes plus the MIME type the provider declared, falling back to the wire format's well-known default (`audio/mpeg` on the OpenAI path; Gemini's `inlineData.mimeType` is always present). A 2xx JSON error envelope from a gateway surfaces as a typed `*APIError` — JSON is never returned as audio. Requests carry the same retry ladder and error taxonomy as chat; empty `Text` or `Voice` fail fast with a `ConfigError`. +## Speech-to-text + +`Transcribe` converts audio bytes to text via the provider's transcription endpoint. v1 supports the OpenAI-compatible wire format (`POST {base}/audio/transcriptions`, multipart/form-data); other formats return a `ConfigError`. The SDK never touches the filesystem — callers own the audio bytes. + +```go +res, err := sdk.Transcribe(ctx, "openai", "whisper-1", llm.TranscribeRequest{ + Audio: audio, // required, non-empty, ≤25MB + Filename: "probe.mp3", // file part name (default "audio.wav") + MIMEType: "audio/mpeg", // audio part content type (default application/octet-stream) + Language: "en", // optional ISO-639-1 hint + Prompt: "context words", // optional conditioning text +}) +// res.Text — recognized text +// res.Model — model that produced the transcription +// res.Language, res.DurationSec — provider-reported, zero when omitted +``` + +Requests carry the same retry ladder and error taxonomy as chat. Oversized audio (>25MB) and empty `Audio`/`Model` fail fast with a `ConfigError` before any network I/O. A 2xx body that is not JSON surfaces as a typed `*APIError`. Fields the provider does not report stay zero — the SDK never guesses. + ## Thread safety `SDK` and `Provider` are safe for concurrent use. `ChatClient` is safe for concurrent `Call`/`CallStream`; `SetRequestTimeout` is race-safe (atomic swap) but should still be called before the first request so in-flight calls use one timeout. Learn-once state is shared per provider via atomics — monotonic, converging, race-free. diff --git a/e2e_test.go b/e2e_test.go index 0f42d0f..2c07c48 100644 --- a/e2e_test.go +++ b/e2e_test.go @@ -462,3 +462,37 @@ func TestE2ETTSGemini(t *testing.T) { } t.Logf("audio=%d bytes mime=%q model=%q", len(res.Audio), res.MIMEType, res.Model) } + +// TestE2ETranscribeOpenAI probes the OpenAI /audio/transcriptions path +// against the live endpoint (built-in registry, key via OPENAI_API_KEY). +// Self-contained: audio is synthesized first via Speak, then transcribed +// (model override via OPENAI_STT_E2E_MODEL, default whisper-1). Asserts +// only SDK guarantees: the call succeeds and non-empty text comes back. +// Transcript accuracy is never asserted. +func TestE2ETranscribeOpenAI(t *testing.T) { + const keyEnv = "OPENAI_API_KEY" + key := e2eEnvKey(t, keyEnv) + sdk := New(WithProvider("openai", WithAPIKey(key))) + spoken, err := sdk.Speak(t.Context(), "openai", e2eTTSModel(t, "openai", "tts-1"), SpeakRequest{ + Text: "go-llm-sdk speech to text probe.", + Voice: "alloy", + }) + if err != nil { + t.Fatalf("Speak: %v", err) + } + sttModel := "whisper-1" + if v := strings.TrimSpace(os.Getenv("OPENAI_STT_E2E_MODEL")); v != "" { + sttModel = v + } + res, err := sdk.Transcribe(t.Context(), "openai", sttModel, TranscribeRequest{ + Audio: spoken.Audio, + Filename: "probe.mp3", + }) + if err != nil { + t.Fatalf("Transcribe: %v", err) + } + if res.Text == "" { + t.Errorf("Transcribe returned empty text") + } + t.Logf("text=%q duration=%.1fs model=%q", res.Text, res.DurationSec, res.Model) +} diff --git a/stt.go b/stt.go new file mode 100644 index 0000000..08b82a5 --- /dev/null +++ b/stt.go @@ -0,0 +1,270 @@ +package llm + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "mime/multipart" + "net/http" + "net/textproto" + "strings" + "time" +) + +// ── STT (speech-to-text) ───────────────────────────────────────────────── +// +// Transcribe turns audio bytes into text. v1 supports the +// OpenAI-compatible wire format: POST {base}/audio/transcriptions, +// multipart/form-data, JSON response. Any other format (gemini, +// anthropic, …) is a ConfigError. The SDK never touches the +// filesystem: callers own the audio bytes. The chat invariants apply +// unchanged — canonical errors, keys never in error text, retry +// ladder identical to the buffered chat path. + +// maxTranscribeAudioBytes bounds the audio payload accepted for one +// transcription request (matches OpenAI's documented 25MB upload limit). +// Oversized audio is rejected with a ConfigError before any network I/O. +const maxTranscribeAudioBytes = 25 << 20 + +// TranscribeRequest describes one speech-to-text conversion. +type TranscribeRequest struct { + Audio []byte // required, non-empty: raw audio bytes (caller-owned) + Filename string // file part name for the audio (default "audio.wav") + MIMEType string // optional content type for the audio part + Language string // optional ISO-639-1 hint, OpenAI-compat only + Prompt string // optional conditioning text, OpenAI-compat only + Format string // optional response_format (provider default when empty) +} + +// TranscribeResult carries the recognized text. Fields the provider +// omits stay zero — unknown data stays unknown (invariant #5). +type TranscribeResult struct { + Text string + Model string + Language string + DurationSec float64 +} + +// Transcribe converts speech to text with the named provider and +// model. Buffered only, with the same retry ladder as chat. +func (s *SDK) Transcribe(ctx context.Context, providerID, model string, req TranscribeRequest) (*TranscribeResult, error) { + if len(req.Audio) == 0 { + return nil, &ConfigError{Msg: "transcribe request requires non-empty Audio"} + } + if len(req.Audio) > maxTranscribeAudioBytes { + return nil, &ConfigError{Msg: "transcribe audio exceeds 25MB limit"} + } + if strings.TrimSpace(model) == "" { + return nil, &ConfigError{Msg: "transcribe request requires a model"} + } + p, err := s.Provider(providerID) + if err != nil { + return nil, err + } + if !p.Authenticated() { + return nil, &ConfigError{Msg: providerID + " has no API key (set " + strings.ToUpper(providerID) + "_API_KEY or use WithAPIKey)"} + } + if p.invalid { + return nil, &ConfigError{Msg: providerID + " has an invalid configuration"} + } + pc := newProviderClient(p.cfg, newBufferedHTTP(s.rt, s.timeout), nil) + return pc.transcribe(ctx, model, req) +} + +// escapeQuotes escapes quotes and backslashes in a multipart filename, +// mirroring mime/multipart's internal escape. +func escapeQuotes(s string) string { + r := strings.NewReplacer("\\", "\\\\", `"`, `\"`) + return r.Replace(s) +} + +// buildTranscribeRequest serializes the multipart body and target URL. +func (pc *providerClient) buildTranscribeRequest(model string, req TranscribeRequest) ([]byte, string, error) { + if pc.cfg.Format != FormatOpenAI { + return nil, "", &ConfigError{Msg: "provider format " + string(pc.cfg.Format) + " does not support transcription"} + } + filename := req.Filename + if filename == "" { + filename = "audio.wav" + } + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + mime := req.MIMEType + if mime == "" { + mime = "application/octet-stream" + } + hdr := textproto.MIMEHeader{} + hdr.Set("Content-Disposition", fmt.Sprintf(`form-data; name="file"; filename="%s"`, escapeQuotes(filename))) + hdr.Set("Content-Type", mime) + fh, err := mw.CreatePart(hdr) + if err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + if _, err := fh.Write(req.Audio); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + if err := mw.WriteField("model", model); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + if req.Language != "" { + if err := mw.WriteField("language", req.Language); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + } + if req.Prompt != "" { + if err := mw.WriteField("prompt", req.Prompt); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + } + if req.Format != "" { + if err := mw.WriteField("response_format", req.Format); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + } + if err := mw.Close(); err != nil { + return nil, "", fmt.Errorf("llm: build transcription request: %w", err) + } + return buf.Bytes(), pc.base + "/audio/transcriptions", nil +} + +// transcribe runs the STT request against one provider with retry +// semantics identical to the buffered chat path. +func (pc *providerClient) transcribe(ctx context.Context, model string, req TranscribeRequest) (*TranscribeResult, error) { + body, url, err := pc.buildTranscribeRequest(model, req) + if err != nil { + return nil, err + } + ctype := "multipart/form-data; boundary=" + multipartBodyBoundary(body) + + var ( + lastErr error + rateErr *APIError + rateRA time.Duration + ) + for attempt := 0; attempt <= maxRetries; attempt++ { + if err := ctx.Err(); err != nil { + return nil, err + } + data, ra, err := pc.postMultipart(ctx, url, body, ctype) + if err != nil { + var apiErr *APIError + if errors.As(err, &apiErr) { + switch { + case apiErr.Status == http.StatusTooManyRequests && billingExhausted(apiErr): + return nil, apiErr + case apiErr.Status == http.StatusTooManyRequests: + rateErr, rateRA, lastErr = apiErr, ra, apiErr + if attempt < maxRetries { + if !retrySleep(ctx, retryDelay(ra, attempt)) { + return nil, &RateLimitError{APIError: *rateErr, Attempts: attempt + 1, RetryAfter: rateRA} + } + continue + } + case apiErr.Retryable && attempt < maxRetries: + lastErr = apiErr + if !retrySleep(ctx, retryDelay(ra, attempt)) { + return nil, ctx.Err() + } + continue + } + if rateErr != nil && !apiErr.Retryable && apiErr.Status != http.StatusTooManyRequests { + return nil, apiErr + } + if rateErr != nil { + return nil, &RateLimitError{APIError: *rateErr, Attempts: attempt + 1, RetryAfter: rateRA} + } + return nil, apiErr + } + // Transport error — retryable. + lastErr = err + if attempt < maxRetries { + if !retrySleep(ctx, retryDelay(0, attempt)) { + return nil, ctx.Err() + } + continue + } + return nil, fmt.Errorf("llm: retry exhausted (%d attempts): %w", maxRetries+1, err) + } + res, perr := parseTranscribeResponse(data) + if perr != nil { + // 2xx with an undecodable body is a provider protocol + // failure — surface it through the typed error taxonomy + // at the actual HTTP status (never a plain fmt.Errorf). + return nil, &APIError{ + Provider: pc.cfg.ID, + Status: http.StatusOK, + Message: perr.Error(), + } + } + res.Model = model + return res, nil + } + return nil, lastErr +} + +// multipartBodyBoundary extracts the boundary from a multipart writer's +// Content-Type header line. +func multipartBodyBoundary(body []byte) string { + // The boundary is the last line of the leading preamble; cheaper and + // stricter: parse it from the first line of the body. + line := body + if i := bytes.IndexByte(body, '\r'); i >= 0 { + line = body[:i] + } else if i := bytes.IndexByte(body, '\n'); i >= 0 { + line = body[:i] + } + return strings.TrimPrefix(string(line), "--") +} + +type transcribeResponse struct { + Text string `json:"text"` + Language string `json:"language"` + Duration float64 `json:"duration"` +} + +// parseTranscribeResponse decodes a buffered transcription response. +// Both plain json ({text}) and verbose_json ({text,language,duration}) +// decode through the same struct — fields the provider omits stay zero. +func parseTranscribeResponse(data []byte) (*TranscribeResult, error) { + var resp transcribeResponse + if err := json.Unmarshal(data, &resp); err != nil { + return nil, fmt.Errorf("llm: decode transcription response: %w", err) + } + return &TranscribeResult{ + Text: resp.Text, + Language: resp.Language, + DurationSec: resp.Duration, + }, nil +} + +// postMultipart sends one transcription request and reads the full JSON +// body (capped). Mirrors postAudio but with the multipart content type. +func (pc *providerClient) postMultipart(ctx context.Context, url string, body []byte, ctype string) (data []byte, ra time.Duration, err error) { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return nil, 0, &ConfigError{Msg: "build request: " + err.Error()} + } + req.Header.Set("Content-Type", ctype) + pc.setAuthHeaders(req.Header) + + resp, err := pc.buffered().Do(req) + if err != nil { + return nil, 0, err + } + defer func() { _ = resp.Body.Close() }() + data, err = io.ReadAll(io.LimitReader(resp.Body, maxResponseSize+1)) + if err != nil { + return nil, 0, err + } + if len(data) > maxResponseSize { + return nil, 0, fmt.Errorf("llm: response exceeds %d bytes", maxResponseSize) + } + ra = parseRetryAfter(resp.Header.Get("Retry-After"), time.Now()) + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return nil, ra, pc.httpError(resp.StatusCode, data) + } + return data, ra, nil +} diff --git a/stt_test.go b/stt_test.go new file mode 100644 index 0000000..4985e49 --- /dev/null +++ b/stt_test.go @@ -0,0 +1,305 @@ +package llm + +import ( + "bytes" + "errors" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +// sttRecord is what a test server captures from one Transcribe request. +type sttRecord struct { + path string + auth string + ct string + body []byte + boundary string +} + +// sttParse splits a captured multipart body into form values and the +// uploaded file part. +func sttParse(t *testing.T, rec sttRecord) (map[string]string, *multipart.FileHeader) { + t.Helper() + mr := multipart.NewReader(bytes.NewReader(rec.body), rec.boundary) + form, err := mr.ReadForm(1 << 20) + if err != nil { + t.Fatalf("parse multipart body: %v", err) + } + vals := map[string]string{} + for k, v := range form.Value { + if len(v) > 0 { + vals[k] = v[0] + } + } + var filePart *multipart.FileHeader + if fh := form.File["file"]; len(fh) > 0 { + filePart = fh[0] + } + return vals, filePart +} + +// newSTTServer records one request, then answers with the given status, +// content type, and body. +func newSTTServer(status int, ctype, body string) (*httptest.Server, *sttRecord) { + rec := &sttRecord{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + rec.path = r.URL.Path + rec.auth = r.Header.Get("Authorization") + rec.ct = r.Header.Get("Content-Type") + if i := strings.Index(rec.ct, "boundary="); i >= 0 { + rec.boundary = rec.ct[i+len("boundary="):] + } + rec.body, _ = io.ReadAll(r.Body) + if ctype != "" { + w.Header().Set("Content-Type", ctype) + } + w.WriteHeader(status) + _, _ = w.Write([]byte(body)) + })) + return srv, rec +} + +// TestTranscribe_MIMETypeHonored asserts the audio part's Content-Type +// comes from TranscribeRequest.MIMEType when set (RED: field was dead — +// CreateFormFile always wrote application/octet-stream). +func TestTranscribe_MIMETypeHonored(t *testing.T) { + srv, rec := newSTTServer(http.StatusOK, "application/json", `{"text":"hi"}`) + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + if _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{ + Audio: []byte{1, 2, 3}, + Filename: "probe.mp3", + MIMEType: "audio/mpeg", + }); err != nil { + t.Fatalf("Transcribe: %v", err) + } + _, fp := sttParse(t, *rec) + if fp == nil { + t.Fatal("no file part") + } + if got := fp.Header.Get("Content-Type"); got != "audio/mpeg" { + t.Errorf("file part Content-Type = %q, want audio/mpeg", got) + } +} + +// TestTranscribe_MIMETypeDefault asserts the audio part falls back to +// application/octet-stream when MIMEType is unset. +func TestTranscribe_MIMETypeDefault(t *testing.T) { + srv, rec := newSTTServer(http.StatusOK, "application/json", `{"text":"hi"}`) + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + if _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{ + Audio: []byte{1, 2, 3}, + Filename: "probe.mp3", + }); err != nil { + t.Fatalf("Transcribe: %v", err) + } + _, fp := sttParse(t, *rec) + if fp == nil { + t.Fatal("no file part") + } + if got := fp.Header.Get("Content-Type"); got != "application/octet-stream" { + t.Errorf("file part Content-Type = %q, want application/octet-stream", got) + } +} + +// TestTranscribe_AudioTooLarge asserts oversized audio is rejected before +// any network round trip (RED: no request-side cap existed). +func TestTranscribe_AudioTooLarge(t *testing.T) { + srv, _ := newSTTServer(http.StatusOK, "application/json", `{"text":"hi"}`) + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{ + Audio: make([]byte, maxTranscribeAudioBytes+1), + Filename: "big.mp3", + }) + var ce *ConfigError + if !errors.As(err, &ce) { + t.Fatalf("err = %T (%v), want *ConfigError", err, err) + } +} + +func TestTranscribe_OpenAI(t *testing.T) { + audio := []byte{0xff, 0xf3, 0x00, 0x01, 0x02} + srv, rec := newSTTServer(http.StatusOK, "application/json", `{"text":"hello world"}`) + defer srv.Close() + + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k-secret"}, srv) + res, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{ + Audio: audio, + Filename: "probe.mp3", + }) + if err != nil { + t.Fatalf("Transcribe: %v", err) + } + if rec.path != "/audio/transcriptions" { + t.Errorf("path = %q, want /audio/transcriptions", rec.path) + } + if rec.auth != "Bearer k-secret" { + t.Errorf("bearer auth not sent correctly: %q", rec.auth) + } + if !strings.HasPrefix(rec.ct, "multipart/form-data") { + t.Errorf("content type = %q, want multipart/form-data", rec.ct) + } + vals, fp := sttParse(t, *rec) + if fp == nil { + t.Fatal("no file part in request") + } + if fp.Filename != "probe.mp3" { + t.Errorf("file filename = %q, want probe.mp3", fp.Filename) + } + f, err := fp.Open() + if err != nil { + t.Fatalf("open file part: %v", err) + } + defer f.Close() + got, err := io.ReadAll(f) + if err != nil { + t.Fatalf("read file part: %v", err) + } + if !bytes.Equal(got, audio) { + t.Errorf("file bytes = %v, want %v", got, audio) + } + if vals["model"] != "whisper-1" { + t.Errorf("model = %q, want whisper-1", vals["model"]) + } + for _, k := range []string{"language", "prompt", "response_format"} { + if v, ok := vals[k]; ok { + t.Errorf("field %q = %q, want omitted when empty", k, v) + } + } + if res.Text != "hello world" { + t.Errorf("Text = %q", res.Text) + } + if res.Model != "whisper-1" { + t.Errorf("Model = %q, want whisper-1", res.Model) + } +} + +func TestTranscribe_OptionalFields(t *testing.T) { + srv, rec := newSTTServer(http.StatusOK, "application/json", + `{"text":"bonjour","language":"fr","duration":1.5}`) + defer srv.Close() + + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + res, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{ + Audio: []byte{1, 2, 3}, + Filename: "a.wav", + Language: "fr", + Prompt: "greeting", + Format: "verbose_json", + }) + if err != nil { + t.Fatalf("Transcribe: %v", err) + } + vals, _ := sttParse(t, *rec) + if vals["language"] != "fr" || vals["prompt"] != "greeting" || vals["response_format"] != "verbose_json" { + t.Errorf("optional fields = %v", vals) + } + if res.Text != "bonjour" || res.Language != "fr" || res.DurationSec != 1.5 { + t.Errorf("result = %+v", res) + } +} + +func TestTranscribe_PlainJSONZeroValues(t *testing.T) { + srv, _ := newSTTServer(http.StatusOK, "application/json", `{"text":"only text"}`) + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + res, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + if err != nil { + t.Fatalf("Transcribe: %v", err) + } + if res.Text != "only text" || res.Language != "" || res.DurationSec != 0 { + t.Errorf("result = %+v, want zero values for absent fields", res) + } +} + +func TestTranscribe_HTTP500JSONEnvelope(t *testing.T) { + srv, _ := newSTTServer(http.StatusInternalServerError, "application/json", + `{"error":{"message":"boom","type":"server_error"}}`) + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + var ae *APIError + if !errors.As(err, &ae) { + t.Fatalf("err = %T (%v), want *APIError", err, err) + } + if ae.Status != http.StatusInternalServerError { + t.Errorf("Status = %d, want 500", ae.Status) + } +} + +func TestTranscribe_2xxNonJSONBody(t *testing.T) { + srv, _ := newSTTServer(http.StatusOK, "text/plain", "not json at all") + defer srv.Close() + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + var ae *APIError + if !errors.As(err, &ae) { + t.Fatalf("err = %T (%v), want *APIError", err, err) + } + if ae.Status != http.StatusOK { + t.Errorf("Status = %d, want 200", ae.Status) + } +} + +func TestTranscribe_ConfigErrors(t *testing.T) { + cfg := ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k-secret"} + cases := []struct { + name string + sdkCfg ProviderConfig + call func(s *SDK) error + }{ + {"empty audio", cfg, func(s *SDK) error { + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Filename: "a.mp3"}) + return err + }}, + {"empty model", cfg, func(s *SDK) error { + _, err := s.Transcribe(t.Context(), "openai", "", TranscribeRequest{Audio: []byte{1}}) + return err + }}, + {"unknown provider", cfg, func(s *SDK) error { + _, err := s.Transcribe(t.Context(), "nope", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + return err + }}, + {"unauthenticated", ProviderConfig{ID: "openai", Format: FormatOpenAI}, func(s *SDK) error { + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + return err + }}, + {"non-openai format", ProviderConfig{ID: "anthropic", Format: FormatAnthropic, APIKey: "k"}, func(s *SDK) error { + _, err := s.Transcribe(t.Context(), "anthropic", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + return err + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + s := newTestSDK(t, tc.sdkCfg, nil) + err := tc.call(s) + var ce *ConfigError + if !errors.As(err, &ce) { + t.Fatalf("err = %T (%v), want *ConfigError", err, err) + } + if strings.Contains(err.Error(), "k-secret") { + t.Errorf("error leaks api key: %v", err) + } + }) + } +} + +func TestTranscribe_ResponseSizeCap(t *testing.T) { + srv, _ := newSTTServer(http.StatusOK, "application/json", string(make([]byte, 100))) + defer srv.Close() + // Respond with more than the cap regardless of the cap value. + srv.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Write(make([]byte, maxResponseSize+2)) + }) + s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k"}, srv) + _, err := s.Transcribe(t.Context(), "openai", "whisper-1", TranscribeRequest{Audio: []byte{1}}) + if err == nil || !strings.Contains(err.Error(), "exceeds") { + t.Fatalf("err = %v, want size-cap error", err) + } +} diff --git a/tts_test.go b/tts_test.go index 618b024..1734434 100644 --- a/tts_test.go +++ b/tts_test.go @@ -279,14 +279,13 @@ func TestSpeak_MissingContentType(t *testing.T) { } defer ln.Close() go func() { - conn, err := ln.Accept() - if err != nil { - return + for { + conn, err := ln.Accept() + if err != nil { + return // listener closed: test over + } + handleConn(conn) } - defer conn.Close() - body := []byte{0xff, 0xf3, 0x00, 0x01} - fmt.Fprintf(conn, "HTTP/1.1 200 OK\r\nContent-Length: %d\r\nConnection: close\r\n\r\n", len(body)) - _, _ = conn.Write(body) }() s := newTestSDK(t, ProviderConfig{ID: "openai", Format: FormatOpenAI, APIKey: "k", BaseURL: "http://" + ln.Addr().String()}, nil) res, err := s.Speak(t.Context(), "openai", "tts-1", SpeakRequest{Text: "hello", Voice: "alloy"}) @@ -298,6 +297,15 @@ func TestSpeak_MissingContentType(t *testing.T) { } } +// handleConn answers one HTTP request with a bare 200 response carrying +// audio bytes and no Content-Type header. +func handleConn(conn net.Conn) { + defer conn.Close() + body := []byte{0xff, 0xf3, 0x00, 0x01} + fmt.Fprintf(conn, "HTTP/1.1 200 OK\r\nContent-Length: %d\r\nConnection: close\r\n\r\n", len(body)) + _, _ = conn.Write(body) +} + // TestSpeak_OpenAI2xxJSONRejected guards against OpenAI-compatible gateways // answering 2xx with a JSON error envelope: the SDK must surface a typed // error, never hand JSON bytes back as audio.