diff --git a/core/http/endpoints/openai/transcription.go b/core/http/endpoints/openai/transcription.go index 49825c3a1..b4be91b59 100644 --- a/core/http/endpoints/openai/transcription.go +++ b/core/http/endpoints/openai/transcription.go @@ -174,7 +174,7 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app } if stream { - return streamTranscription(c, req, ml, *config, appConfig) + return streamTranscription(c, req, input.Model, ml, *config, appConfig) } tr, err := backend.ModelTranscriptionWithOptions(c.Request().Context(), req, ml, *config, appConfig) @@ -195,7 +195,7 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app switch responseFormat { case schema.TranscriptionResponseFormatLrc, schema.TranscriptionResponseFormatText, schema.TranscriptionResponseFormatSrt, schema.TranscriptionResponseFormatVtt: - return c.String(http.StatusOK, schema.TranscriptionResponse(tr, responseFormat)) + err = c.String(http.StatusOK, schema.TranscriptionResponse(tr, responseFormat)) case schema.TranscriptionResponseFormatJson: tr.Segments = nil tr.Words = nil @@ -236,10 +236,16 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app Words: segWords, }) } - return c.JSON(http.StatusOK, trs) + err = c.JSON(http.StatusOK, trs) default: return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format") } + if err == nil { + // Transcription exposes no canonical token counts, but successful + // requests still contribute model usage and elapsed time. + middleware.StampUsage(c, input.Model, 0, 0) + } + return err } } @@ -258,7 +264,7 @@ func validTranscriptionResponseFormat(f schema.TranscriptionResponseFormatType) // `transcript.text.done` with the assembled text, and `[DONE]`. Backends that // can't truly stream still produce a single Final event, which we surface as // one delta + done. -func streamTranscription(c echo.Context, req backend.TranscriptionRequest, ml *model.ModelLoader, config config.ModelConfig, appConfig *config.ApplicationConfig) error { +func streamTranscription(c echo.Context, req backend.TranscriptionRequest, requestedModel string, ml *model.ModelLoader, config config.ModelConfig, appConfig *config.ApplicationConfig) error { c.Response().Header().Set("Content-Type", "text/event-stream") c.Response().Header().Set("Cache-Control", "no-cache") c.Response().Header().Set("Connection", "keep-alive") @@ -278,14 +284,17 @@ func streamTranscription(c echo.Context, req backend.TranscriptionRequest, ml *m var assembled strings.Builder var finalResult *schema.TranscriptionResult + var writeErr error err := backend.ModelTranscriptionStream(c.Request().Context(), req, ml, config, appConfig, func(chunk backend.TranscriptionStreamChunk) { if chunk.Delta != "" { assembled.WriteString(chunk.Delta) - _ = writeEvent(map[string]any{ - "type": "transcript.text.delta", - "delta": chunk.Delta, - }) + if writeErr == nil { + writeErr = writeEvent(map[string]any{ + "type": "transcript.text.delta", + "delta": chunk.Delta, + }) + } } if chunk.Final != nil { finalResult = chunk.Final @@ -304,6 +313,9 @@ func streamTranscription(c echo.Context, req backend.TranscriptionRequest, ml *m c.Response().Flush() return nil } + if writeErr != nil { + return writeErr + } // Build the final event. Prefer the backend-provided final result; if the // backend only emitted deltas, synthesize the result from what we collected. @@ -315,10 +327,12 @@ func streamTranscription(c echo.Context, req backend.TranscriptionRequest, ml *m // If the backend never produced a delta but did return a final text, emit // it as a single delta so clients always see at least one delta event. if assembled.Len() == 0 && finalResult.Text != "" { - _ = writeEvent(map[string]any{ + if err := writeEvent(map[string]any{ "type": "transcript.text.delta", "delta": finalResult.Text, - }) + }); err != nil { + return err + } } // done carries the assembled text plus, when the backend produced them, // per-segment timings, audio duration, and detected language. The OpenAI @@ -353,8 +367,13 @@ func streamTranscription(c echo.Context, req backend.TranscriptionRequest, ml *m } doneEvent["segments"] = segs } - _ = writeEvent(doneEvent) - _, _ = fmt.Fprintf(c.Response().Writer, "data: [DONE]\n\n") + if err := writeEvent(doneEvent); err != nil { + return err + } + if _, err := fmt.Fprintf(c.Response().Writer, "data: [DONE]\n\n"); err != nil { + return err + } c.Response().Flush() + middleware.StampUsage(c, requestedModel, 0, 0) return nil } diff --git a/core/http/middleware/trace.go b/core/http/middleware/trace.go index ec77abbe9..02d2e6437 100644 --- a/core/http/middleware/trace.go +++ b/core/http/middleware/trace.go @@ -3,6 +3,7 @@ package middleware import ( "bufio" "bytes" + "encoding/json" "errors" "io" "mime" @@ -19,6 +20,7 @@ import ( "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/application" "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/schema" "github.com/mudler/LocalAI/core/trace/tracepersist" "github.com/mudler/xlog" ) @@ -234,7 +236,7 @@ func redactSensitiveHeaders(h http.Header) http.Header { return out } -// TraceMiddleware intercepts and logs JSON API requests and responses +// TraceMiddleware logs JSON exchanges and transcription upload metadata. func TraceMiddleware(app *application.Application) echo.MiddlewareFunc { initializeTracing(app.ApplicationConfig().DataPath, app.ApplicationConfig().TracingMaxItems) return func(next echo.HandlerFunc) echo.HandlerFunc { @@ -253,19 +255,26 @@ func TraceMiddleware(app *application.Application) echo.MiddlewareFunc { } ct, _, _ := mime.ParseMediaType(c.Request().Header.Get("Content-Type")) - if ct != "application/json" { + multipartTranscription := ct == "multipart/form-data" && (c.Path() == "/v1/audio/transcriptions" || c.Path() == "/audio/transcriptions") + if ct != "application/json" && !multipartTranscription { return next(c) } - body, err := io.ReadAll(c.Request().Body) - if err != nil { - xlog.Error("Failed to read request body") - return err + var body []byte + if !multipartTranscription { + var err error + body, err = io.ReadAll(c.Request().Body) + if err != nil { + xlog.Error("Failed to read request body") + return err + } + c.Request().Body = io.NopCloser(bytes.NewBuffer(body)) + } else { + // Leave the upload untouched. Only the endpoint should parse or + // spool audio, regardless of the configured trace capture limit. + body = []byte(`{"body_omitted":"multipart upload omitted"}`) } - // Restore the body for downstream handlers - c.Request().Body = io.NopCloser(bytes.NewBuffer(body)) - startTime := time.Now() // Cap captured payload size. Without this, /embeddings and @@ -316,6 +325,17 @@ func TraceMiddleware(app *application.Application) echo.MiddlewareFunc { c.Response().Writer = mw handlerErr := next(c) + if multipartTranscription { + metadata := map[string]string{"body_omitted": "multipart upload omitted"} + if input, ok := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest); ok && input != nil { + metadata["model"] = input.Model + } + body, _ = json.Marshal(metadata) + metadataBody, truncated := truncateForTrace(body, maxBodyBytes) + exchange.Request.Body = &metadataBody + exchange.Request.BodyTruncated = truncated + exchange.Request.BodyBytes = len(body) + } // Restore original writer unconditionally c.Response().Writer = mw.ResponseWriter diff --git a/core/http/middleware/trace_multipart_test.go b/core/http/middleware/trace_multipart_test.go new file mode 100644 index 000000000..7a96d3884 --- /dev/null +++ b/core/http/middleware/trace_multipart_test.go @@ -0,0 +1,73 @@ +// SPDX-License-Identifier: MIT +package middleware + +import ( + "bytes" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/application" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("multipart transcription traces", func() { + DescribeTable("leaves upload consumption to the handler", func(route string, enabled, expected bool) { + root := GinkgoT().TempDir() + app, err := application.New(config.WithDataPath(root), config.WithDisableStats(true), config.WithDisableLocalAIAssistant(true), config.WithSystemState(&system.SystemState{Model: system.Model{ModelsPath: root}, Backend: system.Backend{BackendsPath: root}})) + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { Expect(app.Shutdown()).To(Succeed()) }) + app.ApplicationConfig().EnableTracing = enabled + var body bytes.Buffer + form := multipart.NewWriter(&body) + f, err := form.CreateFormFile("file", "secret.wav") + Expect(err).NotTo(HaveOccurred()) + _, err = f.Write([]byte("private upload")) + Expect(err).NotTo(HaveOccurred()) + Expect(form.Close()).To(Succeed()) + req := httptest.NewRequest(http.MethodPost, route, &body) + req.Header.Set("Content-Type", form.FormDataContentType()) + originalBody := req.Body + originalLen := body.Len() + e := echo.New() + trace := TraceMiddleware(app) + ClearTraces() + e.POST(route, func(c echo.Context) error { + Expect(c.Request().Body).To(BeIdenticalTo(originalBody)) + Expect(body.Len()).To(Equal(originalLen)) + f, err := c.FormFile("file") + Expect(err).NotTo(HaveOccurred()) + reader, err := f.Open() + Expect(err).NotTo(HaveOccurred()) + defer func() { Expect(reader.Close()).To(Succeed()) }() + audio, err := io.ReadAll(reader) + Expect(err).NotTo(HaveOccurred()) + Expect(string(audio)).To(Equal("private upload")) + return c.String(http.StatusOK, "transcript") + }, trace) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + Expect(rec.Code).To(Equal(http.StatusOK)) + if !expected { + Expect(GetTraces()).To(BeEmpty()) + return + } + Eventually(GetTraces).Should(ConsistOf(HaveField("Response.Status", http.StatusOK))) + exchange := GetTraces()[0] + Expect(string(*exchange.Request.Body)).NotTo(ContainSubstring("private upload")) + Expect(string(*exchange.Request.Body)).To(ContainSubstring("multipart upload omitted")) + Expect(string(*exchange.Response.Body)).To(Equal("transcript")) + }, + Entry("v1", "/v1/audio/transcriptions", true, true), + Entry("legacy", "/audio/transcriptions", true, true), + Entry("disabled", "/v1/audio/transcriptions", false, false), + Entry("diarization excluded", "/v1/audio/diarization", true, false), + Entry("registration excluded", "/v1/voice/register", true, false), + Entry("other multipart unchanged", "/other", true, false), + ) +}) diff --git a/core/http/routes/openai.go b/core/http/routes/openai.go index ab9ceb3fa..b041305bc 100644 --- a/core/http/routes/openai.go +++ b/core/http/routes/openai.go @@ -178,6 +178,7 @@ func RegisterOpenAIRoutes(app *echo.Echo, audioHandler := openai.TranscriptEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()) audioMiddleware := []echo.MiddlewareFunc{ nodeHeaderMiddleware, + usageMiddleware, traceMiddleware, re.BuildFilteredFirstAvailableDefaultModel(config.BuildUsecaseFilterFn(config.FLAG_TRANSCRIPT)), re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }), diff --git a/core/http/routes/transcription_observability_test.go b/core/http/routes/transcription_observability_test.go new file mode 100644 index 000000000..45bf99437 --- /dev/null +++ b/core/http/routes/transcription_observability_test.go @@ -0,0 +1,189 @@ +// SPDX-License-Identifier: MIT +package routes_test + +import ( + "bytes" + "context" + "errors" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + + "github.com/mudler/LocalAI/core/application" + "github.com/mudler/LocalAI/core/config" + corehttp "github.com/mudler/LocalAI/core/http" + "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/services/routing/billing" + grpcpkg "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" +) + +type observedTranscriptionBackend struct { + grpcpkg.Backend + err error + audio []byte +} + +type failedTranscriptionWriter struct { + *httptest.ResponseRecorder +} + +func (*failedTranscriptionWriter) Write([]byte) (int, error) { + return 0, errors.New("client disconnected") +} + +func (*observedTranscriptionBackend) HealthCheck(context.Context) (bool, error) { return true, nil } +func (*observedTranscriptionBackend) IsBusy() bool { return false } +func (*observedTranscriptionBackend) Free(context.Context) error { return nil } +func (b *observedTranscriptionBackend) AudioTranscription(_ context.Context, r *pb.TranscriptRequest, _ ...ggrpc.CallOption) (*pb.TranscriptResult, error) { + var err error + b.audio, err = os.ReadFile(r.Dst) + if err != nil { + return nil, err + } + if b.err != nil { + return nil, b.err + } + return &pb.TranscriptResult{Text: "hello world", Segments: []*pb.TranscriptSegment{{Text: "hello world"}}}, nil +} +func (b *observedTranscriptionBackend) AudioTranscriptionStream(ctx context.Context, r *pb.TranscriptRequest, cb func(*pb.TranscriptStreamResponse), _ ...ggrpc.CallOption) error { + result, err := b.AudioTranscription(ctx, r) + cb(&pb.TranscriptStreamResponse{Delta: "hello "}) + if err != nil { + return err + } + cb(&pb.TranscriptStreamResponse{FinalResult: result}) + return nil +} + +var _ = Describe("transcription observability", func() { + var app *application.Application + var handler http.Handler + var fixture *observedTranscriptionBackend + var requestedModel string + BeforeEach(func() { + requestedModel = "speech-test" + root := GinkgoT().TempDir() + var err error + app, err = application.New(config.EnableTracing, config.WithDataPath(root), config.WithDisableLocalAIAssistant(true), config.WithDisableCSRF(true), config.WithSystemState(&system.SystemState{Model: system.Model{ModelsPath: root}, Backend: system.Backend{BackendsPath: root}})) + Expect(err).NotTo(HaveOccurred()) + DeferCleanup(func() { Expect(app.Shutdown()).To(Succeed()) }) + cfg := config.ModelConfig{Name: "speech-test", Backend: "whisper"} + cfg.SetDefaults() + cfg.Model = "speech.bin" + app.ModelConfigLoader().ReplaceModelConfigs([]config.ModelConfig{cfg, {Name: "speech-alias", Alias: "speech-test"}}) + fixture = &observedTranscriptionBackend{} + app.ModelLoader().SetModelRouter(func(_ context.Context, id string, _, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) { + return model.NewModelWithClient(id, "test://speech", fixture), nil + }) + handler, err = corehttp.API(app) + Expect(err).NotTo(HaveOccurred()) + middleware.ClearTraces() + }) + request := func(route, format string, stream bool) *http.Request { + var body bytes.Buffer + form := multipart.NewWriter(&body) + Expect(form.WriteField("model", requestedModel)).To(Succeed()) + Expect(form.WriteField("response_format", format)).To(Succeed()) + if stream { + Expect(form.WriteField("stream", "true")).To(Succeed()) + } + f, err := form.CreateFormFile("file", "sample.wav") + Expect(err).NotTo(HaveOccurred()) + _, err = f.Write([]byte("private audio bytes")) + Expect(err).NotTo(HaveOccurred()) + Expect(form.Close()).To(Succeed()) + req := httptest.NewRequest(http.MethodPost, route, &body) + req.Header.Set("Content-Type", form.FormDataContentType()) + return req + } + DescribeTable("records successful requests with zero tokens and traces multipart uploads", func(route, format string, stream, fail bool) { + if fail { + fixture.err = errors.New("transcription failed") + } + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, request(route, format, stream)) + expectedStatus := http.StatusOK + if fail && !stream { + expectedStatus = http.StatusInternalServerError + } + Expect(rec.Code).To(Equal(expectedStatus), rec.Body.String()) + Expect(fixture.audio).To(Equal([]byte("private audio bytes"))) + if stream { + Expect(rec.Body.String()).To(ContainSubstring("data: [DONE]")) + if fail { + Expect(rec.Body.String()).To(ContainSubstring(`"type":"error"`)) + } else { + Expect(rec.Body.String()).To(ContainSubstring(`"type":"transcript.text.done"`)) + } + } else if !fail { + Expect(rec.Body.String()).To(ContainSubstring("hello world")) + } + buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{UserID: app.FallbackUser().ID, Period: "day"}) + Expect(err).NotTo(HaveOccurred()) + if fail { + Expect(buckets).To(BeEmpty()) + } else { + Expect(buckets).To(HaveLen(1)) + Expect(buckets[0].Model).To(Equal("speech-test")) + Expect(buckets[0].RequestCount).To(Equal(int64(1))) + Expect(buckets[0].TotalTokens).To(BeZero()) + Expect(buckets[0].PromptTokens).To(BeZero()) + Expect(buckets[0].CompletionTokens).To(BeZero()) + } + Eventually(middleware.GetTraces).Should(ConsistOf(HaveField("Response.Status", expectedStatus))) + exchange := middleware.GetTraces()[0] + Expect(exchange.Request.Path).To(Equal(route)) + Expect(exchange.Duration).To(BeNumerically(">", 0)) + Expect(string(*exchange.Request.Body)).NotTo(ContainSubstring("private audio bytes")) + Expect(string(*exchange.Request.Body)).To(ContainSubstring(`"model":"speech-test"`)) + }, + Entry("JSON", "/v1/audio/transcriptions", "json", false, false), + Entry("legacy text", "/audio/transcriptions", "txt", false, false), + Entry("verbose JSON", "/v1/audio/transcriptions", "verbose_json", false, false), + Entry("SRT", "/v1/audio/transcriptions", "srt", false, false), + Entry("VTT", "/v1/audio/transcriptions", "vtt", false, false), + Entry("LRC", "/v1/audio/transcriptions", "lrc", false, false), + Entry("SSE", "/v1/audio/transcriptions", "", true, false), + Entry("backend error", "/v1/audio/transcriptions", "json", false, true), + Entry("SSE error", "/audio/transcriptions", "", true, true), + ) + DescribeTable("attributes alias usage to the requested model", func(stream bool) { + requestedModel = "speech-alias" + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, request("/v1/audio/transcriptions", "json", stream)) + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring("hello world")) + buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{UserID: app.FallbackUser().ID, Period: "day"}) + Expect(err).NotTo(HaveOccurred()) + Expect(buckets).To(HaveLen(1)) + Expect(buckets[0].Model).To(Equal("speech-alias")) + Expect(buckets[0].RequestCount).To(Equal(int64(1))) + Expect(buckets[0].TotalTokens).To(BeZero()) + }, Entry("JSON", false), Entry("SSE", true)) + DescribeTable("does not count responses the client cannot receive", func(stream bool) { + rec := &failedTranscriptionWriter{httptest.NewRecorder()} + handler.ServeHTTP(rec, request("/v1/audio/transcriptions", "json", stream)) + Expect(fixture.audio).To(Equal([]byte("private audio bytes"))) + buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{UserID: app.FallbackUser().ID, Period: "day"}) + Expect(err).NotTo(HaveOccurred()) + Expect(buckets).To(BeEmpty()) + Eventually(middleware.GetTraces).Should(ConsistOf(HaveField("Error", ContainSubstring("client disconnected")))) + }, Entry("JSON", false), Entry("SSE", true)) + It("traces invalid formats without recording successful usage", func() { + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, request("/audio/transcriptions", "invalid", false)) + Expect(rec.Code).To(Equal(http.StatusBadRequest)) + Expect(fixture.audio).To(BeNil()) + buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{UserID: app.FallbackUser().ID, Period: "day"}) + Expect(err).NotTo(HaveOccurred()) + Expect(buckets).To(BeEmpty()) + Eventually(middleware.GetTraces).Should(ConsistOf(HaveField("Response.Status", http.StatusBadRequest))) + }) +}) diff --git a/docs/content/features/audio-to-text.md b/docs/content/features/audio-to-text.md index ceeabc646..8e78f9bcf 100644 --- a/docs/content/features/audio-to-text.md +++ b/docs/content/features/audio-to-text.md @@ -32,6 +32,19 @@ For instance, with cURL: curl http://localhost:8080/v1/audio/transcriptions -H "Content-Type: multipart/form-data" -F file="@" -F model="" ``` +When usage statistics are enabled, successful requests to `/v1/audio/transcriptions` +and `/audio/transcriptions` contribute request counts and elapsed time for the +requested model name, including an alias. This includes streaming responses. Token counts remain zero +because transcription results do not provide canonical token usage. Failed +transcriptions, including errors reported within a stream, do not count as +successful usage. + +When API tracing is enabled, multipart transcription requests appear in Traces +with the model name, response, status, and elapsed time. The captured request +body contains metadata marked `multipart upload omitted`; API traces do not +store the uploaded audio. Response capture follows the configured trace size +limit. Backend traces have their own payload capture behavior. + ## Example Download one of the models from [here](https://huggingface.co/ggerganov/whisper.cpp/tree/main) in the `models` folder,