Stream WhisperX uploads safely
This commit is contained in:
@@ -1,7 +1,6 @@
|
||||
package whisperx
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -14,12 +13,16 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gitea.maximumdirect.net/eric/narratio/internal/fileops"
|
||||
)
|
||||
|
||||
const defaultMaxResponseBytes int64 = 10 * 1024 * 1024
|
||||
const (
|
||||
defaultMaxWhisperXResponseBytes int64 = 10 * 1024 * 1024
|
||||
whisperXUploadBufferSize = 32 * 1024
|
||||
)
|
||||
|
||||
// HTTPClientConfig contains parsed, deterministic WhisperX HTTP client settings.
|
||||
type HTTPClientConfig struct {
|
||||
@@ -41,6 +44,7 @@ type HTTPClient struct {
|
||||
retryDelay time.Duration
|
||||
httpClient *http.Client
|
||||
maxResponseBytes int64
|
||||
openAudio func(string) (io.ReadCloser, error)
|
||||
}
|
||||
|
||||
// NewHTTPClientFromConfigValues builds a client from config values and parses durations once.
|
||||
@@ -74,11 +78,11 @@ func NewHTTPClient(cfg HTTPClientConfig) (*HTTPClient, error) {
|
||||
return nil, fmt.Errorf("whisperx transcribe_url is required")
|
||||
}
|
||||
u, err := url.Parse(cfg.TranscribeURL)
|
||||
if err != nil || u.Scheme == "" || u.Host == "" {
|
||||
if err != nil || !u.IsAbs() || u.Host == "" || !isHTTPURLScheme(u.Scheme) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid whisperx transcribe_url %q: %w", cfg.TranscribeURL, err)
|
||||
}
|
||||
return nil, fmt.Errorf("invalid whisperx transcribe_url %q", cfg.TranscribeURL)
|
||||
return nil, fmt.Errorf("invalid whisperx transcribe_url %q: must be an absolute http or https URL", cfg.TranscribeURL)
|
||||
}
|
||||
if cfg.Timeout <= 0 {
|
||||
return nil, fmt.Errorf("whisperx timeout must be > 0")
|
||||
@@ -100,7 +104,7 @@ func NewHTTPClient(cfg HTTPClientConfig) (*HTTPClient, error) {
|
||||
|
||||
maxBytes := cfg.MaxResponseBytes
|
||||
if maxBytes <= 0 {
|
||||
maxBytes = defaultMaxResponseBytes
|
||||
maxBytes = defaultMaxWhisperXResponseBytes
|
||||
}
|
||||
|
||||
return &HTTPClient{
|
||||
@@ -111,6 +115,7 @@ func NewHTTPClient(cfg HTTPClientConfig) (*HTTPClient, error) {
|
||||
retryDelay: cfg.RetryDelay,
|
||||
httpClient: client,
|
||||
maxResponseBytes: maxBytes,
|
||||
openAudio: func(path string) (io.ReadCloser, error) { return os.Open(path) },
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -185,52 +190,40 @@ func (c *HTTPClient) Transcribe(ctx context.Context, req TranscribeRequest) (Tra
|
||||
}
|
||||
|
||||
func (c *HTTPClient) doTranscribeAttempt(ctx context.Context, audioPath string) (int, []byte, error) {
|
||||
bodyBuf := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(bodyBuf)
|
||||
upload := newMultipartUpload(ctx, audioPath, c.language, c.openAudio)
|
||||
defer upload.Close()
|
||||
|
||||
fileWriter, err := writer.CreateFormFile("file", filepath.Base(audioPath))
|
||||
if err != nil {
|
||||
return 0, nil, fmt.Errorf("create multipart file field: %w", err)
|
||||
}
|
||||
|
||||
audioFile, err := os.Open(audioPath)
|
||||
if err != nil {
|
||||
return 0, nil, fmt.Errorf("open audio file %q: %w", audioPath, err)
|
||||
}
|
||||
if _, err := io.Copy(fileWriter, audioFile); err != nil {
|
||||
_ = audioFile.Close()
|
||||
return 0, nil, fmt.Errorf("copy audio file %q: %w", audioPath, err)
|
||||
}
|
||||
if err := audioFile.Close(); err != nil {
|
||||
return 0, nil, fmt.Errorf("close audio file %q: %w", audioPath, err)
|
||||
}
|
||||
|
||||
if err := writer.WriteField("language", c.language); err != nil {
|
||||
return 0, nil, fmt.Errorf("write language form field: %w", err)
|
||||
}
|
||||
if err := writer.Close(); err != nil {
|
||||
return 0, nil, fmt.Errorf("close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.url.String(), bodyBuf)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.url.String(), upload)
|
||||
if err != nil {
|
||||
return 0, nil, fmt.Errorf("build whisperx request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Content-Type", upload.contentType)
|
||||
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
_ = upload.Close()
|
||||
if producerErr := upload.Wait(); producerErr != nil {
|
||||
return 0, nil, fmt.Errorf("stream whisperx request body: %w", producerErr)
|
||||
}
|
||||
return 0, nil, fmt.Errorf("perform whisperx request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
data, err := readBounded(resp.Body, c.maxResponseBytes)
|
||||
if err != nil {
|
||||
return resp.StatusCode, nil, fmt.Errorf("read whisperx response body: %w", err)
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
_ = upload.Close()
|
||||
if _, err := readWhisperXResponse(resp.Body, c.maxResponseBytes); err != nil {
|
||||
return resp.StatusCode, nil, fmt.Errorf("read whisperx response body: %w", err)
|
||||
}
|
||||
return resp.StatusCode, nil, fmt.Errorf("whisperx returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return resp.StatusCode, nil, fmt.Errorf("whisperx returned status %d", resp.StatusCode)
|
||||
if err := upload.Wait(); err != nil {
|
||||
return resp.StatusCode, nil, fmt.Errorf("stream whisperx request body: %w", err)
|
||||
}
|
||||
|
||||
data, err := readWhisperXResponse(resp.Body, c.maxResponseBytes)
|
||||
if err != nil {
|
||||
return resp.StatusCode, nil, fmt.Errorf("read whisperx response body: %w", err)
|
||||
}
|
||||
return resp.StatusCode, data, nil
|
||||
}
|
||||
@@ -266,18 +259,183 @@ func (c *HTTPClient) shouldRetry(parent context.Context, err error, status int)
|
||||
return false
|
||||
}
|
||||
|
||||
func readBounded(r io.Reader, maxBytes int64) ([]byte, error) {
|
||||
func readWhisperXResponse(r io.Reader, maxBytes int64) ([]byte, error) {
|
||||
limited := io.LimitReader(r, maxBytes+1)
|
||||
data, err := io.ReadAll(limited)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if int64(len(data)) > maxBytes {
|
||||
return nil, fmt.Errorf("response exceeds max size %d bytes", maxBytes)
|
||||
return nil, fmt.Errorf("whisperx response exceeds configured limit of %d bytes", maxBytes)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func isHTTPURLScheme(scheme string) bool {
|
||||
switch strings.ToLower(scheme) {
|
||||
case "http", "https":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
type multipartUpload struct {
|
||||
reader *io.PipeReader
|
||||
writer *io.PipeWriter
|
||||
contentType string
|
||||
done chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
audio io.Closer
|
||||
err error
|
||||
aborted bool
|
||||
}
|
||||
|
||||
func newMultipartUpload(ctx context.Context, audioPath, language string, openAudio func(string) (io.ReadCloser, error)) *multipartUpload {
|
||||
reader, writer := io.Pipe()
|
||||
multipartWriter := multipart.NewWriter(writer)
|
||||
upload := &multipartUpload{
|
||||
reader: reader,
|
||||
writer: writer,
|
||||
contentType: multipartWriter.FormDataContentType(),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
go func() {
|
||||
err := upload.write(ctx, multipartWriter, audioPath, language, openAudio)
|
||||
if err != nil {
|
||||
_ = writer.CloseWithError(err)
|
||||
} else {
|
||||
_ = writer.Close()
|
||||
}
|
||||
upload.mu.Lock()
|
||||
upload.err = err
|
||||
upload.audio = nil
|
||||
upload.mu.Unlock()
|
||||
close(upload.done)
|
||||
}()
|
||||
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
upload.abort()
|
||||
case <-upload.done:
|
||||
}
|
||||
}()
|
||||
|
||||
return upload
|
||||
}
|
||||
|
||||
func (u *multipartUpload) Read(p []byte) (int, error) {
|
||||
return u.reader.Read(p)
|
||||
}
|
||||
|
||||
func (u *multipartUpload) Close() error {
|
||||
u.abort()
|
||||
<-u.done
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *multipartUpload) Wait() error {
|
||||
<-u.done
|
||||
u.mu.Lock()
|
||||
defer u.mu.Unlock()
|
||||
return u.err
|
||||
}
|
||||
|
||||
func (u *multipartUpload) write(ctx context.Context, writer *multipart.Writer, audioPath, language string, openAudio func(string) (io.ReadCloser, error)) error {
|
||||
fileWriter, err := writer.CreateFormFile("file", filepath.Base(audioPath))
|
||||
if err != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("create multipart file field: %w", err))
|
||||
}
|
||||
|
||||
audioFile, err := openAudio(audioPath)
|
||||
if err != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("open audio file %q: %w", audioPath, err))
|
||||
}
|
||||
u.setAudio(audioFile)
|
||||
|
||||
_, copyErr := io.CopyBuffer(fileWriter, &contextReader{ctx: ctx, reader: audioFile}, make([]byte, whisperXUploadBufferSize))
|
||||
closeErr := audioFile.Close()
|
||||
u.clearAudio(audioFile)
|
||||
if copyErr != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("copy audio file %q: %w", audioPath, copyErr))
|
||||
}
|
||||
if closeErr != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("close audio file %q: %w", audioPath, closeErr))
|
||||
}
|
||||
|
||||
if err := writer.WriteField("language", language); err != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("write language form field: %w", err))
|
||||
}
|
||||
if err := writer.Close(); err != nil {
|
||||
return u.producerError(ctx, fmt.Errorf("close multipart writer: %w", err))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *multipartUpload) producerError(ctx context.Context, err error) error {
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
u.mu.Lock()
|
||||
aborted := u.aborted
|
||||
u.mu.Unlock()
|
||||
if aborted {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (u *multipartUpload) setAudio(audio io.Closer) {
|
||||
u.mu.Lock()
|
||||
u.audio = audio
|
||||
aborted := u.aborted
|
||||
u.mu.Unlock()
|
||||
if aborted {
|
||||
_ = audio.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func (u *multipartUpload) clearAudio(audio io.Closer) {
|
||||
u.mu.Lock()
|
||||
if u.audio == audio {
|
||||
u.audio = nil
|
||||
}
|
||||
u.mu.Unlock()
|
||||
}
|
||||
|
||||
func (u *multipartUpload) abort() {
|
||||
u.mu.Lock()
|
||||
if u.aborted {
|
||||
u.mu.Unlock()
|
||||
return
|
||||
}
|
||||
u.aborted = true
|
||||
audio := u.audio
|
||||
u.mu.Unlock()
|
||||
|
||||
_ = u.reader.Close()
|
||||
if audio != nil {
|
||||
_ = audio.Close()
|
||||
}
|
||||
}
|
||||
|
||||
type contextReader struct {
|
||||
ctx context.Context
|
||||
reader io.Reader
|
||||
}
|
||||
|
||||
func (r *contextReader) Read(p []byte) (int, error) {
|
||||
select {
|
||||
case <-r.ctx.Done():
|
||||
return 0, r.ctx.Err()
|
||||
default:
|
||||
return r.reader.Read(p)
|
||||
}
|
||||
}
|
||||
|
||||
func writeFileAtomic(path string, data []byte, perm os.FileMode) error {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
return fmt.Errorf("path is required")
|
||||
|
||||
Reference in New Issue
Block a user