diff --git a/.agents/api-endpoints-and-auth.md b/.agents/api-endpoints-and-auth.md index e3a63b1e5..4ed161725 100644 --- a/.agents/api-endpoints-and-auth.md +++ b/.agents/api-endpoints-and-auth.md @@ -394,6 +394,7 @@ When adding a new endpoint: - [ ] Error responses use `schema.ErrorResponse` format (or `echo.NewHTTPError` with a mapped gRPC status — see the `mapBackendError` helper in `core/http/endpoints/localai/images.go`) - [ ] Tests cover both authenticated and unauthenticated access - [ ] Swagger regenerated (`make swagger`) if you changed any `@Router`/`@Tags`/`@Param` annotation +- [ ] Stateful feature: distributed mode chosen and documented (see [distributed-state.md](distributed-state.md)) ## Companion: MCP admin tool surface diff --git a/.agents/distributed-state.md b/.agents/distributed-state.md new file mode 100644 index 000000000..00173edfe --- /dev/null +++ b/.agents/distributed-state.md @@ -0,0 +1,142 @@ +# Distributed-aware state + +Every frontend process is a stateless replica: any of them can take a +request, and any of them can be the one that starts, stops, or dies first. +Runtime state that lives only in a Go map on one frontend — an in-memory +cache, a set of pins, a scheduler's next-run time, a background probe loop — +diverges silently the moment there is more than one frontend. Each replica +believes its own copy, clients see different answers depending on which +frontend they hit, and nothing errors: the failure mode is quiet disagreement, +not a crash. + +The failover chains feature is the concrete example this rule generalizes +from: chain pins, target health, and the active target of each chain used to +live in a plain map per frontend. Two frontends could disagree about which +target a chain was using, and a pin set through one frontend would be +invisible through another. + +## The rule + +Before you add runtime state (an in-memory map, a cache, a set of pins, a +scheduler, a background loop, a health prober, anything that is not re-derived +fresh from the database or the request itself), pick one of the four modes +below and write down which one you picked on the feature's own docs page +under `docs/content/`. "I didn't think about it" is not one of the modes. + +### 1. Shared + +Use `syncstate.SyncedMap` (`core/services/syncstate`) when every frontend +needs to see the same state and any frontend may be the one that reads or +writes it next. Deltas replicate over the message bus to every subscriber, +including the publisher, so all frontends converge without an extra polling +loop. + +Add a `Store` (`syncstate.Config.Store`) when the state must survive a +restart of the whole cluster — the store is the durable source of truth, and +new frontends rehydrate from it instead of starting empty. Without a `Store`, +give a `Loader` for local, disk-backed rehydration in standalone mode. + +Examples: +- Finetune jobs — `core/services/finetune/service.go`. A `Store` persists jobs + across restarts; without one, a `Loader` reloads from disk. +- Failover pins, targets, and chains — `core/services/failover/distsync`. + `failover.pins` is `Store`-backed (pins must survive a restart); + `failover.targets` and `failover.chains` are ephemeral live-health state + with no `Store`, republished by the leader every 10 s. + +**Gotcha:** `Reconcile` with neither a `Store` nor a `Loader` does **nothing** +— there is nothing for it to pull from, so it cannot help a late joiner catch +up. If your map has no durable backing, a leader (or another privileged +writer) must republish its live state after a reconnect instead of relying on +`Reconcile` to recover it. + +**Gotcha:** a hydrate (on `Start`, after a NATS reconnect, and on every +`Reconcile` tick) replaces the map's contents **without firing `OnApply`**. +If `OnApply` feeds derived state (the failover manager's pins, a cache, a +running process), that state stays stale after a reconnect or a repaired +missed delta. Re-sync the derived state from the map's `Snapshot()`: on a +periodic tick, and/or from an `OnReconnect` callback registered after the +map's `Start` (callbacks run in registration order, so the map has already +re-hydrated). `failover.Manager.ReconcilePins` is the example. + +### 2. Single-runner + +Use `advisorylock.RunLeaderLoop` or `advisorylock.TryWithLockCtx` +(`core/services/advisorylock`) when exactly one frontend in the cluster should +run some work on a schedule, and it does not matter which one. Add your lock +key to `keys.go` — keys are a single global namespace across the database, so +pick a number that is not already taken. + +`RunLeaderLoop`/`TryWithLockCtx` leadership is **not sticky**: the lock is +acquired and released around each tick, so a different frontend can win it +next time. That is fine for idempotent, tick-shaped work (a cleanup pass, a +periodic scan) where it does not matter who ran the last one. + +When leadership must be sticky — the same frontend keeps a role across many +ticks, for example because it holds live connections or in-memory session +state tied to that role — use `advisorylock.HeldLock` instead. It keeps a +dedicated database session and PostgreSQL TCP keepalives so a dead host's +lock is freed in about 30 s rather than after the OS's multi-hour keepalive +default. The failover prober (`core/application/failover_distributed.go`) +uses `HeldLock` for exactly this reason: it needs to stay the same leader +across probe ticks, not re-elect on every one. + +Examples: +- Tick-shaped, non-sticky: the node health monitor + (`core/services/nodes/health.go`), via `TryWithLockCtx`. +- Sticky: the failover prober (`core/application/failover_distributed.go`), + via `HeldLock`. + +### 3. Stateless per request + +Nothing to share: the feature re-derives everything it needs from the +request, the database, or another already-distributed source, so there is no +in-memory state to diverge. Most handlers are this by default — no action +needed beyond noting it if a reviewer might otherwise ask. + +### 4. Per-instance (documented exception) + +Some state is legitimately local to a frontend — a warm in-process cache +that only saves work when it is warm, a local rate limiter for the process's +own outbound connections. This is allowed, but only when the feature's docs +page says so explicitly and states the consequence (which behavior differs +between frontends, and why that is acceptable). An undocumented per-instance +map is a bug, not a design choice. + +## Tests + +A shared or single-runner feature includes a **two-instance test** that +proves the two frontends actually agree, not just that each one works alone. +Build it on `testutil.NewFakeBus()` (`core/services/testutil/fakebus.go`): +construct two instances of the feature against the same fake bus, mutate +state through one, and assert the other observes it. + +`FakeBus` delivers every publish **synchronously**, including back to the +publisher itself — that is what makes the two-instance test deterministic +without polling, but it is also a trap: **never publish while holding a lock +that the apply path also takes.** The publish call re-enters your own apply +handler on the same goroutine before `Publish` returns, so a lock held across +the publish deadlocks against itself. Release the lock, then publish. + +## Checklist + +When your PR adds or changes state that lives longer than a single request: + +- [ ] Mode chosen: shared (`syncstate.SyncedMap`) / single-runner + (`advisorylock`) / stateless / documented per-instance +- [ ] If shared: `Store` added if the state must survive a cluster restart, + or a `Loader` for standalone rehydration; if neither, a leader + republishes after reconnect instead of relying on bare `Reconcile`; + state derived through `OnApply` re-syncs from `Snapshot()` after a + hydrate (hydrate fires no `OnApply`) +- [ ] If single-runner: new lock key added to `keys.go`; `HeldLock` chosen + over `RunLeaderLoop`/`TryWithLockCtx` if leadership must be sticky + across ticks +- [ ] Publishing to the shared state never happens while holding a lock the + apply path also takes (see the `FakeBus` gotcha above) +- [ ] Two-instance test on `testutil.NewFakeBus()` for shared/single-runner + state +- [ ] The chosen mode (and, for per-instance, the reason) is written on the + feature's docs page under `docs/content/` +- [ ] Standalone mode (no NATS/DB) still behaves correctly — most + shared/single-runner code paths must degrade gracefully alone diff --git a/.github/backend-matrix.yml b/.github/backend-matrix.yml index fa61fa514..8649209cf 100644 --- a/.github/backend-matrix.yml +++ b/.github/backend-matrix.yml @@ -5948,6 +5948,35 @@ include: dockerfile: "./backend/Dockerfile.golang" context: "./" ubuntu-version: '2404' + # localai-proxy + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/amd64' + platform-tag: 'amd64' + tag-latest: 'auto' + tag-suffix: '-cpu-localai-proxy' + runs-on: 'ubuntu-latest' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "localai-proxy" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/arm64' + platform-tag: 'arm64' + tag-latest: 'auto' + tag-suffix: '-cpu-localai-proxy' + runs-on: 'ubuntu-24.04-arm' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "localai-proxy" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' # valkey-store - build-type: '' cuda-major-version: "" @@ -6753,6 +6782,10 @@ includeDarwin: tag-suffix: "-metal-darwin-arm64-cloud-proxy" build-type: "metal" lang: "go" + - backend: "localai-proxy" + tag-suffix: "-metal-darwin-arm64-localai-proxy" + build-type: "metal" + lang: "go" - backend: "valkey-store" tag-suffix: "-metal-darwin-arm64-valkey-store" build-type: "metal" diff --git a/.gitignore b/.gitignore index a4c84a0e8..f59bbe453 100644 --- a/.gitignore +++ b/.gitignore @@ -29,6 +29,7 @@ LocalAI # Root-level build artifacts when running `go build ./...` against # Go backend packages whose main lives under backend/go/. /cloud-proxy +/localai-proxy /local-store /valkey-store # prevent above rules from omitting the helm chart @@ -48,6 +49,10 @@ tests/e2e-aio/backends # tests/e2e/mock-backend/.gitignore covers the same binary; kept here too so # the artifact stays ignored if that scoped file is ever removed. /tests/e2e/mock-backend/mock-backend +# The cloud-proxy backend binary the e2e suite runs next to the mock backend. +/tests/e2e/mock-backend/cloud-proxy +# The localai-proxy backend binary, built next to it the same way. +/tests/e2e/mock-backend/localai-proxy release/ diff --git a/AGENTS.md b/AGENTS.md index abbb342df..e3e8b24f7 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -33,6 +33,7 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants] | [.agents/localai-assistant-mcp.md](.agents/localai-assistant-mcp.md) | LocalAI Assistant chat modality — adding admin tools to the in-process MCP server, editing skill prompts, keeping REST + MCP + skills in sync | | [.agents/backend-signing.md](.agents/backend-signing.md) | Backend OCI image signing (keyless cosign + sigstore-go) — producer-side CI setup, consumer-side gallery `verification:` block, strict mode (`LOCALAI_REQUIRE_BACKEND_INTEGRITY`), revocation via `not_before` | | [.agents/preparing-a-release.md](.agents/preparing-a-release.md) | Cutting a release: PR labels, `RELEASE_NOTES_vX.Y.Z.md`, the blog post under `website/content/blog/`, and the demo clips under `website/static/media/` | +| [.agents/distributed-state.md](.agents/distributed-state.md) | Features that keep runtime state — how they must behave with several frontends (syncstate, advisory-lock leaders, fakebus tests) | | [.impeccable.md](.impeccable.md) | Design context for UI/UX work — users, brand personality, aesthetic direction, and design principles | ## Quick Reference @@ -49,3 +50,4 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants] - **Backend OS coverage**: a new backend must target every OS it can build for, not just Linux. `.github/backend-matrix.yml` has two matrices — `include:` (Linux) and `includeDarwin:` (macOS / Apple Silicon). Most C/C++/GGML and many Python backends build on Darwin too — wire the `includeDarwin` entry + `backend/index.yaml` `metal:` entries, or say in the PR why an OS is unsupported. See the darwin checklist in [.agents/adding-backends.md](.agents/adding-backends.md). - **Gallery variant ranking**: a gallery entry can declare `variants` (alternative builds of the same weights), and LocalAI ranks the ones a host can run by engine preference first, size second. A new backend that should be preferred on some hardware must be listed in `engineNamePreferenceRules` in `pkg/system/capabilities.go`; the sibling `backendBuildTagPreferenceRules` speaks build tags rather than engine names, and using the wrong table matches nothing without erroring. See [.agents/adding-backends.md](.agents/adding-backends.md). - **UI**: The active UI is the React app in `core/http/react-ui/`. The older Alpine.js/HTML UI in `core/http/static/` is pending deprecation — all new UI work goes in the React UI +- **Distributed-aware state**: any feature that keeps runtime state (maps, caches, pins, schedulers, background loops) must choose shared (syncstate), single-runner (advisorylock), stateless, or documented per-instance behaviour for multi-frontend clusters. See [.agents/distributed-state.md](.agents/distributed-state.md). diff --git a/Makefile b/Makefile index e39c5dfa0..49812a8d7 100644 --- a/Makefile +++ b/Makefile @@ -1,5 +1,5 @@ # Disable parallel execution for backend builds -.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin +.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/localai-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin .NOTPARALLEL: backends/whisper-medusa .NOTPARALLEL: backends/funasr @@ -76,7 +76,7 @@ else GORELEASER=$(shell which goreleaser) endif -TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/... +TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/localai-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/... ## Coverage output and the committed baseline that CI compares against. ## The gate is strict: total coverage must never decrease (no tolerance). @@ -385,13 +385,14 @@ prepare-e2e: run-e2e-image: docker run -p 5390:8080 -e MODELS_PATH=/models -e THREADS=1 -e DEBUG=true -d --rm -v $(TEST_DIR):/models --name e2e-tests-$(RANDOM) localai-tests -test-e2e: build-mock-backend build-cloud-proxy-backend prepare-e2e run-e2e-image +test-e2e: build-mock-backend build-cloud-proxy-backend build-localai-proxy-backend prepare-e2e run-e2e-image @echo 'Running e2e tests' BUILD_TYPE=$(BUILD_TYPE) \ LOCALAI_API=http://$(E2E_BRIDGE_IP):5390 \ $(GOCMD) run github.com/onsi/ginkgo/v2/ginkgo --flake-attempts $(TEST_FLAKES) -v -r ./tests/e2e $(MAKE) clean-mock-backend $(MAKE) clean-cloud-proxy-backend + $(MAKE) clean-localai-proxy-backend $(MAKE) teardown-e2e docker rmi localai-tests @@ -1330,6 +1331,7 @@ BACKEND_PIPER = piper|golang|.|false|true BACKEND_LOCAL_STORE = local-store|golang|.|false|true BACKEND_VALKEY_STORE = valkey-store|golang|.|false|true BACKEND_CLOUD_PROXY = cloud-proxy|golang|.|false|true +BACKEND_LOCALAI_PROXY = localai-proxy|golang|.|false|true BACKEND_HUGGINGFACE = huggingface|golang|.|false|true BACKEND_SILERO_VAD = silero-vad|golang|.|false|true BACKEND_STABLEDIFFUSION_GGML = stablediffusion-ggml|golang|.|--progress=plain|true @@ -1436,6 +1438,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_PIPER))) $(eval $(call generate-docker-build-target,$(BACKEND_LOCAL_STORE))) $(eval $(call generate-docker-build-target,$(BACKEND_VALKEY_STORE))) $(eval $(call generate-docker-build-target,$(BACKEND_CLOUD_PROXY))) +$(eval $(call generate-docker-build-target,$(BACKEND_LOCALAI_PROXY))) $(eval $(call generate-docker-build-target,$(BACKEND_HUGGINGFACE))) $(eval $(call generate-docker-build-target,$(BACKEND_SILERO_VAD))) $(eval $(call generate-docker-build-target,$(BACKEND_STABLEDIFFUSION_GGML))) @@ -1511,7 +1514,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_SUPERTONIC))) docker-save-%: backend-images docker save local-ai-backend:$* -o backend-images/$*.tar -docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp +docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-localai-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp docker-build-backends: docker-build-whisper-medusa docker-build-backends: docker-build-funasr @@ -1531,6 +1534,12 @@ build-cloud-proxy-backend: protogen-go clean-cloud-proxy-backend: rm -f tests/e2e/mock-backend/cloud-proxy +build-localai-proxy-backend: protogen-go + $(GOCMD) build -o tests/e2e/mock-backend/localai-proxy ./backend/go/localai-proxy + +clean-localai-proxy-backend: + rm -f tests/e2e/mock-backend/localai-proxy + ######################################################## ### UI E2E Test Server ######################################################## diff --git a/backend/go/localai-proxy/Makefile b/backend/go/localai-proxy/Makefile new file mode 100644 index 000000000..94fc5cf48 --- /dev/null +++ b/backend/go/localai-proxy/Makefile @@ -0,0 +1,13 @@ +GOCMD=go + +# Packaged as a standalone gallery backend by backend/Dockerfile.golang. +localai-proxy: + CGO_ENABLED=0 $(GOCMD) build -ldflags "$(LD_FLAGS)" -tags "$(GO_TAGS)" -o localai-proxy ./ + +package: + bash package.sh + +build: localai-proxy package + +clean: + rm -f localai-proxy diff --git a/backend/go/localai-proxy/audio.go b/backend/go/localai-proxy/audio.go new file mode 100644 index 000000000..07e011209 --- /dev/null +++ b/backend/go/localai-proxy/audio.go @@ -0,0 +1,453 @@ +package main + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "io" + "net/url" + "path/filepath" + "strconv" + "strings" + "time" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +const transcriptionsPath = "/v1/audio/transcriptions" + +// ttsRequest is the body of LocalAI's /tts (schema.TTSRequest). +type ttsRequest struct { + Model string `json:"model"` + Input string `json:"input"` + Voice string `json:"voice,omitempty"` + Language string `json:"language,omitempty"` + Instructions string `json:"instructions,omitempty"` + Params map[string]string `json:"params,omitempty"` + Stream bool `json:"stream,omitempty"` +} + +func (p *LocalAIProxy) ttsRequest(req *pb.TTSRequest, stream bool) ttsRequest { + return ttsRequest{ + Model: p.model(""), + Input: req.GetText(), + Voice: req.GetVoice(), + Language: req.GetLanguage(), + Instructions: req.GetInstructions(), + Params: req.GetParams(), + Stream: stream, + } +} + +func (p *LocalAIProxy) TTS(req *pb.TTSRequest) error { + return p.postJSONToFile(context.Background(), "/tts", p.ttsRequest(req, false), req.GetDst()) +} + +// TTSStream forwards the upstream's chunked audio unchanged. That body is +// already what core expects from a streaming backend: a WAV header followed +// by PCM. out is closed on every path because the gRPC server drains it until +// closed and would otherwise hang. +func (p *LocalAIProxy) TTSStream(req *pb.TTSRequest, out chan []byte) error { + defer close(out) + resp, err := p.postStream(context.Background(), "/tts", p.ttsRequest(req, true)) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + buf := make([]byte, 32*1024) + for { + n, err := resp.Body.Read(buf) + if n > 0 { + // The reader reuses buf, so each chunk needs its own copy. + out <- append([]byte(nil), buf[:n]...) + } + if errors.Is(err, io.EOF) { + return nil + } + if err != nil { + // A cut-off stream is a failed synthesis, not a short one. + return transportError("/tts", err) + } + } +} + +// soundGenerationRequest is the body of /v1/sound-generation +// (schema.ElevenLabsSoundGenerationRequest). Pointers keep "unset" distinct +// from zero so the upstream model's defaults apply. +type soundGenerationRequest struct { + ModelID string `json:"model_id"` + Text string `json:"text"` + Duration *float32 `json:"duration_seconds,omitempty"` + Temperature *float32 `json:"prompt_influence,omitempty"` + DoSample *bool `json:"do_sample,omitempty"` + Think *bool `json:"think,omitempty"` + Caption string `json:"caption,omitempty"` + Lyrics string `json:"lyrics,omitempty"` + BPM *int32 `json:"bpm,omitempty"` + Keyscale string `json:"keyscale,omitempty"` + Language string `json:"language,omitempty"` + Timesignature string `json:"timesignature,omitempty"` + Instrumental *bool `json:"instrumental,omitempty"` +} + +// SoundGeneration refuses audio-conditioned requests: the REST endpoint has no +// field for a source clip, and dropping it would return unconditioned audio as +// if it were the answer. +func (p *LocalAIProxy) SoundGeneration(req *pb.SoundGenerationRequest) error { + if req.GetSrc() != "" { + return unimplemented("SoundGeneration with src") + } + body := soundGenerationRequest{ + ModelID: p.model(""), + Text: req.GetText(), + Duration: req.Duration, + Temperature: req.Temperature, + DoSample: req.Sample, + Think: req.Think, + Caption: req.GetCaption(), + Lyrics: req.GetLyrics(), + BPM: req.Bpm, + Keyscale: req.GetKeyscale(), + Language: req.GetLanguage(), + Timesignature: req.GetTimesignature(), + Instrumental: req.Instrumental, + } + return p.postJSONToFile(context.Background(), "/v1/sound-generation", body, req.GetDst()) +} + +// transcriptionForm builds the /v1/audio/transcriptions upload. Dst is the +// input audio (core names it that way). Language and translate are only sent +// when set so the upstream model config's defaults still apply; diarize is +// always sent because the upstream treats a missing field as true. +func (p *LocalAIProxy) transcriptionForm(req *pb.TranscriptRequest, stream bool) multipartForm { + f := url.Values{} + f.Set("model", p.model("")) + f.Set("diarize", strconv.FormatBool(req.GetDiarize())) + // verbose_json keeps segments and words; the default drops nothing today, + // but naming it keeps the reply shape pinned. + f.Set("response_format", "verbose_json") + if v := req.GetLanguage(); v != "" { + f.Set("language", v) + } + if req.GetTranslate() { + f.Set("translate", "true") + } + if v := req.GetPrompt(); v != "" { + f.Set("prompt", v) + } + if v := req.GetTemperature(); v != 0 { + f.Set("temperature", formatFloat(v)) + } + for _, g := range req.GetTimestampGranularities() { + f.Add("timestamp_granularities[]", g) + } + if stream { + f.Set("stream", "true") + } + return multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}} +} + +// transcriptionResult is TranscriptionResultSeconds: the REST API reports +// times in seconds, pb in nanoseconds (core reads them as time.Duration). +type transcriptionResult struct { + Text string `json:"text"` + Language string `json:"language"` + Duration float64 `json:"duration"` + Segments []transcriptionSegment `json:"segments"` +} + +type transcriptionSegment struct { + ID int32 `json:"id"` + Start float64 `json:"start"` + End float64 `json:"end"` + Text string `json:"text"` + Tokens []int32 `json:"tokens"` + Speaker string `json:"speaker"` + Words []struct { + Start float64 `json:"start"` + End float64 `json:"end"` + Text string `json:"text"` + } `json:"words"` +} + +// toProto drops the top-level words: core rebuilds them from segment words. +func (r transcriptionResult) toProto() *pb.TranscriptResult { + out := &pb.TranscriptResult{Text: r.Text, Language: r.Language, Duration: float32(r.Duration)} + for _, s := range r.Segments { + seg := &pb.TranscriptSegment{ + Id: s.ID, Start: nanos(s.Start), End: nanos(s.End), Text: s.Text, + Tokens: s.Tokens, Speaker: s.Speaker, + } + for _, w := range s.Words { + seg.Words = append(seg.Words, &pb.TranscriptWord{Start: nanos(w.Start), End: nanos(w.End), Text: w.Text}) + } + out.Segments = append(out.Segments, seg) + } + return out +} + +func nanos(seconds float64) int64 { + return int64(seconds * float64(time.Second)) +} + +func (p *LocalAIProxy) AudioTranscription(ctx context.Context, req *pb.TranscriptRequest) (pb.TranscriptResult, error) { + var resp transcriptionResult + if err := p.postForm(ctx, transcriptionsPath, p.transcriptionForm(req, false), &resp); err != nil { + return pb.TranscriptResult{}, err + } + return *resp.toProto(), nil +} + +// transcriptEvent covers every frame of the upstream transcription SSE stream. +type transcriptEvent struct { + Type string `json:"type"` + Delta string `json:"delta"` + Error *struct { + Message string `json:"message"` + } `json:"error"` + transcriptionResult +} + +// AudioTranscriptionStream maps upstream SSE frames to stream responses. out +// is closed on every path: the gRPC server drains it until closed. +func (p *LocalAIProxy) AudioTranscriptionStream(ctx context.Context, req *pb.TranscriptRequest, out chan *pb.TranscriptStreamResponse) error { + defer close(out) + resp, err := p.postMultipartStream(ctx, transcriptionsPath, p.transcriptionForm(req, true)) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + scanner := bufio.NewScanner(resp.Body) + // The done frame carries every segment of the recording. + scanner.Buffer(make([]byte, 0, 64*1024), 4<<20) + for scanner.Scan() { + payload, ok := strings.CutPrefix(scanner.Text(), "data:") + if !ok { + continue + } + payload = strings.TrimSpace(payload) + if payload == "" || payload == "[DONE]" { + continue + } + var ev transcriptEvent + if err := json.Unmarshal([]byte(payload), &ev); err != nil { + xlog.Debug("localai-proxy: skip malformed SSE frame", "path", transcriptionsPath, "error", err) + continue + } + switch ev.Type { + case "transcript.text.delta": + out <- &pb.TranscriptStreamResponse{Delta: ev.Delta} + case "transcript.text.done": + out <- &pb.TranscriptStreamResponse{FinalResult: ev.toProto()} + return nil + case "error": + msg := "unknown error" + if ev.Error != nil { + msg = ev.Error.Message + } + xlog.Warn("localai-proxy: upstream stream error", "path", transcriptionsPath, "error", msg) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", transcriptionsPath, msg) + } + } + if err := scanner.Err(); err != nil { + return transportError(transcriptionsPath, err) + } + // The upstream always ends with a done or error frame, so a stream that + // stops without one was cut off; returning nil would pass the partial + // deltas off as the whole transcript. + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream ended before the final transcript", transcriptionsPath) +} + +func (p *LocalAIProxy) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) { + f := url.Values{} + f.Set("model", p.model("")) + // verbose_json keeps per-segment text; core strips it again for callers + // that did not ask for it. + f.Set("response_format", "verbose_json") + if v := req.GetLanguage(); v != "" { + f.Set("language", v) + } + for name, v := range map[string]int32{ + "num_speakers": req.GetNumSpeakers(), "min_speakers": req.GetMinSpeakers(), "max_speakers": req.GetMaxSpeakers(), + } { + if v != 0 { + f.Set(name, strconv.Itoa(int(v))) + } + } + for name, v := range map[string]float32{ + "clustering_threshold": req.GetClusteringThreshold(), + "min_duration_on": req.GetMinDurationOn(), + "min_duration_off": req.GetMinDurationOff(), + } { + if v != 0 { + f.Set(name, formatFloat(v)) + } + } + if req.GetIncludeText() { + f.Set("include_text", "true") + } + + var resp struct { + Duration float64 `json:"duration"` + Language string `json:"language"` + NumSpeakers int32 `json:"num_speakers"` + Segments []struct { + ID int32 `json:"id"` + Speaker string `json:"speaker"` + Label string `json:"label"` + Start float32 `json:"start"` + End float32 `json:"end"` + Text string `json:"text"` + } `json:"segments"` + } + if err := p.postForm(context.Background(), "/v1/audio/diarization", + multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}, &resp); err != nil { + return pb.DiarizeResponse{}, err + } + var segments []*pb.DiarizeSegment + for _, s := range resp.Segments { + // The upstream renames speakers to SPEAKER_NN and keeps the backend's + // own id in label. Core renames again, so hand it the raw id: the + // caller then sees the same speakers and labels a direct call gives. + speaker := s.Label + if speaker == "" { + speaker = s.Speaker + } + segments = append(segments, &pb.DiarizeSegment{ + Id: s.ID, Start: s.Start, End: s.End, Speaker: speaker, Text: s.Text, + }) + } + return pb.DiarizeResponse{ + Segments: segments, NumSpeakers: resp.NumSpeakers, Duration: float32(resp.Duration), Language: resp.Language, + }, nil +} + +func (p *LocalAIProxy) VAD(req *pb.VADRequest) (pb.VADResponse, error) { + var resp struct { + Segments []struct { + Start float32 `json:"start"` + End float32 `json:"end"` + } `json:"segments"` + } + body := map[string]any{"model": p.model(""), "audio": req.GetAudio()} + if err := p.postJSON(context.Background(), "/v1/vad", body, &resp); err != nil { + return pb.VADResponse{}, err + } + var segments []*pb.VADSegment + for _, s := range resp.Segments { + segments = append(segments, &pb.VADSegment{Start: s.Start, End: s.End}) + } + return pb.VADResponse{Segments: segments}, nil +} + +func (p *LocalAIProxy) SoundDetection(ctx context.Context, req *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) { + f := url.Values{} + f.Set("model", p.model("")) + if v := req.GetTopK(); v != 0 { + f.Set("top_k", strconv.Itoa(int(v))) + } + if v := req.GetThreshold(); v != 0 { + f.Set("threshold", formatFloat(v)) + } + var resp struct { + Detections []struct { + Index int32 `json:"index"` + Label string `json:"label"` + Score float32 `json:"score"` + } `json:"detections"` + } + if err := p.postForm(ctx, "/v1/audio/classification", + multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetSrc()}}}, &resp); err != nil { + return nil, err + } + out := &pb.SoundDetectionResponse{} + for _, d := range resp.Detections { + out.Detections = append(out.Detections, &pb.SoundClass{Index: d.Index, Label: d.Label, Score: d.Score}) + } + return out, nil +} + +// stemsHeader names the other outputs of a separation run: the body carries +// one file, and the rest are served under /generated-audio/. +const stemsHeader = "X-Audio-Stems" + +// AudioTransform uploads the input (and reference) to /audio/transformations +// and writes the returned audio to Dst. SampleRate and Samples stay 0: the +// REST reply does not report them, and core only uses them for tracing. +func (p *LocalAIProxy) AudioTransform(req *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) { + f := url.Values{} + f.Set("model", p.model("")) + for k, v := range req.GetParams() { + f.Set("params["+k+"]", v) + } + form := multipartForm{fields: f, files: []formFile{{field: "audio", path: req.GetAudioPath()}}} + if ref := req.GetReferencePath(); ref != "" { + form.files = append(form.files, formFile{field: "reference", path: ref}) + } + + ctx := context.Background() + header, err := p.postMultipartToFile(ctx, "/audio/transformations", form, req.GetDst()) + if err != nil { + return nil, err + } + return &pb.AudioTransformResult{ + Dst: req.GetDst(), + ReferenceProvided: req.GetReferencePath() != "", + Stems: p.fetchStems(ctx, header.Get(stemsHeader), req.GetDst()), + }, nil +} + +// fetchStems downloads each stem the upstream names and writes it beside +// dst, where core looks for them. A stem that cannot be fetched is dropped +// with a warning rather than failing the call: the caller still gets the +// audio it asked for. +func (p *LocalAIProxy) fetchStems(ctx context.Context, header, dst string) []*pb.AudioTransformStem { + if header == "" { + return nil + } + var entries []struct { + Name string `json:"name"` + URL string `json:"url"` + } + if err := json.Unmarshal([]byte(header), &entries); err != nil { + xlog.Warn("localai-proxy: ignore malformed stems header", "error", err) + return nil + } + dir := filepath.Dir(dst) + prefix := strings.TrimSuffix(filepath.Base(dst), filepath.Ext(dst)) + var stems []*pb.AudioTransformStem + for _, e := range entries { + // Only files under /generated-audio/ are stems. Anything else is not + // a path this API hands out, so it is not fetched. + escaped, ok := strings.CutPrefix(e.URL, "/generated-audio/") + if !ok || e.Name == "" { + xlog.Warn("localai-proxy: skip stem with unexpected url", "name", e.Name, "url", e.URL) + continue + } + name, err := url.PathUnescape(escaped) + if err != nil || name == "" || name != filepath.Base(name) || name == "." || name == ".." { + xlog.Warn("localai-proxy: skip stem with unsafe file name", "name", e.Name, "url", e.URL) + continue + } + local := filepath.Join(dir, prefix+"-"+name) + if err := p.getToFile(ctx, e.URL, local); err != nil { + xlog.Warn("localai-proxy: stem download failed", "name", e.Name, "error", err) + continue + } + stems = append(stems, &pb.AudioTransformStem{Name: e.Name, Dst: local}) + } + return stems +} + +// formatFloat prints a float32 form value without float64 noise +// (0.2, not 0.20000000298023224). +func formatFloat(v float32) string { + return strconv.FormatFloat(float64(v), 'g', -1, 32) +} diff --git a/backend/go/localai-proxy/audio_test.go b/backend/go/localai-proxy/audio_test.go new file mode 100644 index 000000000..8f5013c7b --- /dev/null +++ b/backend/go/localai-proxy/audio_test.go @@ -0,0 +1,489 @@ +package main + +import ( + "context" + "fmt" + "net/http" + "os" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// writeInput writes content to a fresh temp file and returns its path, so a +// spec can check the upstream received exactly those bytes. +func writeInput(name, content string) string { + path := filepath.Join(GinkgoT().TempDir(), name) + Expect(os.WriteFile(path, []byte(content), 0o600)).To(Succeed()) + return path +} + +// drainBytes collects every chunk sent on ch until it is closed. +func drainBytes(ch chan []byte) <-chan [][]byte { + done := make(chan [][]byte, 1) + go func() { + var got [][]byte + for c := range ch { + got = append(got, c) + } + done <- got + }() + return done +} + +func drainTranscript(ch chan *pb.TranscriptStreamResponse) <-chan []*pb.TranscriptStreamResponse { + done := make(chan []*pb.TranscriptStreamResponse, 1) + go func() { + var got []*pb.TranscriptStreamResponse + for c := range ch { + got = append(got, c) + } + done <- got + }() + return done +} + +// cutMidStream answers 200, flushes body, then drops the connection without +// finishing the chunked encoding, as a crashed upstream would. +func cutMidStream(contentType, body string) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(body)) + w.(http.Flusher).Flush() + conn, _, err := w.(http.Hijacker).Hijack() + if err == nil { + _ = conn.Close() + } + } +} + +var _ = Describe("audio methods", func() { + var up *fakeUpstream + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.Close) + }) + + Describe("TTS", func() { + It("posts to /tts and writes the upstream audio to Dst", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-audio-bytes"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.wav") + lang := "it" + instr := "cheerful" + + Expect(p.TTS(&pb.TTSRequest{ + Text: "ciao", Model: "local-path", Dst: dst, Voice: "v1", Language: &lang, + Instructions: &instr, Params: map[string]string{"speed": "1.2"}, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-audio-bytes")) + req := up.last() + Expect(req.Path).To(Equal("/tts")) + Expect(req.JSON).To(Equal(map[string]any{ + "model": "remote-model", "input": "ciao", "voice": "v1", "language": "it", + "instructions": "cheerful", "params": map[string]any{"speed": "1.2"}, + })) + }) + + It("maps an upstream failure and leaves no partial file", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.wav") + err := p.TTS(&pb.TTSRequest{Text: "x", Dst: dst}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(dst).NotTo(BeAnExistingFile()) + }) + }) + + Describe("TTSStream", func() { + It("forwards the chunked WAV body in order and closes the channel", func() { + header := "RIFF" + string(make([]byte, 40)) + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "audio/wav") + w.WriteHeader(http.StatusOK) + for _, c := range []string{header, "pcm-1", "pcm-2"} { + _, _ = w.Write([]byte(c)) + w.(http.Flusher).Flush() + time.Sleep(10 * time.Millisecond) + } + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + out := make(chan []byte) + done := drainBytes(out) + Expect(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)).To(Succeed()) + + var chunks [][]byte + Eventually(done).Should(Receive(&chunks)) + Expect(chunks).NotTo(BeEmpty()) + Expect(string(chunks[0][:4])).To(Equal("RIFF")) + var all []byte + for _, c := range chunks { + all = append(all, c...) + } + Expect(string(all)).To(Equal(header + "pcm-1pcm-2")) + }) + + It("asks the upstream to stream", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF"}) + out := make(chan []byte) + done := drainBytes(out) + Expect(p.TTSStream(&pb.TTSRequest{Text: "hi", Voice: "v"}, out)).To(Succeed()) + Eventually(done).Should(Receive()) + Expect(up.last().JSON).To(HaveKeyWithValue("stream", true)) + Expect(up.last().JSON).To(HaveKeyWithValue("input", "hi")) + }) + + It("reports a mid-stream disconnect as Unavailable and still closes the channel", func() { + cut := newFakeUpstreamWithHandler(cutMidStream("audio/wav", "RIFFpartial")) + DeferCleanup(cut.Close) + p := loadProxy(cut, nil) + out := make(chan []byte) + done := drainBytes(out) + err := p.TTSStream(&pb.TTSRequest{Text: "hi"}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Eventually(done).Should(Receive()) + }) + + It("closes the channel when the upstream refuses the request", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusBadRequest, Body: "bad"}) + out := make(chan []byte) + done := drainBytes(out) + Expect(codeOf(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out))).To(Equal(codes.InvalidArgument)) + Eventually(done).Should(Receive()) + }) + }) + + Describe("SoundGeneration", func() { + It("posts the ElevenLabs body to /v1/sound-generation and writes Dst", func() { + p := loadProxy(up, nil) + up.script("/v1/sound-generation", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-sfx"}) + dst := filepath.Join(GinkgoT().TempDir(), "sfx.wav") + dur, temp, bpm := float32(4.5), float32(0.3), int32(120) + sample, instrumental := true, false + + Expect(p.SoundGeneration(&pb.SoundGenerationRequest{ + Text: "rain", Dst: dst, Duration: &dur, Temperature: &temp, Sample: &sample, + Bpm: &bpm, Caption: ptr("cap"), Lyrics: ptr("la"), Keyscale: ptr("C major"), + Language: ptr("en"), Timesignature: ptr("4/4"), Instrumental: &instrumental, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-sfx")) + Expect(up.last().JSON).To(Equal(map[string]any{ + "model_id": "remote-model", "text": "rain", "duration_seconds": 4.5, + "prompt_influence": 0.3, "do_sample": true, "bpm": float64(120), + "caption": "cap", "lyrics": "la", "keyscale": "C major", "language": "en", + "timesignature": "4/4", "instrumental": false, + })) + }) + + It("refuses audio-conditioned generation it cannot upload", func() { + p := loadProxy(up, nil) + err := p.SoundGeneration(&pb.SoundGenerationRequest{Text: "x", Src: ptr("/tmp/in.wav")}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + Expect(up.recorded()).To(BeEmpty()) + }) + }) + + Describe("AudioTranscription", func() { + It("uploads Dst with the request fields and maps the seconds-based result", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/transcriptions", map[string]any{ + "text": "hello world", "language": "en", "duration": 2.5, + "segments": []any{map[string]any{ + "id": 0, "start": 0.5, "end": 1.25, "text": "hello world", "tokens": []int{1, 2}, + "speaker": "SPEAKER_00", + "words": []any{map[string]any{"start": 0.5, "end": 0.75, "text": "hello"}}, + }}, + }) + in := writeInput("speech.wav", "RIFF-speech") + + res, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{ + Dst: in, Language: "en", Translate: true, Diarize: true, Prompt: "names: Ada", + Temperature: 0.2, TimestampGranularities: []string{"word"}, + }) + Expect(err).NotTo(HaveOccurred()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/audio/transcriptions")) + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-speech"})) + Expect(req.Fields).To(Equal(map[string]string{ + "model": "remote-model", "language": "en", "translate": "true", "diarize": "true", + "prompt": "names: Ada", "temperature": "0.2", "timestamp_granularities[]": "word", + "response_format": "verbose_json", + })) + + Expect(res.Text).To(Equal("hello world")) + Expect(res.Language).To(Equal("en")) + Expect(res.Duration).To(BeNumerically("==", 2.5)) + Expect(res.Segments).To(HaveLen(1)) + seg := res.Segments[0] + // pb carries nanoseconds (core reads them as time.Duration). + Expect(seg.Start).To(Equal(int64(500 * time.Millisecond))) + Expect(seg.End).To(Equal(int64(1250 * time.Millisecond))) + Expect(seg.Text).To(Equal("hello world")) + Expect(seg.Tokens).To(Equal([]int32{1, 2})) + Expect(seg.Speaker).To(Equal("SPEAKER_00")) + Expect(seg.Words).To(HaveLen(1)) + Expect(seg.Words[0].Start).To(Equal(int64(500 * time.Millisecond))) + Expect(seg.Words[0].Text).To(Equal("hello")) + }) + + It("sends diarize=false explicitly because the upstream defaults it on", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "x"}) + _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}) + Expect(err).NotTo(HaveOccurred()) + f := up.last().Fields + Expect(f).To(HaveKeyWithValue("diarize", "false")) + // Unset language and translate are left to the upstream model config. + Expect(f).NotTo(HaveKey("language")) + Expect(f).NotTo(HaveKey("translate")) + }) + + It("maps a 4xx to InvalidArgument", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad audio"}) + _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + }) + }) + + Describe("AudioTranscriptionStream", func() { + It("emits deltas, then the final result, and closes the channel", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}), + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "lo"}), + sseJSON(map[string]any{ + "type": "transcript.text.done", "text": "hello", "language": "en", "duration": 1.5, + "segments": []any{map[string]any{"id": 0, "start": 0.0, "end": 1.5, "text": "hello"}}, + }), + "[DONE]", + }}) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + + Expect(p.AudioTranscriptionStream(context.Background(), + &pb.TranscriptRequest{Dst: writeInput("a.wav", "RIFF-a"), Stream: true}, out)).To(Succeed()) + + var got []*pb.TranscriptStreamResponse + Eventually(done).Should(Receive(&got)) + Expect(got).To(HaveLen(3)) + Expect(got[0].Delta).To(Equal("hel")) + Expect(got[1].Delta).To(Equal("lo")) + final := got[2].FinalResult + Expect(final).NotTo(BeNil()) + Expect(final.Text).To(Equal("hello")) + Expect(final.Language).To(Equal("en")) + Expect(final.Duration).To(BeNumerically("==", 1.5)) + Expect(final.Segments).To(HaveLen(1)) + Expect(final.Segments[0].End).To(Equal(int64(1500 * time.Millisecond))) + + req := up.last() + Expect(req.Fields).To(HaveKeyWithValue("stream", "true")) + Expect(req.Files).To(HaveKeyWithValue("file", "RIFF-a")) + }) + + It("returns an upstream error event as Unavailable", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"type": "error", "error": map[string]any{"message": "decoder died"}}), + "[DONE]", + }}) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(err).To(MatchError(ContainSubstring("decoder died"))) + Eventually(done).Should(Receive()) + }) + + It("reports an upstream disconnect before the final result as Unavailable", func() { + frame := "data: " + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}) + "\n\n" + cut := newFakeUpstreamWithHandler(cutMidStream("text/event-stream", frame)) + DeferCleanup(cut.Close) + p := loadProxy(cut, nil) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + var got []*pb.TranscriptStreamResponse + Eventually(done).Should(Receive(&got)) + Expect(got).To(HaveLen(1)) + }) + + It("closes the channel when the local file is missing", func() { + p := loadProxy(up, nil) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: "/nonexistent.wav"}, out) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + Eventually(done).Should(Receive()) + Expect(up.recorded()).To(BeEmpty()) + }) + }) + + Describe("Diarize", func() { + It("uploads Dst with the tuning fields and maps the segments", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/diarization", map[string]any{ + "task": "diarize", "duration": 3.0, "language": "en", "num_speakers": 2, + "segments": []any{ + map[string]any{"id": 0, "speaker": "SPEAKER_00", "label": "spk_a", "start": 0.0, "end": 1.5, "text": "hi"}, + map[string]any{"id": 1, "speaker": "SPEAKER_01", "start": 1.5, "end": 3.0}, + }, + }) + res, err := p.Diarize(&pb.DiarizeRequest{ + Dst: writeInput("talk.wav", "RIFF-talk"), Language: "en", NumSpeakers: 2, MinSpeakers: 1, + MaxSpeakers: 3, ClusteringThreshold: 0.5, MinDurationOn: 0.1, MinDurationOff: 0.2, IncludeText: true, + }) + Expect(err).NotTo(HaveOccurred()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/audio/diarization")) + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-talk"})) + Expect(req.Fields).To(Equal(map[string]string{ + "model": "remote-model", "language": "en", "num_speakers": "2", "min_speakers": "1", + "max_speakers": "3", "clustering_threshold": "0.5", "min_duration_on": "0.1", + "min_duration_off": "0.2", "include_text": "true", "response_format": "verbose_json", + })) + + Expect(res.NumSpeakers).To(Equal(int32(2))) + Expect(res.Duration).To(BeNumerically("==", 3.0)) + Expect(res.Language).To(Equal("en")) + Expect(res.Segments).To(HaveLen(2)) + // The raw upstream label survives, so core's own normalisation + // yields the same speakers a direct call would. + Expect(res.Segments[0].Speaker).To(Equal("spk_a")) + Expect(res.Segments[0].Text).To(Equal("hi")) + Expect(res.Segments[1].Speaker).To(Equal("SPEAKER_01")) + Expect(res.Segments[1].Start).To(BeNumerically("==", 1.5)) + Expect(res.Segments[1].Id).To(Equal(int32(1))) + }) + }) + + Describe("VAD", func() { + It("posts the samples to /v1/vad and maps the segments", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/vad", map[string]any{"segments": []any{map[string]any{"start": 0.25, "end": 1.5}}}) + res, err := p.VAD(&pb.VADRequest{Audio: []float32{0.5, -0.25}}) + Expect(err).NotTo(HaveOccurred()) + Expect(up.last().JSON).To(Equal(map[string]any{"model": "remote-model", "audio": []any{0.5, -0.25}})) + Expect(res.Segments).To(HaveLen(1)) + Expect(res.Segments[0].Start).To(BeNumerically("==", 0.25)) + Expect(res.Segments[0].End).To(BeNumerically("==", 1.5)) + }) + }) + + Describe("SoundDetection", func() { + It("uploads Src to /v1/audio/classification and maps the detections", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/classification", map[string]any{"model": "remote-model", "detections": []any{ + map[string]any{"index": 74, "label": "Dog", "score": 0.9}, + map[string]any{"index": 0, "label": "Speech", "score": 0.4}, + }}) + res, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{ + Src: writeInput("bark.wav", "RIFF-bark"), TopK: 5, Threshold: 0.25, + }) + Expect(err).NotTo(HaveOccurred()) + req := up.last() + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-bark"})) + Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "top_k": "5", "threshold": "0.25"})) + Expect(res.Detections).To(HaveLen(2)) + Expect(res.Detections[0].Label).To(Equal("Dog")) + Expect(res.Detections[0].Index).To(Equal(int32(74))) + Expect(res.Detections[0].Score).To(BeNumerically("~", 0.9, 1e-6)) + }) + }) + + Describe("AudioTransform", func() { + It("uploads audio and reference with params and writes the result to Dst", func() { + p := loadProxy(up, nil) + up.script("/audio/transformations", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-clean"}) + dst := filepath.Join(GinkgoT().TempDir(), "transform.wav") + + res, err := p.AudioTransform(&pb.AudioTransformRequest{ + AudioPath: writeInput("mic.wav", "RIFF-mic"), ReferencePath: writeInput("ref.wav", "RIFF-ref"), + Dst: dst, Params: map[string]string{"noise_gate": "true"}, + }) + Expect(err).NotTo(HaveOccurred()) + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-clean")) + Expect(res.Dst).To(Equal(dst)) + Expect(res.ReferenceProvided).To(BeTrue()) + Expect(res.Stems).To(BeEmpty()) + + req := up.last() + Expect(req.Path).To(Equal("/audio/transformations")) + Expect(req.Files).To(Equal(map[string]string{"audio": "RIFF-mic", "reference": "RIFF-ref"})) + Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "params[noise_gate]": "true"})) + }) + + It("fetches the stems the upstream names and writes them beside Dst", func() { + p := loadProxy(up, nil) + stems := `[{"name":"vocals","url":"/generated-audio/sep%20vocals.wav"},{"name":"evil","url":"/etc/passwd"}]` + up.script("/audio/transformations", scriptedResponse{ + Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-drums", + Header: map[string]string{"X-Audio-Stems": stems}, + }) + up.script("/generated-audio/sep vocals.wav", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-vocals"}) + dir := GinkgoT().TempDir() + dst := filepath.Join(dir, "transform.wav") + + res, err := p.AudioTransform(&pb.AudioTransformRequest{AudioPath: writeInput("song.wav", "RIFF-song"), Dst: dst}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.ReferenceProvided).To(BeFalse()) + Expect(res.Stems).To(HaveLen(1)) + Expect(res.Stems[0].Name).To(Equal("vocals")) + // Core keeps only stems that are direct children of Dst's directory. + Expect(filepath.Dir(res.Stems[0].Dst)).To(Equal(dir)) + got, err := os.ReadFile(res.Stems[0].Dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-vocals")) + + var paths []string + for _, r := range up.recorded() { + paths = append(paths, r.Method+" "+r.Path) + } + Expect(paths).To(Equal([]string{"POST /audio/transformations", "GET /generated-audio/sep vocals.wav"})) + }) + }) + + It("names the upload file after the local file", func() { + // The upstream saves the upload under its base name and some handlers + // pick the decoder by extension, so the name must reach it intact. + var name string + named := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + _, fh, err := r.FormFile("file") + if err == nil { + name = fh.Filename + } + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprint(w, `{"detections":[]}`) + }) + DeferCleanup(named.Close) + p := loadProxy(named, nil) + _, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{Src: writeInput("clip.mp3", "ID3")}) + Expect(err).NotTo(HaveOccurred()) + Expect(name).To(Equal("clip.mp3")) + }) +}) + +func ptr[T any](v T) *T { return &v } diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go new file mode 100644 index 000000000..6b9ecfa18 --- /dev/null +++ b/backend/go/localai-proxy/client.go @@ -0,0 +1,376 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "mime/multipart" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +// maxErrorBody caps the upstream body quoted in an error. It keeps gRPC +// status messages small while leaving room for LocalAI's JSON error text, +// which failover scans for request errors such as context overflows. +const maxErrorBody = 500 + +// postJSON sends body as JSON to path and decodes a 2xx JSON reply into out +// (skipped when out is nil). The request_timeout_seconds limit applies. +func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any) 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") + return p.do(req, path, out) +} + +// postMultipart sends fields and, when fileField is set, the local file at +// filePath as a multipart form to path, decoding a 2xx JSON reply into out +// (skipped when out is nil). Core hands audio and images to backends as local +// 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 + } + ctx, cancel := withTimeout(ctx, cfg) + defer cancel() + req, err := p.newMultipartRequest(ctx, cfg, path, form) + if err != nil { + return err + } + return p.do(req, path, out) +} + +// 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() + } + } + 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 + } + if _, err := io.Copy(part, file); err != nil { + return err + } + } + return mw.Close() +} + +// postStream sends body as JSON to path and returns the open response of a +// 2xx reply; the caller must close its body. No request_timeout_seconds +// limit applies: streams legitimately outlast it, and ctx bounds them. +func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (*http.Response, error) { + cfg, err := p.config() + if err != nil { + return nil, err + } + payload, err := json.Marshal(body) + if err != nil { + return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err) + } + 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) + } + if resp.StatusCode < 200 || resp.StatusCode > 299 { + defer func() { _ = resp.Body.Close() }() + return nil, statusError(path, resp) + } + return resp, nil +} + +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) + } + if cfg.apiKey != "" { + req.Header.Set("Authorization", "Bearer "+cfg.apiKey) + } + return req, nil +} + +// do runs req and decodes a 2xx JSON reply into out. +func (p *LocalAIProxy) do(req *http.Request, path string, out any) error { + resp, err := p.client.Do(req) + if err != nil { + return transportError(path, err) + } + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode < 200 || resp.StatusCode > 299 { + return statusError(path, resp) + } + if out == nil { + _, _ = io.Copy(io.Discard, resp.Body) + return nil + } + if err := json.NewDecoder(resp.Body).Decode(out); err != nil { + if ctxErr := req.Context().Err(); ctxErr != nil { + return transportError(path, ctxErr) + } + return status.Errorf(codes.Internal, "localai-proxy: decode %s response: %v", path, err) + } + return nil +} + +func withTimeout(ctx context.Context, cfg *proxyConfig) (context.Context, context.CancelFunc) { + if cfg.timeout > 0 { + return context.WithTimeout(ctx, cfg.timeout) + } + return context.WithCancel(ctx) +} + +// transportError maps a failed round trip to a gRPC status. A dead or +// unreachable upstream is Unavailable so failover moves to the next target. +func transportError(path string, err error) error { + code := codes.Unavailable + switch { + case errors.Is(err, context.DeadlineExceeded): + code = codes.DeadlineExceeded + case errors.Is(err, context.Canceled): + code = codes.Canceled + } + xlog.Warn("localai-proxy: upstream request failed", "path", path, "error", err) + return status.Errorf(code, "localai-proxy: upstream %s: %v", path, err) +} + +// statusError maps a non-2xx upstream reply to a gRPC status. 5xx means the +// upstream is unhealthy (Unavailable, so failover retries elsewhere); 4xx +// means the request itself is wrong (InvalidArgument, so failover does not +// trip a healthy target over a client error). 429 is the exception, see +// below. 501 is the upstream saying it cannot serve this kind of request, +// which failover treats as a capability gap, like our own Unimplemented +// methods. +func statusError(path string, resp *http.Response) error { + raw, _ := io.ReadAll(io.LimitReader(resp.Body, maxErrorBody+1)) + msg := strings.TrimSpace(string(raw)) + if len(msg) > maxErrorBody { + msg = msg[:maxErrorBody] + "..." + } + // gRPC refuses to send a status message that is not valid UTF-8, and the + // cut above may split a rune. + msg = strings.ToValidUTF8(msg, "") + + var code codes.Code + switch { + case resp.StatusCode == http.StatusNotImplemented: + code = codes.Unimplemented + case resp.StatusCode >= 500: + code = codes.Unavailable + case resp.StatusCode == http.StatusTooManyRequests: + // Rate limited: the request is fine but the upstream is out of + // capacity. Failover retries it on the next target and trips this + // one, so traffic moves off it for a while. + code = codes.ResourceExhausted + case resp.StatusCode >= 400: + code = codes.InvalidArgument + default: + // A 1xx/3xx here means a misbehaving upstream (redirects are refused + // by the client), not a bad request. + code = codes.Unavailable + } + 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() }() + + // #nosec G304 -- dst is the output path core chose for this call (generated content dir), never a caller-supplied path + f, err := os.Create(filepath.Clean(dst)) + if err != nil { + return nil, status.Errorf(codes.Internal, "localai-proxy: create %s: %v", dst, err) + } + _, copyErr := io.Copy(f, resp.Body) + closeErr := f.Close() + if copyErr != nil || closeErr != nil { + _ = os.Remove(dst) + if copyErr != nil { + if ctxErr := req.Context().Err(); ctxErr != nil { + return nil, transportError(path, ctxErr) + } + return nil, transportError(path, copyErr) + } + return nil, status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, closeErr) + } + return resp.Header, nil +} diff --git a/backend/go/localai-proxy/fake_upstream_test.go b/backend/go/localai-proxy/fake_upstream_test.go new file mode 100644 index 000000000..933e012be --- /dev/null +++ b/backend/go/localai-proxy/fake_upstream_test.go @@ -0,0 +1,164 @@ +package main + +import ( + "encoding/json" + "io" + "mime" + "mime/multipart" + "net/http" + "net/http/httptest" + "strings" + "sync" + + . "github.com/onsi/gomega" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// recordedRequest is what the fake upstream saw for one call. JSON bodies land +// in JSON; multipart bodies land in Fields and Files (field name to content). +type recordedRequest struct { + Method string + Path string + Auth string + JSON map[string]any + Fields map[string]string + Files map[string]string +} + +// scriptedResponse is the reply for one path. SSE, when set, is written as +// "data: " events and wins over Body. Header adds response headers. +type scriptedResponse struct { + Status int + ContentType string + Header map[string]string + Body string + SSE []string +} + +// fakeUpstream stands in for a remote LocalAI: it records every request and +// answers each path with the response scripted for it (404 otherwise). +type fakeUpstream struct { + *httptest.Server + + mu sync.Mutex + requests []recordedRequest + responses map[string]scriptedResponse +} + +func newFakeUpstream() *fakeUpstream { + f := &fakeUpstream{responses: map[string]scriptedResponse{}} + f.Server = httptest.NewServer(http.HandlerFunc(f.serve)) + return f +} + +// newFakeUpstreamWithHandler serves every request with h instead of the +// scripted responses, for tests that need to control timing. +func newFakeUpstreamWithHandler(h http.HandlerFunc) *fakeUpstream { + return &fakeUpstream{Server: httptest.NewServer(h), responses: map[string]scriptedResponse{}} +} + +func (f *fakeUpstream) script(path string, r scriptedResponse) { + f.mu.Lock() + defer f.mu.Unlock() + f.responses[path] = r +} + +// replyJSON scripts a 200 JSON response for path. +func (f *fakeUpstream) replyJSON(path string, body any) { + raw, err := json.Marshal(body) + Expect(err).NotTo(HaveOccurred()) + f.script(path, scriptedResponse{Status: http.StatusOK, ContentType: "application/json", Body: string(raw)}) +} + +func (f *fakeUpstream) recorded() []recordedRequest { + f.mu.Lock() + defer f.mu.Unlock() + return append([]recordedRequest(nil), f.requests...) +} + +// last returns the single most recent request, failing when none arrived. +func (f *fakeUpstream) last() recordedRequest { + reqs := f.recorded() + ExpectWithOffset(1, reqs).NotTo(BeEmpty(), "upstream received no request") + return reqs[len(reqs)-1] +} + +func (f *fakeUpstream) serve(w http.ResponseWriter, r *http.Request) { + rec := recordedRequest{Method: r.Method, Path: r.URL.Path, Auth: r.Header.Get("Authorization")} + mediaType, params, _ := mime.ParseMediaType(r.Header.Get("Content-Type")) + switch { + case mediaType == "multipart/form-data": + rec.Fields, rec.Files = map[string]string{}, map[string]string{} + mr := multipart.NewReader(r.Body, params["boundary"]) + for { + part, err := mr.NextPart() + if err != nil { + break + } + data, _ := io.ReadAll(part) + if part.FileName() != "" { + rec.Files[part.FormName()] = string(data) + } else { + rec.Fields[part.FormName()] = string(data) + } + } + default: + raw, _ := io.ReadAll(r.Body) + if len(raw) > 0 { + _ = json.Unmarshal(raw, &rec.JSON) + } + } + + f.mu.Lock() + f.requests = append(f.requests, rec) + resp, ok := f.responses[r.URL.Path] + f.mu.Unlock() + + if !ok { + http.Error(w, "no scripted response for "+r.URL.Path, http.StatusNotFound) + return + } + if resp.SSE != nil { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + flusher, _ := w.(http.Flusher) + for _, frame := range resp.SSE { + _, _ = io.WriteString(w, "data: "+frame+"\n\n") + if flusher != nil { + flusher.Flush() + } + } + return + } + 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) +} + +// loadProxy returns a proxy loaded against the fake upstream with the given +// proxy options merged over sane defaults. +func loadProxy(f *fakeUpstream, mutate func(*pb.ModelOptions)) *LocalAIProxy { + opts := &pb.ModelOptions{ + Model: "local-name", + Proxy: &pb.ProxyOptions{UpstreamUrl: f.URL + "/", UpstreamModel: "remote-model"}, + } + if mutate != nil { + mutate(opts) + } + p := NewLocalAIProxy() + ExpectWithOffset(1, p.Load(opts)).To(Succeed()) + return p +} + +// sseJSON marshals v for use as one SSE frame. +func sseJSON(v any) string { + raw, err := json.Marshal(v) + Expect(err).NotTo(HaveOccurred()) + return strings.TrimSpace(string(raw)) +} diff --git a/backend/go/localai-proxy/live.go b/backend/go/localai-proxy/live.go new file mode 100644 index 000000000..8787ac0fd --- /dev/null +++ b/backend/go/localai-proxy/live.go @@ -0,0 +1,479 @@ +package main + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "math" + "net/http" + "net/url" + "strings" + "time" + + "github.com/gorilla/websocket" + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + "github.com/mudler/LocalAI/pkg/grpc/grpcerrors" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +const ( + realtimePath = "/v1/realtime" + + // defaultLiveSampleRate is what TranscriptLiveConfig.sample_rate 0 means. + defaultLiveSampleRate = 16000 + + // finalWait bounds how long a closing session waits for an utterance the + // upstream still has in flight. The upstream transcribes only after its + // VAD sees the speech end, so a slow model can finish after the client + // stops sending; waiting forever would pin the gRPC call on a hung + // upstream. + finalWait = 5 * time.Second + + // commitGrace is how long a speech_stopped waits for its + // input_audio_buffer.committed. The upstream sends the two back to back, + // so a stop with no commit inside this window is a discarded turn and + // must not hold a closing session for the full finalWait. + commitGrace = 500 * time.Millisecond + + // maxBacklogSeconds caps the audio held before the ready ack. Callers + // wait for the ack before streaming, so more than this means a client + // that ignores the contract, and holding it unbounded while a cold + // upstream loads models would grow memory without limit. + maxBacklogSeconds = 5 + + // liveWriteTimeout turns an upstream that stops reading into an error + // instead of a call blocked on a full socket buffer. + liveWriteTimeout = 10 * time.Second + + // liveHandshakeTimeout bounds the WebSocket upgrade only. The upstream + // warms the pipeline models before it sends session.created, which can + // take minutes on a cold box, so the setup phase after the upgrade has + // its own, longer bound. + liveHandshakeTimeout = 30 * time.Second + + // defaultLiveSetupTimeout bounds the session setup (upgrade to + // session.updated) when request_timeout_seconds is unset. Core waits for + // the ready ack with a plain Recv, so without a bound a hung upstream + // would hold the call, and the failover that should move the stage to + // the next target, forever. It is generous because a cold upstream loads + // the pipeline's VAD and transcription models before it answers. + defaultLiveSetupTimeout = 3 * time.Minute +) + +// liveSetupTimeout is defaultLiveSetupTimeout, as a variable so tests can +// exercise the bound without waiting minutes. +var liveSetupTimeout = defaultLiveSetupTimeout + +// realtimeEvent holds the fields the bridge reads from any upstream server +// event; unrelated events decode into it harmlessly. +type realtimeEvent struct { + Type string `json:"type"` + // ItemID is set on committed, delta, completed and failed events; the + // upstream uses the committed turn's id for its transcription events. + ItemID string `json:"item_id"` + Delta string `json:"delta"` + Transcript string `json:"transcript"` + Error *struct { + Message string `json:"message"` + Code string `json:"code"` + } `json:"error"` +} + +// AudioTranscriptionLive bridges a live transcription session to the +// upstream's /v1/realtime transcription session, so a realtime pipeline on +// this box can use a remote LocalAI's streaming ASR. The upstream runs its own +// server VAD and transcribes per utterance: each completed utterance becomes +// Delta (whatever the streamed deltas did not already carry) plus Eou. Word +// timings and Eob have no upstream counterpart and stay empty. +// +// Contract (pkg/grpc/server.go): this method closes out, in closes when the +// client half-closes, and errors are returned at once because the caller +// blocks on the first Recv for the ready ack. +func (p *LocalAIProxy) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest, out chan<- *pb.TranscriptLiveResponse) error { + defer close(out) + + cfg, err := p.config() + if err != nil { + return err + } + if cfg.realtimePipeline == "" { + return grpcerrors.LiveTranscriptionUnsupported(backendName, "set the realtime_pipeline backend option") + } + + first, ok := <-in + if !ok { + return nil // the caller closed without sending anything + } + lc := first.GetConfig() + if lc == nil { + return status.Error(codes.InvalidArgument, "localai-proxy: the first live transcription message must carry a config") + } + rate := int(lc.GetSampleRate()) + if rate == 0 { + rate = defaultLiveSampleRate + } + + conn, err := p.dialRealtime(cfg) + if err != nil { + return err + } + defer func() { _ = conn.Close() }() + + s := &liveSession{ + conn: conn, + out: out, + pipeline: cfg.realtimePipeline, + language: lc.GetLanguage(), + rate: rate, + sent: map[string]string{}, + committed: map[string]bool{}, + events: make(chan realtimeEvent), + done: make(chan struct{}), + } + // Closing done releases the reader if it is blocked handing over an + // event; closing conn (deferred above, runs after this) unblocks its read. + defer close(s.done) + go s.readLoop() + + setup := cfg.timeout + if setup <= 0 { + setup = liveSetupTimeout + } + return s.run(in, setup) +} + +// dialRealtime opens the upstream WebSocket. A refused upgrade is mapped like +// any other upstream HTTP reply, so a 5xx trips failover and a 4xx does not. +func (p *LocalAIProxy) dialRealtime(cfg *proxyConfig) (*websocket.Conn, error) { + u, err := url.Parse(cfg.base + realtimePath) + if err != nil { + return nil, status.Errorf(codes.Internal, "localai-proxy: build %s URL: %v", realtimePath, err) + } + // Load only accepts http(s) bases. + if u.Scheme == "https" { + u.Scheme = "wss" + } else { + u.Scheme = "ws" + } + u.RawQuery = url.Values{"model": {cfg.realtimePipeline}}.Encode() + + header := http.Header{} + if cfg.apiKey != "" { + header.Set("Authorization", "Bearer "+cfg.apiKey) + } + dialer := websocket.Dialer{ + Proxy: http.ProxyFromEnvironment, + HandshakeTimeout: liveHandshakeTimeout, + } + conn, resp, err := dialer.Dial(u.String(), header) + if err != nil { + if resp != nil && resp.StatusCode != http.StatusSwitchingProtocols { + defer func() { _ = resp.Body.Close() }() + return nil, statusError(realtimePath, resp) + } + return nil, transportError(realtimePath, err) + } + return conn, nil +} + +// liveSession is the state of one bridged session. Only run touches it; the +// reader goroutine only hands events over. +type liveSession struct { + conn *websocket.Conn + out chan<- *pb.TranscriptLiveResponse + pipeline string + language string + rate int + + sent map[string]string // text already sent as deltas, per upstream item + final []string // completed transcripts, in order + + // In-flight tracking, so a closing session waits only for an utterance + // the upstream will still transcribe. The upstream VAD emits + // speech_started, then either speech_stopped plus committed (a turn it + // will transcribe) or nothing at all (a turn it discarded as no speech). + speaking bool // between speech_started and speech_stopped + stopping bool // speech_stopped seen, its committed not yet + committed map[string]bool // committed items not yet completed + + events chan realtimeEvent + readErr error // set before events is closed + done chan struct{} +} + +// readLoop decodes upstream events until the socket fails. It is the only +// reader, as gorilla/websocket requires. +func (s *liveSession) readLoop() { + defer close(s.events) + for { + _, msg, err := s.conn.ReadMessage() + if err != nil { + s.readErr = err + return + } + var ev realtimeEvent + if err := json.Unmarshal(msg, &ev); err != nil { + xlog.Debug("localai-proxy: skipping undecodable realtime event", "error", err) + continue + } + select { + case s.events <- ev: + case <-s.done: + return + } + } +} + +// run drives the session: setup, then audio forwarding and event mapping, +// then the drain once the client closes its side. It is also the only writer +// on the socket, so frames never interleave. +func (s *liveSession) run(in <-chan *pb.TranscriptLiveRequest, setupTimeout time.Duration) error { + var ( + created, ready bool + inClosed bool + backlog []*pb.TranscriptLiveRequest + backlogSamples int + drainTimer <-chan time.Time + graceTimer <-chan time.Time + ) + setup := time.NewTimer(setupTimeout) + defer setup.Stop() + setupTimer := setup.C + drain := time.NewTimer(finalWait) + drain.Stop() + defer drain.Stop() + grace := time.NewTimer(commitGrace) + grace.Stop() + defer grace.Stop() + + for { + select { + case req, ok := <-in: + if !ok { + in, inClosed = nil, true + if !ready { + // The caller gave up waiting for the ready ack (core + // closes its send side on failure or cancel). Canceled + // rather than nil: there is no session to report as + // complete, and failover neither retries nor trips a + // target on a canceled call. + return status.Error(codes.Canceled, "localai-proxy: live transcription closed before the upstream session was ready") + } + if s.inFlight() { + drain.Reset(finalWait) + drainTimer = drain.C + } + } else if !ready { + // Hold audio sent before the ready ack rather than drop it, + // up to a bound. + backlogSamples += len(req.GetAudio().GetPcm()) + if backlogSamples > maxBacklogSeconds*s.rate { + return status.Errorf(codes.InvalidArgument, + "localai-proxy: more than %d s of audio sent before the live transcription session was ready", maxBacklogSeconds) + } + backlog = append(backlog, req) + } else if err := s.forward(req); err != nil { + return err + } + + case ev, ok := <-s.events: + if !ok { + return s.upstreamGone() + } + switch ev.Type { + case "session.created": + if created { + continue + } + created = true + if err := s.write(s.sessionUpdate()); err != nil { + return err + } + case "session.updated": + if ready || !created { + continue + } + ready = true + setupTimer = nil + s.out <- &pb.TranscriptLiveResponse{Ready: true} + for _, req := range backlog { + if err := s.forward(req); err != nil { + return err + } + } + backlog = nil + case "error": + return s.upstreamError("error", ev) + case "conversation.item.input_audio_transcription.failed": + return s.upstreamError("transcription failed", ev) + case "input_audio_buffer.speech_started": + s.speaking, s.stopping = true, false + case "input_audio_buffer.speech_stopped": + if s.speaking { + s.speaking, s.stopping = false, true + grace.Reset(commitGrace) + graceTimer = grace.C + } + case "input_audio_buffer.committed": + s.stopping, graceTimer = false, nil + if ev.ItemID != "" { + s.committed[ev.ItemID] = true + } + case "conversation.item.input_audio_transcription.delta": + s.delta(ev) + case "conversation.item.input_audio_transcription.completed": + s.completed(ev) + } + + case <-graceTimer: + // A stop the upstream never committed: the turn was discarded. + s.stopping, graceTimer = false, nil + + case <-setupTimer: + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s did not set up the transcription session within %s", realtimePath, setupTimeout) + + case <-drainTimer: + xlog.Warn("localai-proxy: upstream did not finish the last utterance in time; finalizing without it", + "pipeline", s.pipeline, "wait", finalWait) + return s.finish() + } + + if inClosed && !s.inFlight() { + return s.finish() + } + } +} + +// inFlight reports an utterance the upstream is still expected to +// transcribe. +func (s *liveSession) inFlight() bool { + return s.speaking || s.stopping || len(s.committed) > 0 +} + +func (s *liveSession) sessionUpdate() map[string]any { + return map[string]any{ + "type": "session.update", + "session": map[string]any{ + "type": "transcription", + "audio": map[string]any{ + "input": map[string]any{ + "format": map[string]any{"type": "audio/pcm", "rate": s.rate}, + "transcription": map[string]any{"model": s.pipeline, "language": s.language}, + "turn_detection": map[string]any{"type": "server_vad"}, + }, + }, + }, + } +} + +// forward sends one client message upstream. The upstream session rate is +// fixed at setup, so a second config cannot be honoured. +func (s *liveSession) forward(req *pb.TranscriptLiveRequest) error { + if req.GetConfig() != nil { + return status.Error(codes.InvalidArgument, "localai-proxy: a live transcription config is only accepted as the first message") + } + pcm := req.GetAudio().GetPcm() + if len(pcm) == 0 { + return nil + } + return s.write(map[string]any{ + "type": "input_audio_buffer.append", + "audio": base64.StdEncoding.EncodeToString(pcm16LE(pcm)), + }) +} + +// pcm16LE converts float samples in [-1, 1] to the little-endian PCM16 the +// realtime API takes. Out-of-range samples are clipped rather than wrapped. +func pcm16LE(pcm []float32) []byte { + buf := make([]byte, len(pcm)*2) + for i, f := range pcm { + v := float64(f) + if math.IsNaN(v) { + // A NaN would convert to an arbitrary int16; silence is the + // only neutral value. + v = 0 + } + v = math.Max(-1, math.Min(1, v)) + // #nosec G115 -- two's-complement reinterpretation for little-endian PCM16 encoding, value range already clamped + binary.LittleEndian.PutUint16(buf[i*2:], uint16(int16(v*math.MaxInt16))) + } + return buf +} + +func (s *liveSession) delta(ev realtimeEvent) { + if ev.Delta == "" { + return + } + s.sent[ev.ItemID] += ev.Delta + s.out <- &pb.TranscriptLiveResponse{Delta: ev.Delta} +} + +// completed ends one utterance. Deltas only arrive when the upstream pipeline +// streams its transcription, so the completion carries whatever text the +// deltas did not; if the final transcript diverged from them there is no way +// to retract, and FinalResult carries the authoritative text. +func (s *liveSession) completed(ev realtimeEvent) { + sent := s.sent[ev.ItemID] + delete(s.sent, ev.ItemID) + rest, ok := strings.CutPrefix(ev.Transcript, sent) + if !ok { + rest = "" + } + delete(s.committed, ev.ItemID) + if t := strings.TrimSpace(ev.Transcript); t != "" { + s.final = append(s.final, t) + } + s.out <- &pb.TranscriptLiveResponse{Delta: rest, Eou: true} +} + +// finish sends the final transcript and closes the upstream session cleanly. +func (s *liveSession) finish() error { + s.out <- &pb.TranscriptLiveResponse{FinalResult: &pb.TranscriptResult{Text: strings.Join(s.final, " ")}} + _ = s.conn.WriteControl(websocket.CloseMessage, + websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second)) + return nil +} + +func (s *liveSession) write(v any) error { + if err := s.conn.SetWriteDeadline(time.Now().Add(liveWriteTimeout)); err != nil { + return s.writeFailed(err) + } + if err := s.conn.WriteJSON(v); err != nil { + return s.writeFailed(err) + } + return nil +} + +func (s *liveSession) writeFailed(err error) error { + xlog.Warn("localai-proxy: realtime write failed", "pipeline", s.pipeline, "error", err) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s: %v", realtimePath, err) +} + +// upstreamGone maps a socket that stopped delivering events. It is +// Unavailable even for a clean close: the client did not end the session, so +// the upstream dropped it, and failover should reopen on the next target. +func (s *liveSession) upstreamGone() error { + err := s.readErr + if err == nil { + err = errors.New("connection closed") + } + xlog.Warn("localai-proxy: realtime upstream disconnected", "pipeline", s.pipeline, "error", err) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s disconnected: %v", realtimePath, err) +} + +func (s *liveSession) upstreamError(what string, ev realtimeEvent) error { + msg := "no details" + if ev.Error != nil && ev.Error.Message != "" { + msg = ev.Error.Message + if ev.Error.Code != "" { + msg = fmt.Sprintf("%s (%s)", msg, ev.Error.Code) + } + } + xlog.Warn("localai-proxy: realtime upstream reported an error", "pipeline", s.pipeline, "event", what, "error", msg) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s %s: %s", realtimePath, what, msg) +} diff --git a/backend/go/localai-proxy/live_test.go b/backend/go/localai-proxy/live_test.go new file mode 100644 index 000000000..c35b4463b --- /dev/null +++ b/backend/go/localai-proxy/live_test.go @@ -0,0 +1,533 @@ +package main + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "math" + "net/http" + "runtime" + "strings" + "time" + + "github.com/gorilla/websocket" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + "github.com/mudler/LocalAI/pkg/grpc/grpcerrors" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// wsUpstream is a fake upstream /v1/realtime endpoint. Each spec scripts the +// server side of the session in script, which runs on the upgraded socket. +func wsUpstream(script func(c *websocket.Conn, r *http.Request)) *fakeUpstream { + upgrader := websocket.Upgrader{} + return newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/realtime" { + http.NotFound(w, r) + return + } + c, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer func() { _ = c.Close() }() + script(c, r) + }) +} + +func wsSend(c *websocket.Conn, v any) { + defer GinkgoRecover() + Expect(c.WriteJSON(v)).To(Succeed()) +} + +// wsRecv reads the next client event; ok is false once the client is gone. +func wsRecv(c *websocket.Conn) (map[string]any, bool) { + var ev map[string]any + if err := c.ReadJSON(&ev); err != nil { + return nil, false + } + return ev, true +} + +// wsHandshake plays the upstream's session setup and returns the client's +// session.update event. +func wsHandshake(c *websocket.Conn) map[string]any { + defer GinkgoRecover() + wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}}) + upd, ok := wsRecv(c) + Expect(ok).To(BeTrue()) + Expect(upd["type"]).To(Equal("session.update")) + wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}}) + return upd +} + +// wsDrain reads client events until the client closes, returning the +// decoded PCM16 samples of every input_audio_buffer.append. +func wsDrain(c *websocket.Conn) [][]int16 { + var frames [][]int16 + for { + ev, ok := wsRecv(c) + if !ok { + return frames + } + if ev["type"] != "input_audio_buffer.append" { + continue + } + raw, err := base64.StdEncoding.DecodeString(ev["audio"].(string)) + if err != nil { + return frames + } + samples := make([]int16, len(raw)/2) + for i := range samples { + samples[i] = int16(binary.LittleEndian.Uint16(raw[i*2:])) + } + frames = append(frames, samples) + } +} + +type liveCall struct { + in chan *pb.TranscriptLiveRequest + out chan *pb.TranscriptLiveResponse + errc chan error +} + +// startLive runs AudioTranscriptionLive the way pkg/grpc/server.go does: +// buffered channels, the caller owning in and the backend owning out. +func startLive(p *LocalAIProxy) *liveCall { + lc := &liveCall{ + in: make(chan *pb.TranscriptLiveRequest, 4), + out: make(chan *pb.TranscriptLiveResponse, 4), + errc: make(chan error, 1), + } + go func() { lc.errc <- p.AudioTranscriptionLive(lc.in, lc.out) }() + return lc +} + +func (lc *liveCall) config(lang string, rate int32) { + lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Config{ + Config: &pb.TranscriptLiveConfig{Language: lang, SampleRate: rate}, + }} +} + +func (lc *liveCall) audio(pcm ...float32) { + lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Audio{ + Audio: &pb.TranscriptLiveAudio{Pcm: pcm}, + }} +} + +func (lc *liveCall) next() *pb.TranscriptLiveResponse { + var r *pb.TranscriptLiveResponse + EventuallyWithOffset(1, lc.out, 2*time.Second).Should(Receive(&r)) + return r +} + +// finish asserts the call returned within 2 s and out is closed, and +// returns the call's error. +func (lc *liveCall) finish() error { + var err error + EventuallyWithOffset(1, lc.errc, 2*time.Second).Should(Receive(&err)) + for range lc.out { + } + return err +} + +// liveGoroutines counts goroutines still running code from live.go, so a +// spec can prove the bridge left nothing blocked behind. +func liveGoroutines() int { + buf := make([]byte, 1<<20) + buf = buf[:runtime.Stack(buf, true)] + n := 0 + for _, g := range strings.Split(string(buf), "\n\n") { + if strings.Contains(g, "localai-proxy/live.go:") { + n++ + } + } + return n +} + +func ev(typ string) map[string]any { return map[string]any{"type": typ} } + +func item(typ, id string) map[string]any { return map[string]any{"type": typ, "item_id": id} } + +func completed(id, transcript string) map[string]any { + return map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": id, "transcript": transcript} +} + +// stallAfterUpdate plays session.created, reads the session.update and then +// never answers, as a hung upstream would, until the spec ends. +func stallAfterUpdate(c *websocket.Conn, gone chan<- struct{}) { + wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}}) + wsRecv(c) + for { + if _, ok := wsRecv(c); !ok { + close(gone) + return + } + } +} + +func withPipeline(o *pb.ModelOptions) { + o.Options = append(o.Options, "realtime_pipeline:remote-pipe") +} + +var _ = Describe("AudioTranscriptionLive", func() { + It("reports live transcription unsupported without realtime_pipeline", func() { + up := newFakeUpstream() + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, nil)) + + err := lc.finish() + Expect(grpcerrors.IsLiveTranscriptionUnsupported(err)).To(BeTrue()) + Expect(err.Error()).To(ContainSubstring("realtime_pipeline")) + Expect(up.recorded()).To(BeEmpty()) + }) + + It("rejects a first message that is not a config", func() { + up := wsUpstream(func(*websocket.Conn, *http.Request) {}) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.audio(0.1) + + Expect(status.Code(lc.finish())).To(Equal(codes.InvalidArgument)) + }) + + It("opens the pipeline session and acks ready only after session.updated", func() { + gotURL := make(chan string, 1) + gotAuth := make(chan string, 1) + gotUpdate := make(chan map[string]any, 1) + release := make(chan struct{}) + up := wsUpstream(func(c *websocket.Conn, r *http.Request) { + gotURL <- r.URL.RequestURI() + gotAuth <- r.Header.Get("Authorization") + wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}}) + upd, _ := wsRecv(c) + gotUpdate <- upd + <-release + wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}}) + wsDrain(c) + }) + DeferCleanup(up.Close) + keyFile := writeInput("key", "sekret\n") + p := loadProxy(up, func(o *pb.ModelOptions) { + withPipeline(o) + o.Proxy.ApiKeyFile = keyFile + }) + + lc := startLive(p) + lc.config("it", 24000) + + Eventually(gotURL, 2*time.Second).Should(Receive(Equal("/v1/realtime?model=remote-pipe"))) + Expect(gotAuth).To(Receive(Equal("Bearer sekret"))) + var upd map[string]any + Eventually(gotUpdate, 2*time.Second).Should(Receive(&upd)) + raw, err := json.Marshal(upd) + Expect(err).NotTo(HaveOccurred()) + Expect(raw).To(MatchJSON(`{"type":"session.update","session":{"type":"transcription","audio":{"input":{ + "format":{"type":"audio/pcm","rate":24000}, + "transcription":{"model":"remote-pipe","language":"it"}, + "turn_detection":{"type":"server_vad"}}}}}`)) + + Consistently(lc.out, 200*time.Millisecond).ShouldNot(Receive()) + close(release) + Expect(lc.next().GetReady()).To(BeTrue()) + + close(lc.in) + Expect(lc.finish()).To(Succeed()) + }) + + It("defaults the session rate to 16000", func() { + gotUpdate := make(chan map[string]any, 1) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + gotUpdate <- wsHandshake(c) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("", 0) + Expect(lc.next().GetReady()).To(BeTrue()) + + var upd map[string]any + Expect(gotUpdate).To(Receive(&upd)) + format := upd["session"].(map[string]any)["audio"].(map[string]any)["input"].(map[string]any)["format"] + Expect(format).To(HaveKeyWithValue("rate", BeNumerically("==", 16000))) + + close(lc.in) + Expect(lc.finish()).To(Succeed()) + }) + + It("forwards audio as base64 PCM16 appends", func() { + frames := make(chan [][]int16, 1) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + frames <- wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + + lc.audio(0, 0.5, -1, 1, 2) + lc.audio(-0.25) + lc.audio(float32(math.NaN())) + close(lc.in) + Expect(lc.finish()).To(Succeed()) + + Eventually(frames, 2*time.Second).Should(Receive(Equal([][]int16{ + {0, 16383, -32767, 32767, 32767}, + {-8191}, + {0}, + }))) + }) + + It("maps deltas and completions to Delta and Eou, and finishes with the full text", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + wsSend(c, ev("input_audio_buffer.speech_started")) + wsSend(c, ev("input_audio_buffer.speech_stopped")) + wsSend(c, item("input_audio_buffer.committed", "a")) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "hel"}) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "lo"}) + wsSend(c, completed("a", "hello world")) + wsSend(c, ev("input_audio_buffer.speech_started")) + wsSend(c, ev("input_audio_buffer.speech_stopped")) + wsSend(c, item("input_audio_buffer.committed", "b")) + wsSend(c, completed("b", "again")) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + + Expect(lc.next()).To(SatisfyAll( + WithTransform((*pb.TranscriptLiveResponse).GetDelta, Equal("hel")), + WithTransform((*pb.TranscriptLiveResponse).GetEou, BeFalse()))) + Expect(lc.next().GetDelta()).To(Equal("lo")) + r := lc.next() + Expect(r.GetDelta()).To(Equal(" world")) + Expect(r.GetEou()).To(BeTrue()) + r = lc.next() + Expect(r.GetDelta()).To(Equal("again")) + Expect(r.GetEou()).To(BeTrue()) + + close(lc.in) + Expect(lc.next().GetFinalResult().GetText()).To(Equal("hello world again")) + Expect(lc.finish()).To(Succeed()) + }) + + It("waits for a committed utterance before the final result", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + wsSend(c, ev("input_audio_buffer.speech_started")) + wsSend(c, ev("input_audio_buffer.speech_stopped")) + wsSend(c, item("input_audio_buffer.committed", "a")) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "late"}) + // The client closes its side once it sees the delta; the + // transcription completes later. + time.Sleep(300 * time.Millisecond) + wsSend(c, completed("a", "late words")) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + Expect(lc.next().GetDelta()).To(Equal("late")) + close(lc.in) + + r := lc.next() + Expect(r.GetDelta()).To(Equal(" words")) + Expect(r.GetEou()).To(BeTrue()) + Expect(lc.next().GetFinalResult().GetText()).To(Equal("late words")) + Expect(lc.finish()).To(Succeed()) + }) + + It("does not hold the close for a turn the upstream discarded", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + wsSend(c, ev("input_audio_buffer.speech_started")) + wsSend(c, ev("input_audio_buffer.speech_stopped")) + wsSend(c, item("input_audio_buffer.committed", "a")) + wsSend(c, completed("a", "kept")) + // A stop that is never committed, then nothing more. + wsSend(c, ev("input_audio_buffer.speech_started")) + wsSend(c, ev("input_audio_buffer.speech_stopped")) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + Expect(lc.next().GetDelta()).To(Equal("kept")) + close(lc.in) + + start := time.Now() + Expect(lc.next().GetFinalResult().GetText()).To(Equal("kept")) + Expect(lc.finish()).To(Succeed()) + Expect(time.Since(start)).To(BeNumerically("<", finalWait/2)) + }) + + It("ends with Unavailable on an upstream error event during setup", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}}) + wsRecv(c) + wsSend(c, map[string]any{"type": "error", "error": map[string]any{ + "type": "invalid_request_error", "code": "session_update_error", + "message": "model is not a valid pipeline model: remote-pipe", + }}) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + + err := lc.finish() + Expect(status.Code(err)).To(Equal(codes.Unavailable)) + Expect(err.Error()).To(ContainSubstring("not a valid pipeline model")) + }) + + It("ends with Unavailable when a transcription fails", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.failed", "item_id": "a", + "error": map[string]any{"message": "backend crashed"}}) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + + err := lc.finish() + Expect(status.Code(err)).To(Equal(codes.Unavailable)) + Expect(err.Error()).To(ContainSubstring("backend crashed")) + }) + + It("maps a refused upgrade to the upstream status", func() { + up := newFakeUpstreamWithHandler(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "loading", http.StatusServiceUnavailable) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + + Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable)) + }) + + It("upstream disconnect ends the stream", func() { + before := liveGoroutines() + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + // Drop the socket mid-session, as a crashed upstream would. + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + // The caller keeps its send side open; the bridge must not wait on it. + + Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable)) + Eventually(liveGoroutines, time.Second).Should(Equal(before)) + close(lc.in) + }) + + It("gives up with Canceled when the caller closes before the ready ack", func() { + before := liveGoroutines() + gone := make(chan struct{}) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + lc.audio(0.1) + close(lc.in) + + Expect(status.Code(lc.finish())).To(Equal(codes.Canceled)) + Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open") + Eventually(liveGoroutines, time.Second).Should(Equal(before)) + }) + + It("bounds a hung setup by request_timeout_seconds", func() { + before := liveGoroutines() + gone := make(chan struct{}) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) }) + DeferCleanup(up.Close) + p := loadProxy(up, func(o *pb.ModelOptions) { + withPipeline(o) + o.Proxy.RequestTimeoutSeconds = 1 + }) + lc := startLive(p) + lc.config("en", 16000) + + var err error + Eventually(lc.errc, 3*time.Second).Should(Receive(&err)) + Expect(status.Code(err)).To(Equal(codes.Unavailable)) + Expect(lc.out).To(BeClosed()) + Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open") + Eventually(liveGoroutines, time.Second).Should(Equal(before)) + close(lc.in) + }) + + It("bounds a hung setup by default when request_timeout_seconds is unset", func() { + Expect(defaultLiveSetupTimeout).To(BeNumerically(">=", 2*time.Minute)) + saved := liveSetupTimeout + liveSetupTimeout = 300 * time.Millisecond + DeferCleanup(func() { liveSetupTimeout = saved }) + + before := liveGoroutines() + gone := make(chan struct{}) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + + Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable)) + Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open") + Eventually(liveGoroutines, time.Second).Should(Equal(before)) + close(lc.in) + }) + + It("forwards audio sent before the ready ack once ready", func() { + release := make(chan struct{}) + frames := make(chan [][]int16, 1) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}}) + wsRecv(c) + <-release + wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}}) + frames <- wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + lc.audio(0.5) + close(release) + Expect(lc.next().GetReady()).To(BeTrue()) + close(lc.in) + Expect(lc.finish()).To(Succeed()) + Eventually(frames, 2*time.Second).Should(Receive(Equal([][]int16{{16383}}))) + }) + + It("refuses more than the backlog cap of audio before the ready ack", func() { + gone := make(chan struct{}) + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + second := make([]float32, 16000) + for range maxBacklogSeconds + 1 { + select { + case lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Audio{Audio: &pb.TranscriptLiveAudio{Pcm: second}}}: + case <-time.After(2 * time.Second): + Fail("bridge stopped reading audio before the cap") + } + } + + err := lc.finish() + Expect(status.Code(err)).To(Equal(codes.InvalidArgument)) + Expect(err.Error()).To(ContainSubstring("before the live transcription session was ready")) + close(lc.in) + }) +}) diff --git a/backend/go/localai-proxy/localai_proxy_suite_test.go b/backend/go/localai-proxy/localai_proxy_suite_test.go new file mode 100644 index 000000000..62edc82f3 --- /dev/null +++ b/backend/go/localai-proxy/localai_proxy_suite_test.go @@ -0,0 +1,17 @@ +package main + +import ( + "testing" + + "github.com/mudler/xlog" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestLocalAIProxy(t *testing.T) { + RegisterFailHandler(Fail) + // The specs drive upstream failures on purpose; their warnings are + // expected and would only bury real failures in the output. + xlog.SetLogger(xlog.NewLogger(xlog.LogLevelError, xlog.TextFormat)) + RunSpecs(t, "localai-proxy specs") +} diff --git a/backend/go/localai-proxy/main.go b/backend/go/localai-proxy/main.go new file mode 100644 index 000000000..93e5534fa --- /dev/null +++ b/backend/go/localai-proxy/main.go @@ -0,0 +1,36 @@ +package main + +// localai-proxy is a LocalAI backend that serves backend gRPC methods by +// calling the REST API of another LocalAI instance. It lets a model config +// (and so a failover chain target) live on a remote LocalAI while callers +// keep using the local backend interface for every modality, not only chat. + +import ( + "flag" + "os" + + grpc "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/xlog" + "golang.org/x/term" +) + +var addr = flag.String("addr", "localhost:50051", "the address to listen on") + +func main() { + // xlog's default handler emits ANSI color codes, which are unreadable once + // LocalAI captures the backend's stdout into a log file. Force plain text + // when LOCALAI_LOG_FORMAT is unset and stdout is not a terminal. + format := os.Getenv("LOCALAI_LOG_FORMAT") + if format == "" && !term.IsTerminal(int(os.Stdout.Fd())) { + format = xlog.TextFormat + } + xlog.SetLogger(xlog.NewLogger(xlog.LogLevel(os.Getenv("LOCALAI_LOG_LEVEL")), format)) + flag.Parse() + if err := grpc.StartServer(*addr, NewLocalAIProxy()); err != nil { + panic(err) + } +} + +// The gRPC server only passes the call's context to backends that implement +// this, so keep the proxy on it. +var _ grpc.AIModelRichContext = (*LocalAIProxy)(nil) diff --git a/backend/go/localai-proxy/media.go b/backend/go/localai-proxy/media.go new file mode 100644 index 000000000..8b598c01d --- /dev/null +++ b/backend/go/localai-proxy/media.go @@ -0,0 +1,862 @@ +package main + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/url" + "os" + "path/filepath" + "strconv" + "strings" + + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// fileToBase64 reads path and base64-encodes its contents, or returns "" for +// an empty path. Core stages image/video/3D conditioning inputs to local +// files before calling the backend, but the REST endpoints on the other side +// take the same media inline as base64 (or a URL/data-URI, neither of which +// a local path is), so every media method needs this same conversion. +func fileToBase64(path string) (string, error) { + if path == "" { + return "", nil + } + // #nosec G304 -- path is a staging file core wrote for this call, never a caller-supplied path + data, err := os.ReadFile(filepath.Clean(path)) + if err != nil { + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: read %s: %v", path, err) + } + return base64.StdEncoding.EncodeToString(data), nil +} + +// toDataURI wraps a base64 payload as a data URI. Detect/Depth/FaceVerify/ +// FaceAnalyze's REST endpoints decode their image field with +// utils.GetContentURIAsBase64, which only accepts an http(s) URL or a +// `data:;base64,` string — never bare base64. That is exactly +// what core hands the backend for these methods (see +// core/http/endpoints/localai/images.go's decodeImageInput, which already +// stripped any data: prefix off before we ever see it), so forwarding it +// unwrapped 400s on every real call. The MIME type isn't carried alongside +// the payload, so it's sniffed from the decoded bytes. +func toDataURI(b64 string) (string, error) { + if b64 == "" { + return "", nil + } + data, err := base64.StdEncoding.DecodeString(b64) + if err != nil { + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: decode base64 image: %v", err) + } + mime := http.DetectContentType(data) + if mime == "application/octet-stream" { + mime = "image/png" + } + return "data:" + mime + ";base64," + b64, nil +} + +// genItem is one schema.Item as the image/video/3D generation endpoints +// return it: either inline base64 data, or a URL to download the asset from. +type genItem struct { + URL string `json:"url,omitempty"` + B64JSON string `json:"b64_json,omitempty"` +} + +type genResponse struct { + Data []genItem `json:"data"` +} + +// writeGenItem writes the first item of a generation reply to dst: it +// base64-decodes inline data directly, or downloads the URL relative to the +// configured upstream (same client, same auth) when the upstream sent one +// instead. +func (p *LocalAIProxy) writeGenItem(ctx context.Context, path string, items []genItem, dst string) error { + if len(items) == 0 { + return status.Errorf(codes.Internal, "localai-proxy: upstream %s returned no data", path) + } + item := items[0] + if item.B64JSON != "" { + data, err := base64.StdEncoding.DecodeString(item.B64JSON) + if err != nil { + return status.Errorf(codes.Internal, "localai-proxy: decode %s b64_json: %v", path, err) + } + // 0o600: core runs as the same user and serves the file itself, so no + // one else needs to read generated media. + if err := os.WriteFile(dst, data, 0o600); err != nil { + _ = os.Remove(dst) + return status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, err) + } + return nil + } + if item.URL == "" { + return status.Errorf(codes.Internal, "localai-proxy: upstream %s returned neither b64_json nor url", path) + } + rel, err := generatedContentPath(path, item.URL) + if err != nil { + return err + } + return p.getToFile(ctx, rel, dst) +} + +// generatedContentPrefixes are the static paths a LocalAI instance serves +// generated media under (core/http/app.go's e.Static calls). A generation +// reply's URL is only ever safe to re-fetch through this proxy's own +// authenticated client when it resolves to one of these. +var generatedContentPrefixes = []string{"/generated-images/", "/generated-videos/", "/generated-audio/", "/generated-3d/"} + +// generatedContentPath extracts the path to re-download raw from, ignoring +// whatever host is in it. A literal-prefix strip of the configured upstream +// base (the previous approach) breaks the moment the upstream advertises a +// different host than the one this proxy is configured with — LOCALAI_BASE_URL, +// a reverse proxy, or X-Forwarded-Host can all change it — so only the path is +// trusted, and only when it is one this proxy's upstream is actually known to +// serve generated media under; anything else could point anywhere, and taking +// it on faith would let a compromised or misconfigured upstream make this +// proxy fetch (with its bearer key) whatever URL it likes. +func generatedContentPath(callPath, raw string) (string, error) { + u, err := url.Parse(raw) + if err != nil { + return "", status.Errorf(codes.Internal, "localai-proxy: upstream %s returned an invalid url %q: %v", callPath, raw, err) + } + for _, prefix := range generatedContentPrefixes { + if strings.HasPrefix(u.Path, prefix) { + if u.RawQuery != "" { + return u.Path + "?" + u.RawQuery, nil + } + return u.Path, nil + } + } + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: upstream %s returned an unexpected url %q", callPath, raw) +} + +// --- Images --------------------------------------------------------- + +type imageGenerationRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + NegativePrompt string `json:"negative_prompt,omitempty"` + Size string `json:"size,omitempty"` + Step int32 `json:"step,omitempty"` + Seed int32 `json:"seed,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + File string `json:"file,omitempty"` + RefImages []string `json:"ref_images,omitempty"` +} + +// GenerateImage sends Src and RefImages (local paths) as base64, asks for a +// b64_json reply, and writes the result to Dst. +func (p *LocalAIProxy) GenerateImage(req *pb.GenerateImageRequest) error { + ctx := context.Background() + file, err := fileToBase64(req.GetSrc()) + if err != nil { + return err + } + var refImages []string + for _, ref := range req.GetRefImages() { + encoded, err := fileToBase64(ref) + if err != nil { + return err + } + refImages = append(refImages, encoded) + } + body := imageGenerationRequest{ + Model: p.model(""), + Prompt: req.GetPositivePrompt(), + NegativePrompt: req.GetNegativePrompt(), + Size: fmt.Sprintf("%dx%d", req.GetWidth(), req.GetHeight()), + Step: req.GetStep(), + Seed: req.GetSeed(), + ResponseFormat: "b64_json", + File: file, + RefImages: refImages, + } + var resp genResponse + if err := p.postJSON(ctx, "/v1/images/generations", body, &resp); err != nil { + return err + } + return p.writeGenItem(ctx, "/v1/images/generations", resp.Data, req.GetDst()) +} + +// UpscaleImage uploads Src as a multipart file, matching the REST endpoint's +// upload-only contract, and writes the (always URL) reply to Dst. +func (p *LocalAIProxy) UpscaleImage(req *pb.UpscaleImageRequest) error { + ctx := context.Background() + fields := url.Values{} + fields.Set("model", p.model("")) + if req.GetScale() > 0 { + fields.Set("scale", strconv.Itoa(int(req.GetScale()))) + } + form := multipartForm{fields: fields, files: []formFile{{field: "image", path: req.GetSrc()}}} + var resp genResponse + if err := p.postForm(ctx, "/v1/images/upscale", form, &resp); err != nil { + return err + } + return p.writeGenItem(ctx, "/v1/images/upscale", resp.Data, req.GetDst()) +} + +// --- Video ------------------------------------------------------------ + +type videoGenerationRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + NegativePrompt string `json:"negative_prompt,omitempty"` + StartImage string `json:"start_image,omitempty"` + EndImage string `json:"end_image,omitempty"` + Audio string `json:"audio,omitempty"` + Width int32 `json:"width,omitempty"` + Height int32 `json:"height,omitempty"` + NumFrames int32 `json:"num_frames,omitempty"` + FPS int32 `json:"fps,omitempty"` + Seed int32 `json:"seed,omitempty"` + CFGScale float32 `json:"cfg_scale,omitempty"` + Step int32 `json:"step,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + Params map[string]string `json:"params,omitempty"` +} + +// GenerateVideo sends the staged local media (StartImage, EndImage, Audio) as +// base64, asks for a b64_json reply, and writes the result to Dst. +func (p *LocalAIProxy) GenerateVideo(req *pb.GenerateVideoRequest) error { + ctx := context.Background() + startImage, err := fileToBase64(req.GetStartImage()) + if err != nil { + return err + } + endImage, err := fileToBase64(req.GetEndImage()) + if err != nil { + return err + } + audio, err := fileToBase64(req.GetAudio()) + if err != nil { + return err + } + body := videoGenerationRequest{ + Model: p.model(""), + Prompt: req.GetPrompt(), + NegativePrompt: req.GetNegativePrompt(), + StartImage: startImage, + EndImage: endImage, + Audio: audio, + Width: req.GetWidth(), + Height: req.GetHeight(), + NumFrames: req.GetNumFrames(), + FPS: req.GetFps(), + Seed: req.GetSeed(), + CFGScale: req.GetCfgScale(), + Step: req.GetStep(), + ResponseFormat: "b64_json", + Params: req.GetParams(), + } + var resp genResponse + if err := p.postJSON(ctx, "/video", body, &resp); err != nil { + return err + } + return p.writeGenItem(ctx, "/video", resp.Data, req.GetDst()) +} + +// --- 3D ----------------------------------------------------------------- + +type model3DRequest struct { + Model string `json:"model"` + Image string `json:"image"` + Seed int32 `json:"seed,omitempty"` + Step int32 `json:"step,omitempty"` + CFGScale float32 `json:"cfg_scale,omitempty"` + TextureSteps int32 `json:"texture_steps,omitempty"` + Quality string `json:"quality,omitempty"` + Background string `json:"background,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + Params map[string]string `json:"params,omitempty"` +} + +// Generate3D sends Src (the staged conditioning image, a local path) as +// base64, asks for a b64_json reply, and writes the result to Dst. +func (p *LocalAIProxy) Generate3D(req *pb.Generate3DRequest) error { + ctx := context.Background() + image, err := fileToBase64(req.GetSrc()) + if err != nil { + return err + } + body := model3DRequest{ + Model: p.model(""), + Image: image, + Seed: req.GetSeed(), + Step: req.GetStep(), + CFGScale: req.GetCfgScale(), + TextureSteps: req.GetTextureSteps(), + Quality: req.GetQuality(), + Background: req.GetBackground(), + ResponseFormat: "b64_json", + Params: req.GetParams(), + } + var resp genResponse + if err := p.postJSON(ctx, "/3d/generations", body, &resp); err != nil { + return err + } + return p.writeGenItem(ctx, "/3d/generations", resp.Data, req.GetDst()) +} + +type animationInputBody struct { + Type string `json:"type"` + Data string `json:"data"` +} + +type animate3DRequestBody struct { + Model string `json:"model"` + Inputs map[string]animationInputBody `json:"inputs"` + Params map[string]string `json:"params,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` +} + +type animate3DResponseBody struct { + genResponse + Metadata json.RawMessage `json:"metadata,omitempty"` +} + +// animate3D does the shared work behind Animate3D and Animate3DWithMetadata: +// non-text inputs are staged local paths, so they go over the wire as +// base64, like every other media field; text inputs travel verbatim. +func (p *LocalAIProxy) animate3D(ctx context.Context, req *pb.Animate3DRequest) (json.RawMessage, error) { + inputs := make(map[string]animationInputBody, len(req.GetInputs())) + for name, in := range req.GetInputs() { + data := in.GetData() + if in.GetType() != "text" { + encoded, err := fileToBase64(data) + if err != nil { + return nil, err + } + data = encoded + } + inputs[name] = animationInputBody{Type: in.GetType(), Data: data} + } + body := animate3DRequestBody{ + Model: p.model(""), + Inputs: inputs, + Params: req.GetParams(), + ResponseFormat: "b64_json", + } + var resp animate3DResponseBody + if err := p.postJSON(ctx, "/3d/animate", body, &resp); err != nil { + return nil, err + } + if err := p.writeGenItem(ctx, "/3d/animate", resp.Data, req.GetDst()); err != nil { + return nil, err + } + return resp.Metadata, nil +} + +func (p *LocalAIProxy) Animate3D(req *pb.Animate3DRequest) error { + _, err := p.animate3D(context.Background(), req) + return err +} + +// Animate3DWithMetadata implements grpc.AnimationMetadataModel: the REST +// reply's metadata field carries whatever the animation backend reported, and +// the gRPC server prefers this method over Animate3D when it is implemented. +func (p *LocalAIProxy) Animate3DWithMetadata(req *pb.Animate3DRequest) ([]byte, error) { + return p.animate3D(context.Background(), req) +} + +// --- Vision (detection, depth) ------------------------------------------ + +type detectRequestBody struct { + Model string `json:"model"` + Image string `json:"image"` + Prompt string `json:"prompt,omitempty"` + Points []float32 `json:"points,omitempty"` + Boxes []float32 `json:"boxes,omitempty"` + Threshold float32 `json:"threshold,omitempty"` +} + +type detectionBody struct { + X float32 `json:"x"` + Y float32 `json:"y"` + Width float32 `json:"width"` + Height float32 `json:"height"` + ClassName string `json:"class_name"` + Confidence float32 `json:"confidence,omitempty"` + Mask string `json:"mask,omitempty"` +} + +type detectResponseBody struct { + Detections []detectionBody `json:"detections"` +} + +// Detect posts Src, which core already carries as a bare base64 payload (the +// same convention DetectionEndpoint uses to call this method locally) — but +// the upstream's own endpoint only accepts a URL or a data URI, so it is +// re-wrapped as one (see toDataURI) — and maps each detection, decoding its +// PNG mask. +func (p *LocalAIProxy) Detect(req *pb.DetectOptions) (pb.DetectResponse, error) { + image, err := toDataURI(req.GetSrc()) + if err != nil { + return pb.DetectResponse{}, err + } + body := detectRequestBody{ + Model: p.model(""), + Image: image, + Prompt: req.GetPrompt(), + Points: req.GetPoints(), + Boxes: req.GetBoxes(), + Threshold: req.GetThreshold(), + } + var resp detectResponseBody + if err := p.postJSON(context.Background(), "/v1/detection", body, &resp); err != nil { + return pb.DetectResponse{}, err + } + var detections []*pb.Detection + for _, d := range resp.Detections { + det := &pb.Detection{ + X: d.X, Y: d.Y, Width: d.Width, Height: d.Height, + Confidence: d.Confidence, ClassName: d.ClassName, + } + if d.Mask != "" { + mask, err := base64.StdEncoding.DecodeString(d.Mask) + if err != nil { + return pb.DetectResponse{}, status.Errorf(codes.Internal, "localai-proxy: decode /v1/detection mask: %v", err) + } + det.Mask = mask + } + detections = append(detections, det) + } + // A composite literal here, not a variable of type pb.DetectResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.DetectResponse{Detections: detections}, nil +} + +// depthRequestBody has no Dst/Exports fields: Depth refuses those requests +// before building the body (see the Depth doc comment), so they never reach +// the upstream. +type depthRequestBody struct { + Model string `json:"model"` + Image string `json:"image"` + IncludeDepth bool `json:"include_depth,omitempty"` + IncludeConfidence bool `json:"include_confidence,omitempty"` + IncludePose bool `json:"include_pose,omitempty"` + IncludeSky bool `json:"include_sky,omitempty"` + IncludePoints bool `json:"include_points,omitempty"` + PointsConfThresh float32 `json:"points_conf_thresh,omitempty"` +} + +type depthResponseBody struct { + Width int32 `json:"width"` + Height int32 `json:"height"` + Depth []float32 `json:"depth,omitempty"` + Confidence []float32 `json:"confidence,omitempty"` + Sky []float32 `json:"sky,omitempty"` + Extrinsics []float32 `json:"extrinsics,omitempty"` + Intrinsics []float32 `json:"intrinsics,omitempty"` + NumPoints int32 `json:"num_points,omitempty"` + Points []float32 `json:"points,omitempty"` + PointColors string `json:"point_colors,omitempty"` + ExportPaths []string `json:"export_paths,omitempty"` + IsMetric bool `json:"is_metric"` +} + +// Depth posts Src (a bare base64 payload, per the same convention as Detect, +// re-wrapped as a data URI via toDataURI) and maps the full response, +// decoding the point-cloud color bytes. A request for exports or a dst +// directory is refused: those would be written to the upstream's own local +// disk, and ExportPaths would name files this proxy (and whatever asked it +// for them) can never reach. +func (p *LocalAIProxy) Depth(req *pb.DepthRequest) (pb.DepthResponse, error) { + if req.GetDst() != "" || len(req.GetExports()) > 0 { + return pb.DepthResponse{}, unimplemented("Depth exports (written to the upstream's own local disk, unreachable from here)") + } + image, err := toDataURI(req.GetSrc()) + if err != nil { + return pb.DepthResponse{}, err + } + body := depthRequestBody{ + Model: p.model(""), + Image: image, + IncludeDepth: req.GetIncludeDepth(), + IncludeConfidence: req.GetIncludeConfidence(), + IncludePose: req.GetIncludePose(), + IncludeSky: req.GetIncludeSky(), + IncludePoints: req.GetIncludePoints(), + PointsConfThresh: req.GetPointsConfThresh(), + } + var resp depthResponseBody + if err := p.postJSON(context.Background(), "/v1/depth", body, &resp); err != nil { + return pb.DepthResponse{}, err + } + var colors []byte + if resp.PointColors != "" { + var err error + colors, err = base64.StdEncoding.DecodeString(resp.PointColors) + if err != nil { + return pb.DepthResponse{}, status.Errorf(codes.Internal, "localai-proxy: decode /v1/depth point_colors: %v", err) + } + } + // A composite literal here, not a variable of type pb.DepthResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.DepthResponse{ + Width: resp.Width, Height: resp.Height, Depth: resp.Depth, Confidence: resp.Confidence, + Sky: resp.Sky, Extrinsics: resp.Extrinsics, Intrinsics: resp.Intrinsics, + NumPoints: resp.NumPoints, Points: resp.Points, ExportPaths: resp.ExportPaths, IsMetric: resp.IsMetric, + PointColors: colors, + }, nil +} + +// --- Face recognition ----------------------------------------------------- + +type facialAreaBody struct { + X float32 `json:"x"` + Y float32 `json:"y"` + W float32 `json:"w"` + H float32 `json:"h"` +} + +func (a facialAreaBody) toProto() *pb.FacialArea { + return &pb.FacialArea{X: a.X, Y: a.Y, W: a.W, H: a.H} +} + +type faceVerifyRequestBody struct { + Model string `json:"model"` + Img1 string `json:"img1"` + Img2 string `json:"img2"` + Threshold float32 `json:"threshold,omitempty"` + AntiSpoofing bool `json:"anti_spoofing,omitempty"` +} + +type faceVerifyResponseBody struct { + Verified bool `json:"verified"` + Distance float32 `json:"distance"` + Threshold float32 `json:"threshold"` + Confidence float32 `json:"confidence"` + Model string `json:"model"` + Img1Area facialAreaBody `json:"img1_area"` + Img2Area facialAreaBody `json:"img2_area"` + ProcessingTimeMs float32 `json:"processing_time_ms,omitempty"` + Img1IsReal *bool `json:"img1_is_real,omitempty"` + Img1AntispoofScore *float32 `json:"img1_antispoof_score,omitempty"` + Img2IsReal *bool `json:"img2_is_real,omitempty"` + Img2AntispoofScore *float32 `json:"img2_antispoof_score,omitempty"` +} + +// FaceVerify posts Img1/Img2, which core already carries as bare base64 (the +// same convention FaceVerifyEndpoint uses to call this method locally) — +// re-wrapped as data URIs via toDataURI, since the upstream's own endpoint +// only accepts a URL or a data URI. +func (p *LocalAIProxy) FaceVerify(req *pb.FaceVerifyRequest) (pb.FaceVerifyResponse, error) { + img1, err := toDataURI(req.GetImg1()) + if err != nil { + return pb.FaceVerifyResponse{}, err + } + img2, err := toDataURI(req.GetImg2()) + if err != nil { + return pb.FaceVerifyResponse{}, err + } + body := faceVerifyRequestBody{ + Model: p.model(""), Img1: img1, Img2: img2, + Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(), + } + var resp faceVerifyResponseBody + if err := p.postJSON(context.Background(), "/v1/face/verify", body, &resp); err != nil { + return pb.FaceVerifyResponse{}, err + } + // A composite literal here, not a variable of type pb.FaceVerifyResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.FaceVerifyResponse{ + Verified: resp.Verified, Distance: resp.Distance, Threshold: resp.Threshold, + Confidence: resp.Confidence, Model: resp.Model, + Img1Area: resp.Img1Area.toProto(), Img2Area: resp.Img2Area.toProto(), + ProcessingTimeMs: resp.ProcessingTimeMs, + Img1IsReal: boolValue(resp.Img1IsReal), + Img1AntispoofScore: float32Value(resp.Img1AntispoofScore), + Img2IsReal: boolValue(resp.Img2IsReal), + Img2AntispoofScore: float32Value(resp.Img2AntispoofScore), + }, nil +} + +// boolValue and float32Value read the liveness fields the REST endpoints only +// populate when anti_spoofing was requested; proto keeps them as plain +// bool/float32 (no "not checked" state), so an absent pointer becomes zero. +func boolValue(b *bool) bool { + return b != nil && *b +} + +func float32Value(v *float32) float32 { + if v == nil { + return 0 + } + return *v +} + +type faceAnalyzeRequestBody struct { + Model string `json:"model"` + Img string `json:"img"` + Actions []string `json:"actions,omitempty"` + AntiSpoofing bool `json:"anti_spoofing,omitempty"` +} + +type faceAnalysisBody struct { + Region facialAreaBody `json:"region"` + FaceConfidence float32 `json:"face_confidence"` + Age float32 `json:"age,omitempty"` + DominantGender string `json:"dominant_gender,omitempty"` + Gender map[string]float32 `json:"gender,omitempty"` + DominantEmotion string `json:"dominant_emotion,omitempty"` + Emotion map[string]float32 `json:"emotion,omitempty"` + DominantRace string `json:"dominant_race,omitempty"` + Race map[string]float32 `json:"race,omitempty"` + IsReal *bool `json:"is_real,omitempty"` + AntispoofScore *float32 `json:"antispoof_score,omitempty"` +} + +type faceAnalyzeResponseBody struct { + Faces []faceAnalysisBody `json:"faces"` +} + +// FaceAnalyze posts Img, which core already carries as bare base64 — +// re-wrapped as a data URI via toDataURI, since the upstream's own endpoint +// only accepts a URL or a data URI. +func (p *LocalAIProxy) FaceAnalyze(req *pb.FaceAnalyzeRequest) (pb.FaceAnalyzeResponse, error) { + img, err := toDataURI(req.GetImg()) + if err != nil { + return pb.FaceAnalyzeResponse{}, err + } + body := faceAnalyzeRequestBody{ + Model: p.model(""), Img: img, Actions: req.GetActions(), AntiSpoofing: req.GetAntiSpoofing(), + } + var resp faceAnalyzeResponseBody + if err := p.postJSON(context.Background(), "/v1/face/analyze", body, &resp); err != nil { + return pb.FaceAnalyzeResponse{}, err + } + var faces []*pb.FaceAnalysis + for _, f := range resp.Faces { + faces = append(faces, &pb.FaceAnalysis{ + Region: f.Region.toProto(), FaceConfidence: f.FaceConfidence, Age: f.Age, + DominantGender: f.DominantGender, Gender: f.Gender, + DominantEmotion: f.DominantEmotion, Emotion: f.Emotion, + DominantRace: f.DominantRace, Race: f.Race, + IsReal: boolValue(f.IsReal), AntispoofScore: float32Value(f.AntispoofScore), + }) + } + // A composite literal here, not a variable of type pb.FaceAnalyzeResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.FaceAnalyzeResponse{Faces: faces}, nil +} + +// --- Voice (speaker) recognition ----------------------------------------- + +type voiceVerifyRequestBody struct { + Model string `json:"model"` + Audio1 string `json:"audio1"` + Audio2 string `json:"audio2"` + Threshold float32 `json:"threshold,omitempty"` + AntiSpoofing bool `json:"anti_spoofing,omitempty"` +} + +type voiceVerifyResponseBody struct { + Verified bool `json:"verified"` + Distance float32 `json:"distance"` + Threshold float32 `json:"threshold"` + Confidence float32 `json:"confidence"` + Model string `json:"model"` + ProcessingTimeMs float32 `json:"processing_time_ms,omitempty"` +} + +// VoiceVerify sends Audio1/Audio2 (staged local paths) as base64, matching +// VoiceVerifyRequest's URL/base64/data-URI contract. +func (p *LocalAIProxy) VoiceVerify(req *pb.VoiceVerifyRequest) (pb.VoiceVerifyResponse, error) { + audio1, err := fileToBase64(req.GetAudio1()) + if err != nil { + return pb.VoiceVerifyResponse{}, err + } + audio2, err := fileToBase64(req.GetAudio2()) + if err != nil { + return pb.VoiceVerifyResponse{}, err + } + body := voiceVerifyRequestBody{ + Model: p.model(""), Audio1: audio1, Audio2: audio2, + Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(), + } + var resp voiceVerifyResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/verify", body, &resp); err != nil { + return pb.VoiceVerifyResponse{}, err + } + return pb.VoiceVerifyResponse{ + Verified: resp.Verified, Distance: resp.Distance, Threshold: resp.Threshold, + Confidence: resp.Confidence, Model: resp.Model, ProcessingTimeMs: resp.ProcessingTimeMs, + }, nil +} + +type voiceAnalyzeRequestBody struct { + Model string `json:"model"` + Audio string `json:"audio"` + Actions []string `json:"actions,omitempty"` +} + +type voiceAnalysisBody struct { + Start float32 `json:"start"` + End float32 `json:"end"` + Age float32 `json:"age,omitempty"` + DominantGender string `json:"dominant_gender,omitempty"` + Gender map[string]float32 `json:"gender,omitempty"` + DominantEmotion string `json:"dominant_emotion,omitempty"` + Emotion map[string]float32 `json:"emotion,omitempty"` +} + +type voiceAnalyzeResponseBody struct { + Segments []voiceAnalysisBody `json:"segments"` +} + +// VoiceAnalyze sends Audio (a staged local path) as base64. +func (p *LocalAIProxy) VoiceAnalyze(req *pb.VoiceAnalyzeRequest) (pb.VoiceAnalyzeResponse, error) { + audio, err := fileToBase64(req.GetAudio()) + if err != nil { + return pb.VoiceAnalyzeResponse{}, err + } + body := voiceAnalyzeRequestBody{Model: p.model(""), Audio: audio, Actions: req.GetActions()} + var resp voiceAnalyzeResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/analyze", body, &resp); err != nil { + return pb.VoiceAnalyzeResponse{}, err + } + var segments []*pb.VoiceAnalysis + for _, s := range resp.Segments { + segments = append(segments, &pb.VoiceAnalysis{ + Start: s.Start, End: s.End, Age: s.Age, + DominantGender: s.DominantGender, Gender: s.Gender, + DominantEmotion: s.DominantEmotion, Emotion: s.Emotion, + }) + } + // A composite literal here, not a variable of type pb.VoiceAnalyzeResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.VoiceAnalyzeResponse{Segments: segments}, nil +} + +type voiceEmbedRequestBody struct { + Model string `json:"model"` + Audio string `json:"audio"` +} + +type voiceEmbedResponseBody struct { + Embedding []float32 `json:"embedding"` + Model string `json:"model,omitempty"` +} + +// VoiceEmbed sends Audio (a staged local path) as base64. +func (p *LocalAIProxy) VoiceEmbed(req *pb.VoiceEmbedRequest) (pb.VoiceEmbedResponse, error) { + audio, err := fileToBase64(req.GetAudio()) + if err != nil { + return pb.VoiceEmbedResponse{}, err + } + body := voiceEmbedRequestBody{Model: p.model(""), Audio: audio} + var resp voiceEmbedResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/embed", body, &resp); err != nil { + return pb.VoiceEmbedResponse{}, err + } + return pb.VoiceEmbedResponse{Embedding: resp.Embedding, Model: resp.Model}, nil +} + +// --- Stores --------------------------------------------------------------- + +// storeKeysToFloats and storeFloatsToKeys convert between the proto's boxed +// StoresKey/StoresValue slices and the plain [][]float32 / []string the REST +// stores endpoints take, per schema.StoresSet and friends. + +func storeKeysToFloats(keys []*pb.StoresKey) [][]float32 { + out := make([][]float32, len(keys)) + for i, k := range keys { + out[i] = k.GetFloats() + } + return out +} + +func storeFloatsToKeys(keys [][]float32) []*pb.StoresKey { + out := make([]*pb.StoresKey, len(keys)) + for i, k := range keys { + out[i] = &pb.StoresKey{Floats: k} + } + return out +} + +func storeValuesToStrings(values []*pb.StoresValue) []string { + out := make([]string, len(values)) + for i, v := range values { + out[i] = string(v.GetBytes()) + } + return out +} + +func storeStringsToValues(values []string) []*pb.StoresValue { + out := make([]*pb.StoresValue, len(values)) + for i, v := range values { + out[i] = &pb.StoresValue{Bytes: []byte(v)} + } + return out +} + +type storesSetRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` +} + +// StoresSet uses the configured upstream model name as the store name: the +// proxy is loaded per store, the same way it is loaded per model for every +// other method. +func (p *LocalAIProxy) StoresSet(req *pb.StoresSetOptions) error { + body := storesSetRequestBody{ + Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys()), Values: storeValuesToStrings(req.GetValues()), + } + return p.postJSON(context.Background(), "/stores/set", body, nil) +} + +type storesDeleteRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` +} + +func (p *LocalAIProxy) StoresDelete(req *pb.StoresDeleteOptions) error { + body := storesDeleteRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())} + return p.postJSON(context.Background(), "/stores/delete", body, nil) +} + +type storesGetRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` +} + +type storesGetResponseBody struct { + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` +} + +func (p *LocalAIProxy) StoresGet(req *pb.StoresGetOptions) (pb.StoresGetResult, error) { + body := storesGetRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())} + var resp storesGetResponseBody + if err := p.postJSON(context.Background(), "/stores/get", body, &resp); err != nil { + return pb.StoresGetResult{}, err + } + return pb.StoresGetResult{Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values)}, nil +} + +type storesFindRequestBody struct { + Store string `json:"store,omitempty"` + Key []float32 `json:"key"` + Topk int `json:"topk,omitempty"` +} + +type storesFindResponseBody struct { + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` + Similarities []float32 `json:"similarities"` +} + +func (p *LocalAIProxy) StoresFind(req *pb.StoresFindOptions) (pb.StoresFindResult, error) { + body := storesFindRequestBody{Store: p.model(""), Key: req.GetKey().GetFloats(), Topk: int(req.GetTopK())} + var resp storesFindResponseBody + if err := p.postJSON(context.Background(), "/stores/find", body, &resp); err != nil { + return pb.StoresFindResult{}, err + } + return pb.StoresFindResult{ + Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values), Similarities: resp.Similarities, + }, nil +} diff --git a/backend/go/localai-proxy/media_test.go b/backend/go/localai-proxy/media_test.go new file mode 100644 index 000000000..491845fe3 --- /dev/null +++ b/backend/go/localai-proxy/media_test.go @@ -0,0 +1,525 @@ +package main + +import ( + "encoding/base64" + "encoding/json" + "net/http" + "os" + "path/filepath" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/utils" +) + +// b64 is the base64 encoding a spec expects the proxy to have produced from +// a local file's content, matching what the REST endpoints take inline. +func b64(content string) string { + return base64.StdEncoding.EncodeToString([]byte(content)) +} + +var _ = Describe("media methods", func() { + var up *fakeUpstream + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.Close) + }) + + Describe("GenerateImage", func() { + It("posts base64 src/ref_images and decodes data[0].b64_json into Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/images/generations", map[string]any{ + "data": []map[string]any{{"b64_json": b64("png-bytes")}}, + }) + src := writeInput("src.png", "src-bytes") + ref := writeInput("ref.png", "ref-bytes") + dst := filepath.Join(GinkgoT().TempDir(), "out.png") + + Expect(p.GenerateImage(&pb.GenerateImageRequest{ + Width: 64, Height: 32, Step: 20, Seed: 7, + PositivePrompt: "a cat", NegativePrompt: "blurry", + Src: src, RefImages: []string{ref}, Dst: dst, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("png-bytes")) + + req := up.last() + Expect(req.Path).To(Equal("/v1/images/generations")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "a cat")) + Expect(req.JSON).To(HaveKeyWithValue("negative_prompt", "blurry")) + Expect(req.JSON).To(HaveKeyWithValue("size", "64x32")) + Expect(req.JSON).To(HaveKeyWithValue("step", BeNumerically("==", 20))) + Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 7))) + Expect(req.JSON).To(HaveKeyWithValue("response_format", "b64_json")) + Expect(req.JSON).To(HaveKeyWithValue("file", b64("src-bytes"))) + Expect(req.JSON).To(HaveKeyWithValue("ref_images", ConsistOf(b64("ref-bytes")))) + }) + + It("downloads a URL reply relative to the upstream into Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/images/generations", map[string]any{ + "data": []map[string]any{{"url": up.URL + "/generated-images/out.png"}}, + }) + up.script("/generated-images/out.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "downloaded-bytes"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.png") + + Expect(p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst})).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("downloaded-bytes")) + }) + + It("maps an upstream failure and leaves no partial file", func() { + p := loadProxy(up, nil) + up.script("/v1/images/generations", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.png") + + err := p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "x", Dst: dst}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(dst).NotTo(BeAnExistingFile()) + }) + + It("downloads a reply URL from the configured upstream even when its host does not match", func() { + // A reverse proxy, LOCALAI_BASE_URL, or X-Forwarded-Host can all make + // the upstream hand back a URL on a different host than the one this + // proxy is configured with. Only the path should matter: the proxy + // must still fetch it from cfg.base (this fake upstream), with its + // own bearer key, never from the host named in the URL. + p := loadProxy(up, nil) + up.replyJSON("/v1/images/generations", map[string]any{ + "data": []map[string]any{{"url": "http://mismatched-host.invalid:1/generated-images/out.png"}}, + }) + up.script("/generated-images/out.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "downloaded-bytes"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.png") + + Expect(p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst})).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("downloaded-bytes")) + }) + + It("refuses a reply URL whose path is not a known generated-content path", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/images/generations", map[string]any{ + "data": []map[string]any{{"url": up.URL + "/etc/passwd"}}, + }) + dst := filepath.Join(GinkgoT().TempDir(), "out.png") + + err := p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst}) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + Expect(dst).NotTo(BeAnExistingFile()) + }) + }) + + Describe("UpscaleImage", func() { + It("uploads the image as multipart and downloads the returned URL into Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/images/upscale", map[string]any{ + "data": []map[string]any{{"url": up.URL + "/generated-images/big.png"}}, + }) + up.script("/generated-images/big.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "big-bytes"}) + src := writeInput("small.png", "small-bytes") + dst := filepath.Join(GinkgoT().TempDir(), "big.png") + + Expect(p.UpscaleImage(&pb.UpscaleImageRequest{Src: src, Dst: dst, Scale: 4})).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("big-bytes")) + + req := up.recorded()[0] + Expect(req.Path).To(Equal("/v1/images/upscale")) + Expect(req.Fields).To(HaveKeyWithValue("model", "remote-model")) + Expect(req.Fields).To(HaveKeyWithValue("scale", "4")) + Expect(req.Files).To(HaveKeyWithValue("image", "small-bytes")) + }) + }) + + Describe("GenerateVideo", func() { + It("posts base64 media fields and decodes data[0].b64_json into Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/video", map[string]any{ + "data": []map[string]any{{"b64_json": b64("mp4-bytes")}}, + }) + start := writeInput("start.png", "start-bytes") + dst := filepath.Join(GinkgoT().TempDir(), "out.mp4") + + Expect(p.GenerateVideo(&pb.GenerateVideoRequest{ + Prompt: "a dog running", NegativePrompt: "static", StartImage: start, + Width: 512, Height: 288, NumFrames: 24, Fps: 8, Seed: 3, CfgScale: 5.5, Step: 10, + Dst: dst, Params: map[string]string{"motion": "high"}, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("mp4-bytes")) + + req := up.last() + Expect(req.Path).To(Equal("/video")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "a dog running")) + Expect(req.JSON).To(HaveKeyWithValue("start_image", b64("start-bytes"))) + Expect(req.JSON).To(HaveKeyWithValue("num_frames", BeNumerically("==", 24))) + Expect(req.JSON).To(HaveKeyWithValue("fps", BeNumerically("==", 8))) + Expect(req.JSON).To(HaveKeyWithValue("params", HaveKeyWithValue("motion", "high"))) + }) + }) + + Describe("Generate3D", func() { + It("posts the base64 conditioning image and decodes data[0].b64_json into Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/3d/generations", map[string]any{ + "data": []map[string]any{{"b64_json": b64("glb-bytes")}}, + }) + src := writeInput("cond.png", "cond-bytes") + dst := filepath.Join(GinkgoT().TempDir(), "out.glb") + + Expect(p.Generate3D(&pb.Generate3DRequest{ + Src: src, Dst: dst, Seed: 1, Step: 12, CfgScale: 7.5, TextureSteps: 12, Quality: "auto", Background: "keep", + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("glb-bytes")) + + req := up.last() + Expect(req.Path).To(Equal("/3d/generations")) + Expect(req.JSON).To(HaveKeyWithValue("image", b64("cond-bytes"))) + Expect(req.JSON).To(HaveKeyWithValue("quality", "auto")) + Expect(req.JSON).To(HaveKeyWithValue("background", "keep")) + }) + }) + + Describe("Animate3D", func() { + It("base64-encodes non-text inputs, keeps text raw, and writes the reply to Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/3d/animate", map[string]any{ + "data": []map[string]any{{"b64_json": b64("anim-bytes")}}, + "metadata": map[string]any{"frames": float64(30)}, + }) + mesh := writeInput("mesh.glb", "mesh-bytes") + dst := filepath.Join(GinkgoT().TempDir(), "out.glb") + + Expect(p.Animate3D(&pb.Animate3DRequest{ + Inputs: map[string]*pb.AnimationInput{ + "mesh": {Type: "mesh", Data: mesh}, + "prompt": {Type: "text", Data: "wave hello"}, + }, + Dst: dst, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("anim-bytes")) + + req := up.last() + Expect(req.Path).To(Equal("/3d/animate")) + Expect(req.JSON["inputs"]).To(HaveKeyWithValue("mesh", map[string]any{"type": "mesh", "data": b64("mesh-bytes")})) + Expect(req.JSON["inputs"]).To(HaveKeyWithValue("prompt", map[string]any{"type": "text", "data": "wave hello"})) + }) + + It("Animate3DWithMetadata returns the upstream metadata and writes Dst", func() { + p := loadProxy(up, nil) + up.replyJSON("/3d/animate", map[string]any{ + "data": []map[string]any{{"b64_json": b64("anim-bytes")}}, + "metadata": map[string]any{"frames": float64(30)}, + }) + dst := filepath.Join(GinkgoT().TempDir(), "out.glb") + + metadata, err := p.Animate3DWithMetadata(&pb.Animate3DRequest{ + Inputs: map[string]*pb.AnimationInput{"prompt": {Type: "text", Data: "wave"}}, + Dst: dst, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(string(metadata)).To(MatchJSON(`{"frames":30}`)) + Expect(os.ReadFile(dst)).To(BeEquivalentTo("anim-bytes")) + }) + }) + + // pngBytes is a minimal payload with a real PNG magic number, so + // http.DetectContentType (which toDataURI uses) sniffs "image/png" the + // way it would for a real image, and decodeAndCheckImage below can tell + // this test wrote a valid data URI apart from one that merely echoes bare + // base64. + pngBytes := append([]byte("\x89PNG\r\n\x1a\n"), []byte("fake-png-body")...) + + // decodeAndCheckImage extracts field from an upstream request body and + // decodes it exactly the way the real REST handlers do — with + // utils.GetContentURIAsBase64 (core/http/endpoints/localai/images.go's + // decodeImageInput calls the same function) — so a spec here fails the + // same way a live upstream would if the proxy ever regressed to sending + // bare base64, which that function rejects outright. + decodeAndCheckImage := func(body map[string]any, field string, want []byte) { + raw, _ := body[field].(string) + decoded, err := utils.GetContentURIAsBase64(raw) + ExpectWithOffset(1, err).NotTo(HaveOccurred(), "the real upstream would reject %s the same way", field) + got, err := base64.StdEncoding.DecodeString(decoded) + ExpectWithOffset(1, err).NotTo(HaveOccurred()) + ExpectWithOffset(1, got).To(Equal(want)) + } + + Describe("Detect", func() { + It("wraps the bare base64 image as a data URI the real upstream decoder accepts, and maps detections", func() { + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + var body map[string]any + Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed()) + decodeAndCheckImage(body, "image", pngBytes) + Expect(body["prompt"]).To(Equal("cat")) + + mask := base64.StdEncoding.EncodeToString([]byte("png-mask")) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "detections": []map[string]any{ + {"x": 1.0, "y": 2.0, "width": 3.0, "height": 4.0, "confidence": 0.9, "class_name": "cat", "mask": mask}, + }, + }) + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + res, err := p.Detect(&pb.DetectOptions{ + Src: base64.StdEncoding.EncodeToString(pngBytes), Prompt: "cat", Threshold: 0.5, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Detections).To(HaveLen(1)) + Expect(res.Detections[0].ClassName).To(Equal("cat")) + Expect(res.Detections[0].Confidence).To(BeNumerically("~", 0.9, 1e-6)) + Expect(res.Detections[0].Mask).To(Equal([]byte("png-mask"))) + }) + }) + + Describe("Depth", func() { + It("wraps the bare base64 image as a data URI and maps the full depth response", func() { + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + var body map[string]any + Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed()) + decodeAndCheckImage(body, "image", pngBytes) + Expect(body["include_depth"]).To(Equal(true)) + + colors := base64.StdEncoding.EncodeToString([]byte("rgb")) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "width": 2, "height": 1, "depth": []float64{0.1, 0.2}, + "point_colors": colors, "is_metric": true, + }) + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + res, err := p.Depth(&pb.DepthRequest{Src: base64.StdEncoding.EncodeToString(pngBytes), IncludeDepth: true}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Width).To(Equal(int32(2))) + Expect(res.Height).To(Equal(int32(1))) + Expect(res.Depth).To(Equal([]float32{0.1, 0.2})) + Expect(res.PointColors).To(Equal([]byte("rgb"))) + Expect(res.IsMetric).To(BeTrue()) + }) + + It("refuses a request for exports or a dst directory without calling the upstream, so failover moves on", func() { + p := loadProxy(up, nil) + + _, exportsErr := p.Depth(&pb.DepthRequest{Src: "irrelevant", Exports: []string{"glb"}}) + Expect(codeOf(exportsErr)).To(Equal(codes.Unimplemented)) + + _, dstErr := p.Depth(&pb.DepthRequest{Src: "irrelevant", Dst: "/some/output/dir"}) + Expect(codeOf(dstErr)).To(Equal(codes.Unimplemented)) + + Expect(up.recorded()).To(BeEmpty(), "an exports/dst request must never reach the upstream") + }) + }) + + Describe("FaceVerify", func() { + It("wraps both bare base64 images as data URIs and maps the response, including liveness fields", func() { + img2Bytes := append([]byte("\x89PNG\r\n\x1a\n"), []byte("other-face")...) + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + var body map[string]any + Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed()) + decodeAndCheckImage(body, "img1", pngBytes) + decodeAndCheckImage(body, "img2", img2Bytes) + Expect(body["anti_spoofing"]).To(Equal(true)) + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "verified": true, "distance": 0.1, "threshold": 0.4, "confidence": 92.0, "model": "buffalo_l", + "img1_area": map[string]any{"x": 1.0, "y": 2.0, "w": 3.0, "h": 4.0}, + "img2_area": map[string]any{"x": 5.0, "y": 6.0, "w": 7.0, "h": 8.0}, + "img1_is_real": true, "img1_antispoof_score": 0.99, + "img2_is_real": false, "img2_antispoof_score": 0.1, + }) + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + res, err := p.FaceVerify(&pb.FaceVerifyRequest{ + Img1: base64.StdEncoding.EncodeToString(pngBytes), Img2: base64.StdEncoding.EncodeToString(img2Bytes), + AntiSpoofing: true, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Verified).To(BeTrue()) + Expect(res.Model).To(Equal("buffalo_l")) + Expect(res.Img1Area.W).To(BeNumerically("==", 3)) + Expect(res.Img1IsReal).To(BeTrue()) + Expect(res.Img2IsReal).To(BeFalse()) + }) + }) + + Describe("FaceAnalyze", func() { + It("wraps the bare base64 image as a data URI and maps per-face demographic attributes", func() { + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + var body map[string]any + Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed()) + decodeAndCheckImage(body, "img", pngBytes) + Expect(body["actions"]).To(ConsistOf("age", "gender")) + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "faces": []map[string]any{ + { + "region": map[string]any{"x": 1.0, "y": 2.0, "w": 3.0, "h": 4.0}, + "face_confidence": 0.95, "age": 30.0, "dominant_gender": "Man", + "gender": map[string]any{"Man": 0.9, "Woman": 0.1}, + }, + }, + }) + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + res, err := p.FaceAnalyze(&pb.FaceAnalyzeRequest{ + Img: base64.StdEncoding.EncodeToString(pngBytes), Actions: []string{"age", "gender"}, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Faces).To(HaveLen(1)) + Expect(res.Faces[0].DominantGender).To(Equal("Man")) + Expect(res.Faces[0].Gender).To(HaveKeyWithValue("Man", Equal(float32(0.9)))) + }) + }) + + Describe("VoiceVerify", func() { + It("base64-encodes both audio clips and maps the response", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/voice/verify", map[string]any{ + "verified": true, "distance": 0.2, "threshold": 0.5, "confidence": 88.0, "model": "ecapa", + }) + a1 := writeInput("a1.wav", "audio-one") + a2 := writeInput("a2.wav", "audio-two") + + res, err := p.VoiceVerify(&pb.VoiceVerifyRequest{Audio1: a1, Audio2: a2, Threshold: 0.5}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Verified).To(BeTrue()) + Expect(res.Model).To(Equal("ecapa")) + + req := up.last() + Expect(req.Path).To(Equal("/v1/voice/verify")) + Expect(req.JSON).To(HaveKeyWithValue("audio1", b64("audio-one"))) + Expect(req.JSON).To(HaveKeyWithValue("audio2", b64("audio-two"))) + }) + }) + + Describe("VoiceAnalyze", func() { + It("base64-encodes the audio clip and maps demographic segments", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/voice/analyze", map[string]any{ + "segments": []map[string]any{ + {"start": 0.0, "end": 1.5, "age": 25.0, "dominant_gender": "Woman"}, + }, + }) + audio := writeInput("clip.wav", "clip-bytes") + + res, err := p.VoiceAnalyze(&pb.VoiceAnalyzeRequest{Audio: audio, Actions: []string{"age"}}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Segments).To(HaveLen(1)) + Expect(res.Segments[0].DominantGender).To(Equal("Woman")) + + req := up.last() + Expect(req.Path).To(Equal("/v1/voice/analyze")) + Expect(req.JSON).To(HaveKeyWithValue("audio", b64("clip-bytes"))) + }) + }) + + Describe("VoiceEmbed", func() { + It("base64-encodes the audio clip and returns the embedding", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/voice/embed", map[string]any{ + "embedding": []float64{0.1, 0.2, 0.3}, "dim": 3.0, "model": "ecapa", + }) + audio := writeInput("clip.wav", "clip-bytes") + + res, err := p.VoiceEmbed(&pb.VoiceEmbedRequest{Audio: audio}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Embedding).To(Equal([]float32{0.1, 0.2, 0.3})) + Expect(res.Model).To(Equal("ecapa")) + + req := up.last() + Expect(req.Path).To(Equal("/v1/voice/embed")) + Expect(req.JSON).To(HaveKeyWithValue("audio", b64("clip-bytes"))) + }) + }) + + Describe("Stores", func() { + It("StoresSet posts the store name, keys and values", func() { + p := loadProxy(up, nil) + up.script("/stores/set", scriptedResponse{Status: http.StatusOK}) + + Expect(p.StoresSet(&pb.StoresSetOptions{ + Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}}, + Values: []*pb.StoresValue{{Bytes: []byte("v1")}}, + })).To(Succeed()) + + req := up.last() + Expect(req.Path).To(Equal("/stores/set")) + Expect(req.JSON).To(HaveKeyWithValue("store", "remote-model")) + Expect(req.JSON).To(HaveKeyWithValue("values", ConsistOf("v1"))) + }) + + It("StoresDelete posts the store name and keys", func() { + p := loadProxy(up, nil) + up.script("/stores/delete", scriptedResponse{Status: http.StatusOK}) + + Expect(p.StoresDelete(&pb.StoresDeleteOptions{Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}}})).To(Succeed()) + + req := up.last() + Expect(req.Path).To(Equal("/stores/delete")) + Expect(req.JSON).To(HaveKeyWithValue("store", "remote-model")) + }) + + It("StoresGet posts keys and maps the returned keys/values", func() { + p := loadProxy(up, nil) + up.replyJSON("/stores/get", map[string]any{ + "keys": [][]float64{{1, 2}}, "values": []string{"v1"}, + }) + + res, err := p.StoresGet(&pb.StoresGetOptions{Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}}}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Keys).To(HaveLen(1)) + Expect(res.Keys[0].Floats).To(Equal([]float32{1, 2})) + Expect(res.Values).To(HaveLen(1)) + Expect(res.Values[0].Bytes).To(Equal([]byte("v1"))) + }) + + It("StoresFind posts the query key/topk and maps keys/values/similarities", func() { + p := loadProxy(up, nil) + up.replyJSON("/stores/find", map[string]any{ + "keys": [][]float64{{1, 2}}, "values": []string{"v1"}, "similarities": []float64{0.9}, + }) + + res, err := p.StoresFind(&pb.StoresFindOptions{Key: &pb.StoresKey{Floats: []float32{1, 2}}, TopK: 5}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.Similarities).To(Equal([]float32{0.9})) + + req := up.last() + Expect(req.Path).To(Equal("/stores/find")) + Expect(req.JSON).To(HaveKeyWithValue("topk", BeNumerically("==", 5))) + Expect(req.JSON).To(HaveKeyWithValue("key", ConsistOf(BeNumerically("==", 1), BeNumerically("==", 2)))) + }) + }) +}) diff --git a/backend/go/localai-proxy/package.sh b/backend/go/localai-proxy/package.sh new file mode 100755 index 000000000..7bc9f64aa --- /dev/null +++ b/backend/go/localai-proxy/package.sh @@ -0,0 +1,13 @@ +#!/bin/bash + +# Script to copy the localai-proxy binary into the package dir for the +# final Dockerfile stage. Mirrors backend/go/local-store/package.sh — +# no extra runtime libs needed since the backend is pure Go. + +set -e + +CURDIR=$(dirname "$(realpath $0)") + +mkdir -p $CURDIR/package +cp -avf $CURDIR/localai-proxy $CURDIR/package/ +cp -rfv $CURDIR/run.sh $CURDIR/package/ diff --git a/backend/go/localai-proxy/proxy.go b/backend/go/localai-proxy/proxy.go new file mode 100644 index 000000000..a1d58bb02 --- /dev/null +++ b/backend/go/localai-proxy/proxy.go @@ -0,0 +1,239 @@ +package main + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + "sync/atomic" + "time" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + "github.com/mudler/LocalAI/pkg/grpc/base" + "github.com/mudler/LocalAI/pkg/grpc/grpcerrors" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/httpclient" +) + +const ( + backendName = "localai-proxy" + + // realtimePipelineOption names the upstream realtime pipeline that serves + // live transcription sessions (options: ["realtime_pipeline:"]). + realtimePipelineOption = "realtime_pipeline:" +) + +// LocalAIProxy serves backend methods by calling a remote LocalAI's REST API. +// base.SingleThread is not embedded: every call is an independent HTTP +// request, so serialising them would only add latency. +type LocalAIProxy struct { + base.Base + + cfg atomic.Pointer[proxyConfig] + client *http.Client +} + +type proxyConfig struct { + base string // upstream base URL without a trailing slash + upstreamModel string // model name sent upstream + apiKey string + realtimePipeline string + timeout time.Duration // per-request limit for non-streaming calls; 0 = none +} + +func NewLocalAIProxy() *LocalAIProxy { + // httpclient.New refuses redirects: the upstream is one configured + // LocalAI, so a 3xx means misconfiguration or a hijacked host, and + // following it would replay the bearer key to an unvetted host. It also + // sets no body deadline, so long SSE streams are not cut short. + return &LocalAIProxy{client: httpclient.New()} +} + +// Load refuses a model without proxy options so greedy backend probing, +// which tries every installed backend on a model file, never selects it. +func (p *LocalAIProxy) Load(opts *pb.ModelOptions) error { + po := opts.GetProxy() + if po == nil { + return errors.New("localai-proxy: Load requires proxy options (proxy.upstream_url)") + } + raw := po.GetUpstreamUrl() + if raw == "" { + return errors.New("localai-proxy: proxy.upstream_url is required") + } + u, err := url.ParseRequestURI(raw) + if err != nil { + return fmt.Errorf("localai-proxy: proxy.upstream_url %q invalid: %w", raw, err) + } + if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" { + return fmt.Errorf("localai-proxy: proxy.upstream_url %q must be an http(s) URL with a host", raw) + } + + // Every request path starts with /v1, so upstream_url is the server's + // root. A URL copied from an OpenAI-style config ends in /v1 (or a full + // endpoint): cut the path at /v1 as the failover prober does, or the + // prober reports the target healthy while every request 404s. + base := strings.TrimRight(raw, "/") + if i := strings.Index(u.Path, "/v1"); i >= 0 { + base = strings.TrimRight(u.Scheme+"://"+u.Host+u.Path[:i], "/") + xlog.Warn("localai-proxy: proxy.upstream_url should be the server root; ignoring its /v1 path", + "upstream_url", raw, "using", base) + } + + // There is no translate mode: the upstream always speaks LocalAI's API. + if po.GetMode() != "" || po.GetProvider() != "" { + xlog.Warn("localai-proxy: proxy.mode and proxy.provider are ignored", + "mode", po.GetMode(), "provider", po.GetProvider()) + } + + key, err := resolveAPIKey(po.GetApiKeyEnv(), po.GetApiKeyFile()) + if err != nil { + return err + } + + model := po.GetUpstreamModel() + if model == "" { + model = opts.GetModel() + } + if model == "" { + xlog.Warn("localai-proxy: no upstream model name; set proxy.upstream_model") + } + + var pipeline string + for _, o := range opts.GetOptions() { + if v, ok := strings.CutPrefix(o, realtimePipelineOption); ok { + pipeline = strings.TrimSpace(v) + } + } + + var timeout time.Duration + if s := po.GetRequestTimeoutSeconds(); s > 0 { + timeout = time.Duration(s) * time.Second + } + + p.cfg.Store(&proxyConfig{ + base: base, + upstreamModel: model, + apiKey: key, + realtimePipeline: pipeline, + timeout: timeout, + }) + xlog.Info("localai-proxy: ready", "upstream", base, "upstream_model", model, + "has_key", key != "", "realtime_pipeline", pipeline) + return nil +} + +// config returns the loaded configuration, or the typed not-loaded error so +// callers see FailedPrecondition instead of a nil dereference. +func (p *LocalAIProxy) config() (*proxyConfig, error) { + cfg := p.cfg.Load() + if cfg == nil { + return nil, grpcerrors.ModelNotLoaded(backendName) + } + return cfg, nil +} + +// model returns the model name to send upstream. The configured name wins so +// every method targets the same upstream model; req (a model named by the +// request itself) is only a fallback for configs that resolved no name. +func (p *LocalAIProxy) model(req string) string { + if cfg := p.cfg.Load(); cfg != nil && cfg.upstreamModel != "" { + return cfg.upstreamModel + } + return req +} + +// resolveAPIKey mirrors config.ProxyConfig.ResolveAPIKey (and cloud-proxy's +// copy). Duplicated so the backend binary does not depend on core's layout. +func resolveAPIKey(envName, filePath string) (string, error) { + if envName != "" { + v := os.Getenv(envName) + if v == "" { + return "", fmt.Errorf("localai-proxy: api_key_env %q is unset", envName) + } + return v, nil + } + if filePath != "" { + // #nosec G304 -- api_key_file comes from the operator's model config (passed by core as a backend option), not from a request + b, err := os.ReadFile(filepath.Clean(filePath)) + if err != nil { + return "", fmt.Errorf("localai-proxy: read api_key_file %q: %w", filePath, err) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} + +// unimplemented is the error for methods LocalAI's REST API cannot serve. +// Failover reads gRPC Unimplemented as a capability gap and moves to the next +// target without marking this one unhealthy. +func unimplemented(method string) error { + return status.Errorf(codes.Unimplemented, "localai-proxy: %s has no upstream counterpart", method) +} + +func (p *LocalAIProxy) AudioEncode(*pb.AudioEncodeRequest) (*pb.AudioEncodeResult, error) { + return nil, unimplemented("AudioEncode") +} + +func (p *LocalAIProxy) AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error) { + return nil, unimplemented("AudioDecode") +} + +// AudioToAudioStream closes out because the gRPC server drains it until +// closed; leaving it open would hang the call. +func (p *LocalAIProxy) AudioToAudioStream(_ <-chan *pb.AudioToAudioRequest, out chan<- *pb.AudioToAudioResponse) error { + close(out) + return unimplemented("AudioToAudioStream") +} + +func (p *LocalAIProxy) TokenClassify(context.Context, *pb.TokenClassifyRequest) (*pb.TokenClassifyResponse, error) { + return nil, unimplemented("TokenClassify") +} + +func (p *LocalAIProxy) ModelMetadata(*pb.ModelOptions) (*pb.ModelMetadataResponse, error) { + return nil, unimplemented("ModelMetadata") +} + +func (p *LocalAIProxy) StartFineTune(*pb.FineTuneRequest) (*pb.FineTuneJobResult, error) { + return nil, unimplemented("StartFineTune") +} + +// FineTuneProgress closes the channel: the gRPC server waits for it to close +// before returning, and base.Base leaves it open. +func (p *LocalAIProxy) FineTuneProgress(_ *pb.FineTuneProgressRequest, updates chan *pb.FineTuneProgressUpdate) error { + close(updates) + return unimplemented("FineTuneProgress") +} + +func (p *LocalAIProxy) StopFineTune(*pb.FineTuneStopRequest) error { + return unimplemented("StopFineTune") +} + +func (p *LocalAIProxy) ListCheckpoints(*pb.ListCheckpointsRequest) (*pb.ListCheckpointsResponse, error) { + return nil, unimplemented("ListCheckpoints") +} + +func (p *LocalAIProxy) ExportModel(*pb.ExportModelRequest) error { + return unimplemented("ExportModel") +} + +func (p *LocalAIProxy) StartQuantization(*pb.QuantizationRequest) (*pb.QuantizationJobResult, error) { + return nil, unimplemented("StartQuantization") +} + +// QuantizationProgress closes the channel for the same reason as +// FineTuneProgress. +func (p *LocalAIProxy) QuantizationProgress(_ *pb.QuantizationProgressRequest, updates chan *pb.QuantizationProgressUpdate) error { + close(updates) + return unimplemented("QuantizationProgress") +} + +func (p *LocalAIProxy) StopQuantization(*pb.QuantizationStopRequest) error { + return unimplemented("StopQuantization") +} diff --git a/backend/go/localai-proxy/run.sh b/backend/go/localai-proxy/run.sh new file mode 100755 index 000000000..f2023e3ec --- /dev/null +++ b/backend/go/localai-proxy/run.sh @@ -0,0 +1,6 @@ +#!/bin/bash +set -ex + +CURDIR=$(dirname "$(realpath "$0")") + +exec "$CURDIR"/localai-proxy "$@" diff --git a/backend/go/localai-proxy/text.go b/backend/go/localai-proxy/text.go new file mode 100644 index 000000000..531718267 --- /dev/null +++ b/backend/go/localai-proxy/text.go @@ -0,0 +1,447 @@ +package main + +import ( + "bufio" + "context" + "encoding/json" + "math" + "strings" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// textRequest is the body for /v1/chat/completions (Messages) and +// /v1/completions (Prompt). Zero sampling values are omitted so the upstream +// model's own config defaults apply, as they would for a direct caller. +// Temperature is the exception: 0 is a real choice (greedy decoding), and +// core always fills it from the model config, so it is always sent. +type textRequest struct { + Model string `json:"model"` + Messages []chatMessage `json:"messages,omitempty"` + Prompt string `json:"prompt,omitempty"` + Stream bool `json:"stream,omitempty"` + MaxTokens int32 `json:"max_tokens,omitempty"` + Temperature float32 `json:"temperature"` + TopP float32 `json:"top_p,omitempty"` + TopK int32 `json:"top_k,omitempty"` + Seed int32 `json:"seed,omitempty"` + Stop []string `json:"stop,omitempty"` + Tools json.RawMessage `json:"tools,omitempty"` + ToolChoice json.RawMessage `json:"tool_choice,omitempty"` + // StreamOptions is set on streamed requests: the upstream sends the + // usage trailer only when include_usage asks for it. + StreamOptions *streamOptions `json:"stream_options,omitempty"` +} + +type streamOptions struct { + IncludeUsage bool `json:"include_usage"` +} + +type chatMessage struct { + Role string `json:"role"` + Content string `json:"content"` + Name string `json:"name,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolCalls []toolCall `json:"tool_calls,omitempty"` +} + +type toolCall struct { + Index int `json:"index"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function struct { + Name string `json:"name,omitempty"` + Arguments string `json:"arguments,omitempty"` + } `json:"function"` +} + +// textChoice covers both endpoints and both shapes: chat replies fill Message +// (or Delta when streaming), completions fill Text. +type textChoice struct { + Text string `json:"text"` + Message choiceDelta `json:"message"` + Delta choiceDelta `json:"delta"` +} + +type choiceDelta struct { + Content string `json:"content"` + // LocalAI names the field "reasoning"; other OpenAI-compatible servers + // use "reasoning_content". + Reasoning string `json:"reasoning"` + ReasoningContent string `json:"reasoning_content"` + ToolCalls []toolCall `json:"tool_calls"` +} + +type textResponse struct { + // Error is set on the frame LocalAI sends when generation fails after + // the stream has started (followed by [DONE]). + Error *struct { + Message string `json:"message"` + } `json:"error"` + Choices []textChoice `json:"choices"` + Usage *struct { + PromptTokens int32 `json:"prompt_tokens"` + CompletionTokens int32 `json:"completion_tokens"` + } `json:"usage"` +} + +// textRequest picks the chat endpoint when core sent structured messages +// (the model uses the tokenizer template, so the upstream must template +// too); otherwise core already rendered the prompt and completions takes it +// verbatim. +func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string, textRequest) { + if dropped := unforwardedFields(opts); len(dropped) > 0 { + xlog.Warn("localai-proxy: request fields are not forwarded upstream", "fields", dropped) + } + req := textRequest{ + Model: p.model(""), + Stream: stream, + MaxTokens: opts.GetTokens(), + Temperature: opts.GetTemperature(), + TopP: opts.GetTopP(), + TopK: opts.GetTopK(), + Seed: opts.GetSeed(), + Stop: opts.GetStopPrompts(), + Tools: rawJSON(opts.GetTools()), + ToolChoice: rawJSON(opts.GetToolChoice()), + } + if stream { + req.StreamOptions = &streamOptions{IncludeUsage: true} + } + if len(opts.GetMessages()) == 0 { + req.Prompt = opts.GetPrompt() + return "/v1/completions", req + } + for _, m := range opts.GetMessages() { + msg := chatMessage{ + Role: m.GetRole(), + Content: m.GetContent(), + Name: m.GetName(), + ToolCallID: m.GetToolCallId(), + } + // A previous assistant turn carries its tool calls as a JSON string. + if tc := m.GetToolCalls(); tc != "" { + if err := json.Unmarshal([]byte(tc), &msg.ToolCalls); err != nil { + xlog.Debug("localai-proxy: drop malformed tool_calls on message", "error", err) + } + } + req.Messages = append(req.Messages, msg) + } + return "/v1/chat/completions", req +} + +// unforwardedFields names the request inputs the REST text endpoints cannot +// carry from here: a grammar core compiled locally, and media that core hands +// over as local paths or base64 outside the messages. They are dropped, so +// say so instead of letting the answer silently ignore them. +func unforwardedFields(opts *pb.PredictOptions) []string { + var out []string + if opts.GetGrammar() != "" { + out = append(out, "grammar") + } + if len(opts.GetImages()) > 0 { + out = append(out, "images") + } + if len(opts.GetAudios()) > 0 { + out = append(out, "audios") + } + if len(opts.GetVideos()) > 0 { + out = append(out, "videos") + } + return out +} + +// rawJSON passes a JSON string through untouched, or omits it when it is +// empty or invalid rather than failing the whole request. +func rawJSON(s string) json.RawMessage { + if s == "" || !json.Valid([]byte(s)) { + return nil + } + return json.RawMessage(s) +} + +// replyFromChoice builds the Reply for one choice. Reasoning and tool calls +// travel as ChatDeltas, the same shape the llama.cpp autoparser emits, so +// core handles them without parsing the text again. +func replyFromChoice(c textChoice, streaming bool) *pb.Reply { + d := c.Message + if streaming { + d = c.Delta + } + content := d.Content + if content == "" { + content = c.Text + } + reasoning := d.Reasoning + if reasoning == "" { + reasoning = d.ReasoningContent + } + + reply := &pb.Reply{Message: []byte(content)} + // Streaming chunks always carry a delta, like the autoparser's. A + // complete reply only needs one for reasoning or tool calls; its + // content must then be in the delta too, because core takes the + // content from the deltas once they carry anything. + if reasoning == "" && len(d.ToolCalls) == 0 && (!streaming || content == "") { + return reply + } + delta := &pb.ChatDelta{Content: content, ReasoningContent: reasoning} + for _, tc := range d.ToolCalls { + delta.ToolCalls = append(delta.ToolCalls, &pb.ToolCallDelta{ + Index: clampInt32(tc.Index), + Id: tc.ID, + Name: tc.Function.Name, + Arguments: tc.Function.Arguments, + }) + } + reply.ChatDeltas = []*pb.ChatDelta{delta} + return reply +} + +func (p *LocalAIProxy) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) { + return p.PredictRichContext(context.Background(), opts) +} + +// PredictRichContext is PredictRich bound to the gRPC call: when the caller +// goes away, the upstream request is cancelled too. +func (p *LocalAIProxy) PredictRichContext(ctx context.Context, opts *pb.PredictOptions) (*pb.Reply, error) { + path, body := p.textRequest(opts, false) + var resp textResponse + if err := p.postJSON(ctx, path, body, &resp); err != nil { + return nil, err + } + if len(resp.Choices) == 0 { + return nil, status.Errorf(codes.Internal, "localai-proxy: upstream %s returned no choices", path) + } + reply := replyFromChoice(resp.Choices[0], false) + if resp.Usage != nil { + reply.PromptTokens = resp.Usage.PromptTokens + reply.Tokens = resp.Usage.CompletionTokens + } + return reply, nil +} + +// PredictStreamRich sends one Reply per upstream SSE delta. It does not close +// results: the gRPC server does, after this returns. +func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<- *pb.Reply) error { + return p.PredictStreamRichContext(context.Background(), opts, results) +} + +// PredictStreamRichContext is PredictStreamRich bound to the gRPC stream: a +// client that disconnects, or a failover that abandons this target, stops the +// upstream generation instead of letting it run (and bill) to the end. +func (p *LocalAIProxy) PredictStreamRichContext(ctx context.Context, opts *pb.PredictOptions, results chan<- *pb.Reply) error { + path, body := p.textRequest(opts, true) + resp, err := p.postStream(ctx, path, body) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + scanner := bufio.NewScanner(resp.Body) + // A single frame can carry a long tool-call argument chunk. + 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 == "[DONE]" { + return nil + } + if payload == "" { + continue + } + var chunk textResponse + if err := json.Unmarshal([]byte(payload), &chunk); err != nil { + xlog.Debug("localai-proxy: skip malformed SSE frame", "path", path, "error", err) + continue + } + if chunk.Error != nil { + // The upstream failed mid-generation. Returning nil would turn a + // cut-off answer into a success; Unavailable lets failover and + // the client see the failure. + xlog.Warn("localai-proxy: upstream stream error", "path", path, "error", chunk.Error.Message) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", path, chunk.Error.Message) + } + if chunk.Usage != nil && len(chunk.Choices) == 0 { + results <- &pb.Reply{PromptTokens: chunk.Usage.PromptTokens, Tokens: chunk.Usage.CompletionTokens} + continue + } + for _, c := range chunk.Choices { + reply := replyFromChoice(c, true) + if len(reply.GetMessage()) == 0 && len(reply.GetChatDeltas()) == 0 { + continue // role-only or finish frames carry nothing to emit + } + results <- reply + } + } + if err := scanner.Err(); err != nil { + return transportError(path, err) + } + return nil +} + +// Predict is the legacy string path; the gRPC server prefers PredictRich. +func (p *LocalAIProxy) Predict(opts *pb.PredictOptions) (string, error) { + reply, err := p.PredictRich(opts) + if err != nil { + return "", err + } + return string(reply.GetMessage()), nil +} + +// PredictStream is the legacy string stream. Unlike PredictStreamRich it +// owns and closes results, per the AIModel contract. +func (p *LocalAIProxy) PredictStream(opts *pb.PredictOptions, results chan string) error { + defer close(results) + rich := make(chan *pb.Reply) + errCh := make(chan error, 1) + go func() { + errCh <- p.PredictStreamRich(opts, rich) + close(rich) + }() + for reply := range rich { + if msg := reply.GetMessage(); len(msg) > 0 { + results <- string(msg) + } + } + return <-errCh +} + +func (p *LocalAIProxy) Embeddings(opts *pb.PredictOptions) ([]float32, error) { + var resp struct { + Data []struct { + Embedding []float32 `json:"embedding"` + } `json:"data"` + } + body := map[string]any{"model": p.model(""), "input": opts.GetEmbeddings()} + // Core sends tokenized input in EmbeddingTokens and leaves Embeddings + // empty; a list of token lists is how the REST API takes tokens. + if tokens := opts.GetEmbeddingTokens(); len(tokens) > 0 { + body["input"] = [][]int32{tokens} + } + if err := p.postJSON(context.Background(), "/v1/embeddings", body, &resp); err != nil { + return nil, err + } + if len(resp.Data) == 0 || len(resp.Data[0].Embedding) == 0 { + return nil, status.Error(codes.Internal, "localai-proxy: upstream /v1/embeddings returned no embedding") + } + return resp.Data[0].Embedding, nil +} + +func (p *LocalAIProxy) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) { + body := map[string]any{ + "model": p.model(""), + "query": in.GetQuery(), + "documents": in.GetDocuments(), + } + // TopN 0 means "score every document" (the router's reranker sends it); + // upstream rejects top_n < 1, and an absent top_n means the same thing. + if n := in.GetTopN(); n > 0 { + body["top_n"] = n + } + var resp struct { + Usage struct { + TotalTokens int32 `json:"total_tokens"` + PromptTokens int32 `json:"prompt_tokens"` + } `json:"usage"` + Results []struct { + Index int32 `json:"index"` + Document struct { + Text string `json:"text"` + } `json:"document"` + RelevanceScore float32 `json:"relevance_score"` + } `json:"results"` + } + if err := p.postJSON(ctx, "/v1/rerank", body, &resp); err != nil { + return nil, err + } + out := &pb.RerankResult{Usage: &pb.Usage{TotalTokens: resp.Usage.TotalTokens, PromptTokens: resp.Usage.PromptTokens}} + for _, r := range resp.Results { + out.Results = append(out.Results, &pb.DocumentResult{Index: r.Index, Text: r.Document.Text, RelevanceScore: r.RelevanceScore}) + } + return out, nil +} + +// TokenizeString uses the upstream model's tokenizer, which is the one that +// matters: the upstream is where the tokens will be spent. +func (p *LocalAIProxy) TokenizeString(opts *pb.PredictOptions) (pb.TokenizationResponse, error) { + var resp struct { + Tokens []int32 `json:"tokens"` + } + body := map[string]any{"model": p.model(""), "content": opts.GetPrompt()} + if err := p.postJSON(context.Background(), "/v1/tokenize", body, &resp); err != nil { + return pb.TokenizationResponse{}, err + } + return pb.TokenizationResponse{Length: clampInt32(len(resp.Tokens)), Tokens: resp.Tokens}, nil +} + +func (p *LocalAIProxy) Detokenize(in *pb.DetokenizeRequest) (pb.DetokenizeResponse, error) { + var resp struct { + Content string `json:"content"` + } + body := map[string]any{"model": p.model(""), "tokens": in.GetTokens()} + if err := p.postJSON(context.Background(), "/v1/detokenize", body, &resp); err != nil { + return pb.DetokenizeResponse{}, err + } + return pb.DetokenizeResponse{Content: resp.Content}, nil +} + +// Score forwards plain candidate scoring to /api/score. That endpoint has no +// decision-pipeline fields, so a question_type request would silently become +// plain scoring upstream; refuse it as a capability gap instead. +func (p *LocalAIProxy) Score(ctx context.Context, in *pb.ScoreRequest) (*pb.ScoreResponse, error) { + if in.GetQuestionType() != "" { + return nil, unimplemented("Score with question_type") + } + body := map[string]any{ + "model": p.model(""), + "prompt": in.GetPrompt(), + "candidates": in.GetCandidates(), + "include_token_logprobs": in.GetIncludeTokenLogprobs(), + "length_normalize": in.GetLengthNormalize(), + } + var resp struct { + Candidates []struct { + LogProb float64 `json:"log_prob"` + LengthNormalizedLogProb float64 `json:"length_normalized_log_prob"` + NumTokens int32 `json:"num_tokens"` + Tokens []struct { + Token string `json:"token"` + LogProb float64 `json:"log_prob"` + } `json:"tokens"` + } `json:"candidates"` + } + if err := p.postJSON(ctx, "/api/score", body, &resp); err != nil { + return nil, err + } + out := &pb.ScoreResponse{} + for _, c := range resp.Candidates { + cs := &pb.CandidateScore{LogProb: c.LogProb, LengthNormalizedLogProb: c.LengthNormalizedLogProb, NumTokens: c.NumTokens} + for _, t := range c.Tokens { + cs.Tokens = append(cs.Tokens, &pb.TokenLogProb{Token: t.Token, LogProb: t.LogProb}) + } + out.Candidates = append(out.Candidates, cs) + } + return out, nil +} + +// clampInt32 narrows an upstream-supplied int (token counts, tool-call +// indexes) to the int32 the gRPC protocol carries. The upstream is another +// server, so an absurd value must saturate rather than wrap to a negative or +// small number that core would take at face value. +func clampInt32(n int) int32 { + switch { + case n > math.MaxInt32: + return math.MaxInt32 + case n < math.MinInt32: + return math.MinInt32 + } + return int32(n) +} diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go new file mode 100644 index 000000000..c8e8fad0d --- /dev/null +++ b/backend/go/localai-proxy/text_test.go @@ -0,0 +1,629 @@ +package main + +import ( + "context" + "math" + "net/http" + "os" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + grpc "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +func codeOf(err error) codes.Code { + st, ok := status.FromError(err) + if !ok { + return codes.Unknown + } + return st.Code() +} + +var _ = Describe("localai-proxy", func() { + var up *fakeUpstream + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.Close) + }) + + Describe("Load", func() { + It("refuses a model without proxy options", func() { + err := NewLocalAIProxy().Load(&pb.ModelOptions{Model: "m"}) + Expect(err).To(MatchError(ContainSubstring("proxy"))) + }) + + It("refuses a missing or invalid upstream_url", func() { + p := NewLocalAIProxy() + Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{}})).NotTo(Succeed()) + Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "not a url"}})).NotTo(Succeed()) + Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "ftp://host"}})).NotTo(Succeed()) + }) + + It("refuses an api_key_env that is unset", func() { + err := NewLocalAIProxy().Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{ + UpstreamUrl: up.URL, ApiKeyEnv: "LOCALAI_PROXY_TEST_UNSET_KEY", + }}) + Expect(err).To(MatchError(ContainSubstring("LOCALAI_PROXY_TEST_UNSET_KEY"))) + }) + + It("parses realtime_pipeline, strips the trailing slash and keeps the timeout", func() { + p := loadProxy(up, func(o *pb.ModelOptions) { + o.Options = []string{"other:1", "realtime_pipeline:my-pipe"} + o.Proxy.RequestTimeoutSeconds = 7 + }) + cfg := p.cfg.Load() + Expect(cfg.realtimePipeline).To(Equal("my-pipe")) + Expect(cfg.base).To(Equal(up.URL)) + Expect(cfg.timeout).To(Equal(7 * time.Second)) + }) + + DescribeTable("strips an OpenAI-style /v1 suffix, as the failover prober does", + func(suffix string) { + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamUrl = up.URL + suffix }) + Expect(p.cfg.Load().base).To(Equal(up.URL)) + }, + Entry("/v1", "/v1"), + Entry("/v1/", "/v1/"), + Entry("a full endpoint path", "/v1/chat/completions"), + ) + + It("keeps a path prefix in front of /v1", func() { + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamUrl = up.URL + "/localai/v1" }) + Expect(p.cfg.Load().base).To(Equal(up.URL + "/localai")) + }) + + It("falls back to the model name when upstream_model is unset", func() { + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamModel = "" }) + Expect(p.model("")).To(Equal("local-name")) + }) + }) + + Describe("PredictRich", func() { + It("sends messages to /v1/chat/completions with the upstream model and key", func() { + GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-test") + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" }) + up.replyJSON("/v1/chat/completions", map[string]any{ + "choices": []any{map[string]any{"message": map[string]any{"role": "assistant", "content": "hello back"}}}, + "usage": map[string]any{"prompt_tokens": 3, "completion_tokens": 2}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{ + Messages: []*pb.Message{{Role: "user", Content: "hello"}}, + Tokens: 32, + Temperature: 0.5, + TopK: 40, + StopPrompts: []string{""}, + Seed: 9, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(string(reply.GetMessage())).To(Equal("hello back")) + Expect(reply.GetPromptTokens()).To(Equal(int32(3))) + Expect(reply.GetTokens()).To(Equal(int32(2))) + + req := up.last() + Expect(req.Method).To(Equal(http.MethodPost)) + Expect(req.Path).To(Equal("/v1/chat/completions")) + Expect(req.Auth).To(Equal("Bearer sk-test")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + Expect(req.JSON).To(HaveKeyWithValue("max_tokens", BeNumerically("==", 32))) + Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0.5))) + Expect(req.JSON).To(HaveKeyWithValue("top_k", BeNumerically("==", 40))) + Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 9))) + Expect(req.JSON).To(HaveKeyWithValue("stop", ConsistOf(""))) + Expect(req.JSON).NotTo(HaveKey("stream")) + Expect(req.JSON["messages"]).To(ConsistOf(HaveKeyWithValue("content", "hello"))) + }) + + It("returns upstream tool calls as chat deltas", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/chat/completions", map[string]any{ + "choices": []any{map[string]any{"message": map[string]any{ + "role": "assistant", + "tool_calls": []any{map[string]any{ + "id": "call_1", "type": "function", + "function": map[string]any{"name": "get_weather", "arguments": `{"city":"Rome"}`}, + }}, + }}}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{ + Messages: []*pb.Message{{Role: "user", Content: "weather?"}}, + Tools: `[{"type":"function","function":{"name":"get_weather"}}]`, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(reply.GetChatDeltas()).To(HaveLen(1)) + tc := reply.GetChatDeltas()[0].GetToolCalls() + Expect(tc).To(HaveLen(1)) + Expect(tc[0].GetName()).To(Equal("get_weather")) + Expect(tc[0].GetArguments()).To(Equal(`{"city":"Rome"}`)) + Expect(up.last().JSON).To(HaveKey("tools")) + }) + + It("sends a bare prompt to /v1/completions", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/completions", map[string]any{ + "choices": []any{map[string]any{"text": "completed"}}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{Prompt: "once upon"}) + Expect(err).NotTo(HaveOccurred()) + Expect(string(reply.GetMessage())).To(Equal("completed")) + req := up.last() + Expect(req.Path).To(Equal("/v1/completions")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "once upon")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + + It("maps a 5xx upstream to Unavailable with the body in the message", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusServiceUnavailable, Body: "backend is down"}) + + _, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(err.Error()).To(ContainSubstring("backend is down")) + }) + + It("maps a 4xx upstream to InvalidArgument", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad request"}) + + _, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}}) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + Expect(err.Error()).To(ContainSubstring("bad request")) + }) + + It("truncates a long upstream error body", func() { + p := loadProxy(up, nil) + long := make([]byte, 2000) + for i := range long { + long[i] = 'a' + } + up.script("/v1/completions", scriptedResponse{Status: http.StatusInternalServerError, Body: string(long)}) + + _, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700)) + }) + + It("maps a 429 upstream to ResourceExhausted so failover moves to the next target", func() { + p := loadProxy(up, nil) + up.script("/v1/completions", scriptedResponse{Status: http.StatusTooManyRequests, Body: "slow down"}) + + _, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(codeOf(err)).To(Equal(codes.ResourceExhausted)) + Expect(err.Error()).To(ContainSubstring("slow down")) + }) + + It("always sends temperature, even 0, so greedy decoding survives", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "t"}}}) + + _, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(err).NotTo(HaveOccurred()) + req := up.last() + Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0))) + Expect(req.JSON).NotTo(HaveKey("top_p")) + Expect(req.JSON).NotTo(HaveKey("top_k")) + }) + + It("maps an unreachable upstream to Unavailable", func() { + p := loadProxy(up, nil) + up.Close() + + _, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + }) + + It("reports an unloaded proxy as FailedPrecondition", func() { + _, err := NewLocalAIProxy().PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(codeOf(err)).To(Equal(codes.FailedPrecondition)) + }) + }) + + Describe("PredictStreamRich", func() { + It("streams SSE deltas in order and leaves the channel open", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"role": "assistant"}}}}), + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "Hel"}}}}), + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "lo"}}}}), + "[DONE]", + }}) + + results := make(chan *pb.Reply, 10) + err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results) + Expect(err).NotTo(HaveOccurred()) + + var got []string + for len(results) > 0 { + got = append(got, string((<-results).GetMessage())) + } + Expect(got).To(Equal([]string{"Hel", "lo"})) + // The gRPC server closes the channel; closing it here must not panic. + close(results) + Expect(up.last().JSON).To(HaveKeyWithValue("stream", true)) + }) + + It("asks the upstream for the usage trailer and reports its token counts", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "hi"}}}}), + sseJSON(map[string]any{"choices": []any{}, "usage": map[string]any{"prompt_tokens": 7, "completion_tokens": 1}}), + "[DONE]", + }}) + + results := make(chan *pb.Reply, 10) + Expect(p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)).To(Succeed()) + // LocalAI (like OpenAI) only sends the usage trailer on request. + Expect(up.last().JSON).To(HaveKeyWithValue("stream_options", HaveKeyWithValue("include_usage", true))) + Expect(results).To(HaveLen(2)) + <-results + usage := <-results + Expect(usage.GetPromptTokens()).To(Equal(int32(7))) + Expect(usage.GetTokens()).To(Equal(int32(1))) + }) + + It("streams /v1/completions text for a bare prompt", func() { + p := loadProxy(up, nil) + up.script("/v1/completions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"choices": []any{map[string]any{"text": "a"}}}), + sseJSON(map[string]any{"choices": []any{map[string]any{"text": "b"}}}), + "[DONE]", + }}) + + results := make(chan *pb.Reply, 10) + Expect(p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, results)).To(Succeed()) + Expect(results).To(HaveLen(2)) + Expect(string((<-results).GetMessage())).To(Equal("a")) + Expect(string((<-results).GetMessage())).To(Equal("b")) + }) + + It("stops the upstream request when the caller cancels", func() { + upstreamGone := make(chan struct{}) + slow := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: " + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "a"}}}}) + "\n\n")) + w.(http.Flusher).Flush() + // A generation that outlasts the spec unless the proxy hangs up. + select { + case <-r.Context().Done(): + close(upstreamGone) + case <-time.After(8 * time.Second): + } + }) + DeferCleanup(slow.Close) + p := loadProxy(slow, nil) + + addr := "test://localai-proxy-cancel" + grpc.Provide(addr, p) + client := grpc.NewClient(addr, true, nil, false) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + errCh := make(chan error, 1) + first := make(chan struct{}, 1) + go func() { + errCh <- client.PredictStream(ctx, &pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, func(*pb.Reply) { + select { + case first <- struct{}{}: + default: + } + }) + }() + + Eventually(first, 5*time.Second).Should(Receive()) + cancel() + Eventually(upstreamGone, 5*time.Second).Should(BeClosed(), "the upstream generation must stop with the caller") + Eventually(errCh, 5*time.Second).Should(Receive(HaveOccurred())) + }) + + It("returns a mid-stream upstream error frame as Unavailable", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "partial"}}}}), + sseJSON(map[string]any{"error": map[string]any{"message": "backend crashed", "type": "server_error", "code": "server_error"}}), + "[DONE]", + }}) + + results := make(chan *pb.Reply, 10) + err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(err.Error()).To(ContainSubstring("backend crashed")) + Expect(results).To(HaveLen(1)) + }) + + It("maps a failing upstream to a gRPC code", func() { + p := loadProxy(up, nil) + up.script("/v1/completions", scriptedResponse{Status: http.StatusBadGateway, Body: "gateway"}) + + err := p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, make(chan *pb.Reply, 1)) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + }) + }) + + Describe("legacy Predict and PredictStream", func() { + It("wrap the rich variants", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "plain"}}}) + out, err := p.Predict(&pb.PredictOptions{Prompt: "x"}) + Expect(err).NotTo(HaveOccurred()) + Expect(out).To(Equal("plain")) + + up.script("/v1/completions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"choices": []any{map[string]any{"text": "s1"}}}), + "[DONE]", + }}) + results := make(chan string, 10) + Expect(p.PredictStream(&pb.PredictOptions{Prompt: "x"}, results)).To(Succeed()) + var got []string + for s := range results { // PredictStream closes the channel + got = append(got, s) + } + Expect(got).To(Equal([]string{"s1"})) + }) + }) + + It("names the request fields it cannot forward", func() { + Expect(unforwardedFields(&pb.PredictOptions{Prompt: "x"})).To(BeEmpty()) + Expect(unforwardedFields(&pb.PredictOptions{ + Grammar: "root ::= x", Images: []string{"i"}, Audios: []string{"a"}, Videos: []string{"v"}, + })).To(Equal([]string{"grammar", "images", "audios", "videos"})) + }) + + Describe("Embeddings", func() { + It("posts the input to /v1/embeddings and returns the first vector", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/embeddings", map[string]any{ + "data": []any{map[string]any{"embedding": []float32{0.1, 0.2, 0.3}}}, + }) + + vec, err := p.Embeddings(&pb.PredictOptions{Embeddings: "embed me"}) + Expect(err).NotTo(HaveOccurred()) + Expect(vec).To(Equal([]float32{0.1, 0.2, 0.3})) + req := up.last() + Expect(req.Path).To(Equal("/v1/embeddings")) + Expect(req.JSON).To(HaveKeyWithValue("input", "embed me")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + + It("sends token input as a token array, not as an empty string", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/embeddings", map[string]any{ + "data": []any{map[string]any{"embedding": []float32{0.5}}}, + }) + + _, err := p.Embeddings(&pb.PredictOptions{EmbeddingTokens: []int32{1, 2, 3}}) + Expect(err).NotTo(HaveOccurred()) + Expect(up.last().JSON).To(HaveKeyWithValue("input", []any{[]any{1.0, 2.0, 3.0}})) + }) + + It("fails when the upstream returns no vector", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/embeddings", map[string]any{"data": []any{}}) + _, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"}) + Expect(err).To(HaveOccurred()) + }) + }) + + Describe("Rerank", func() { + It("posts to /v1/rerank and maps the results", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/rerank", map[string]any{ + "model": "remote-model", + "usage": map[string]any{"total_tokens": 12, "prompt_tokens": 10}, + "results": []any{ + map[string]any{"index": 1, "document": map[string]any{"text": "b"}, "relevance_score": 0.9}, + map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.1}, + }, + }) + + res, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}, TopN: 2}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.GetUsage().GetTotalTokens()).To(Equal(int32(12))) + Expect(res.GetUsage().GetPromptTokens()).To(Equal(int32(10))) + Expect(res.GetResults()).To(HaveLen(2)) + Expect(res.GetResults()[0].GetIndex()).To(Equal(int32(1))) + Expect(res.GetResults()[0].GetText()).To(Equal("b")) + Expect(res.GetResults()[0].GetRelevanceScore()).To(BeNumerically("~", 0.9, 1e-6)) + + req := up.last() + Expect(req.Path).To(Equal("/v1/rerank")) + Expect(req.JSON).To(HaveKeyWithValue("query", "q")) + Expect(req.JSON).To(HaveKeyWithValue("top_n", BeNumerically("==", 2))) + Expect(req.JSON).To(HaveKeyWithValue("documents", ConsistOf("a", "b"))) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + }) + + It("Rerank omits top_n when it is 0, which means score every document", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/rerank", map[string]any{"results": []any{}}) + + _, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}}) + Expect(err).NotTo(HaveOccurred()) + Expect(up.last().JSON).NotTo(HaveKey("top_n")) + }) + + Describe("TokenizeString and Detokenize", func() { + It("posts the prompt to /v1/tokenize", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/tokenize", map[string]any{"tokens": []int32{5, 6, 7}}) + + res, err := p.TokenizeString(&pb.PredictOptions{Prompt: "abc"}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.GetTokens()).To(Equal([]int32{5, 6, 7})) + Expect(res.GetLength()).To(Equal(int32(3))) + req := up.last() + Expect(req.Path).To(Equal("/v1/tokenize")) + Expect(req.JSON).To(HaveKeyWithValue("content", "abc")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + + It("posts tokens to /v1/detokenize", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/detokenize", map[string]any{"content": "abc"}) + + res, err := p.Detokenize(&pb.DetokenizeRequest{Tokens: []int32{5, 6}}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.GetContent()).To(Equal("abc")) + Expect(up.last().JSON).To(HaveKeyWithValue("tokens", ConsistOf(BeNumerically("==", 5), BeNumerically("==", 6)))) + }) + }) + + Describe("Score", func() { + It("posts to /api/score and maps the candidates", func() { + p := loadProxy(up, nil) + up.replyJSON("/api/score", map[string]any{ + "model": "remote-model", + "candidates": []any{map[string]any{ + "log_prob": -1.5, "length_normalized_log_prob": -0.75, "num_tokens": 2, + "tokens": []any{map[string]any{"token": "yes", "log_prob": -1.5}}, + }}, + }) + + res, err := p.Score(context.Background(), &pb.ScoreRequest{ + Prompt: "p", Candidates: []string{"yes"}, IncludeTokenLogprobs: true, LengthNormalize: true, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(res.GetCandidates()).To(HaveLen(1)) + c := res.GetCandidates()[0] + Expect(c.GetLogProb()).To(Equal(-1.5)) + Expect(c.GetLengthNormalizedLogProb()).To(Equal(-0.75)) + Expect(c.GetNumTokens()).To(Equal(int32(2))) + Expect(c.GetTokens()).To(HaveLen(1)) + Expect(c.GetTokens()[0].GetToken()).To(Equal("yes")) + + req := up.last() + Expect(req.Path).To(Equal("/api/score")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "p")) + Expect(req.JSON).To(HaveKeyWithValue("include_token_logprobs", true)) + Expect(req.JSON).To(HaveKeyWithValue("length_normalize", true)) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + + It("refuses decision-pipeline requests it cannot forward", func() { + p := loadProxy(up, nil) + _, err := p.Score(context.Background(), &pb.ScoreRequest{Prompt: "{}", QuestionType: "systemone"}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + Expect(up.recorded()).To(BeEmpty()) + }) + }) + + Describe("helpers", func() { + It("postMultipart sends fields and the file with auth", func() { + GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-mp") + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" }) + up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "ok"}) + + path := filepath.Join(GinkgoT().TempDir(), "a.wav") + Expect(os.WriteFile(path, []byte("RIFFDATA"), 0o600)).To(Succeed()) + + var out struct { + Text string `json:"text"` + } + err := p.postMultipart(context.Background(), "/v1/audio/transcriptions", + map[string]string{"model": "remote-model", "language": "it"}, "file", path, &out) + Expect(err).NotTo(HaveOccurred()) + Expect(out.Text).To(Equal("ok")) + req := up.last() + Expect(req.Auth).To(Equal("Bearer sk-mp")) + Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "language": "it"})) + Expect(req.Files).To(HaveKeyWithValue("file", "RIFFDATA")) + }) + + It("postMultipart reports a missing local file without calling upstream", func() { + p := loadProxy(up, nil) + err := p.postMultipart(context.Background(), "/v1/audio/transcriptions", nil, "file", "/nonexistent/a.wav", nil) + Expect(err).To(HaveOccurred()) + Expect(up.recorded()).To(BeEmpty()) + }) + + It("applies request_timeout_seconds to non-streaming calls", func() { + slow := make(chan struct{}) + hang := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { <-slow }) + // Cleanups run last-in first-out: release the handler before + // Close, which waits for in-flight requests. + DeferCleanup(hang.Close) + DeferCleanup(func() { close(slow) }) + + p := NewLocalAIProxy() + Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{ + UpstreamUrl: hang.URL, UpstreamModel: "m", RequestTimeoutSeconds: 1, + }})).To(Succeed()) + _, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"}) + Expect(codeOf(err)).To(Equal(codes.DeadlineExceeded)) + }) + + It("postStream returns the open response for a 2xx", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "WAVBYTES"}) + resp, err := p.postStream(context.Background(), "/tts", map[string]any{"input": "hi"}) + Expect(err).NotTo(HaveOccurred()) + defer func() { _ = resp.Body.Close() }() + Expect(resp.Header.Get("Content-Type")).To(Equal("audio/wav")) + }) + }) + + Describe("through the gRPC server", func() { + It("dispatches Rerank and keeps the Unimplemented code end to end", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/rerank", map[string]any{"results": []any{ + map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.5}, + }}) + addr := "test://localai-proxy-grpc" + grpc.Provide(addr, p) + client := grpc.NewClient(addr, true, nil, false) + + res, err := client.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a"}}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.GetResults()).To(HaveLen(1)) + + _, err = client.AudioEncode(context.Background(), &pb.AudioEncodeRequest{}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart")) + }) + }) + + Describe("methods with no upstream counterpart", func() { + It("return Unimplemented with the exact message", func() { + p := loadProxy(up, nil) + _, err := p.AudioEncode(&pb.AudioEncodeRequest{}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart")) + + _, err = p.TokenClassify(context.Background(), &pb.TokenClassifyRequest{}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + _, err = p.ModelMetadata(&pb.ModelOptions{}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + _, err = p.StartFineTune(&pb.FineTuneRequest{}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + }) + + It("close the output channel of streaming stubs so the server does not hang", func() { + p := loadProxy(up, nil) + updates := make(chan *pb.FineTuneProgressUpdate) + Expect(codeOf(p.FineTuneProgress(&pb.FineTuneProgressRequest{}, updates))).To(Equal(codes.Unimplemented)) + Eventually(updates).Should(BeClosed()) + + out := make(chan *pb.AudioToAudioResponse) + Expect(codeOf(p.AudioToAudioStream(make(chan *pb.AudioToAudioRequest), out))).To(Equal(codes.Unimplemented)) + Eventually(out).Should(BeClosed()) + }) + }) +}) + +var _ = Describe("clampInt32", func() { + It("passes in-range values through and saturates the rest", func() { + Expect(clampInt32(0)).To(Equal(int32(0))) + Expect(clampInt32(42)).To(Equal(int32(42))) + Expect(clampInt32(-7)).To(Equal(int32(-7))) + Expect(clampInt32(math.MaxInt32 + 1)).To(Equal(int32(math.MaxInt32))) + Expect(clampInt32(math.MinInt32 - 1)).To(Equal(int32(math.MinInt32))) + }) +}) diff --git a/backend/index.yaml b/backend/index.yaml index 8b2efd417..d399d6498 100644 --- a/backend/index.yaml +++ b/backend/index.yaml @@ -2038,6 +2038,21 @@ capabilities: default: "cpu-cloud-proxy" metal: "metal-cloud-proxy" +- &localai-proxy + name: "localai-proxy" + alias: "localai-proxy" + urls: + - https://github.com/mudler/LocalAI/tree/master/backend/go/localai-proxy + description: | + Serve a model from another LocalAI instance: text, embeddings, rerank, audio, image and video requests are forwarded to its REST API. + tags: + - text-to-text + - proxy + - CPU + license: MIT + capabilities: + default: "cpu-localai-proxy" + metal: "metal-localai-proxy" - &valkey-store name: "valkey-store" urls: @@ -2623,6 +2638,31 @@ uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-cloud-proxy" mirrors: - localai/localai-backends:master-metal-darwin-arm64-cloud-proxy +- !!merge <<: *localai-proxy + name: "cpu-localai-proxy" + uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-localai-proxy" + mirrors: + - localai/localai-backends:latest-cpu-localai-proxy +- !!merge <<: *localai-proxy + name: "cpu-localai-proxy-development" + uri: "quay.io/go-skynet/local-ai-backends:master-cpu-localai-proxy" + mirrors: + - localai/localai-backends:master-cpu-localai-proxy +- !!merge <<: *localai-proxy + name: "localai-proxy-development" + capabilities: + default: "cpu-localai-proxy-development" + metal: "metal-localai-proxy-development" +- !!merge <<: *localai-proxy + name: "metal-localai-proxy" + uri: "quay.io/go-skynet/local-ai-backends:latest-metal-darwin-arm64-localai-proxy" + mirrors: + - localai/localai-backends:latest-metal-darwin-arm64-localai-proxy +- !!merge <<: *localai-proxy + name: "metal-localai-proxy-development" + uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-localai-proxy" + mirrors: + - localai/localai-backends:master-metal-darwin-arm64-localai-proxy - !!merge <<: *valkey-store name: "cpu-valkey-store" alias: "valkey-store" diff --git a/core/application/application.go b/core/application/application.go index b49851d47..56e22eeed 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -13,9 +13,12 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/http/auth" mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" + "github.com/mudler/LocalAI/core/services/advisorylock" "github.com/mudler/LocalAI/core/services/agentpool" "github.com/mudler/LocalAI/core/services/cloudproxy/mitm" "github.com/mudler/LocalAI/core/services/facerecognition" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/core/services/failover/distsync" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/monitoring" "github.com/mudler/LocalAI/core/services/nodes" @@ -82,6 +85,7 @@ type Application struct { routerRegistry *router.Registry routerCorpus *corpus.Manager admissionLimiter *admission.Limiter + failoverManager *failover.Manager watchdogMutex sync.Mutex watchdogStop chan bool p2pMutex sync.Mutex @@ -92,6 +96,13 @@ type Application struct { // Distributed mode services (nil when not in distributed mode) distributed *DistributedServices + // failoverSync shares failover state between frontends; nil in + // standalone mode or when it could not start. + failoverSync *distsync.Sync + // failoverLock is the probe-leader lock this frontend holds or tries + // for; nil when failoverSync is nil. + failoverLock *advisorylock.HeldLock + // Upgrade checker (background service for detecting backend upgrades) upgradeChecker *UpgradeChecker @@ -99,6 +110,13 @@ type Application struct { // is set; otherwise initialised in start() after galleryService. localAIAssistant *mcpTools.LocalAIAssistantHolder + // assistantClient is the concrete inproc client backing localAIAssistant. + // start() constructs it before failoverManager exists (see New() in + // startup.go), so New() sets assistantClient.Failover once the manager + // is built, using this field to reach back into the already-registered + // MCP tool set. nil when DisableLocalAIAssistant is set. + assistantClient *localaiInproc.Client + // startupComplete flips to true once New() has finished its whole startup // sequence. It backs the /readyz probe. // @@ -478,6 +496,9 @@ func (a *Application) AdmissionLimiter() *admission.Limiter { return a.admissionLimiter } +// FailoverManager serves failover chains. Never nil after New. +func (a *Application) FailoverManager() *failover.Manager { return a.failoverManager } + // StartupConfig returns the original startup configuration (from env vars, before file loading) func (a *Application) StartupConfig() *config.ApplicationConfig { return a.startupConfig @@ -503,6 +524,9 @@ func (a *Application) IsDistributed() bool { func (a *Application) Shutdown() error { var err error a.shutdownOnce.Do(func() { + // Before distributed shutdown: the sync's subscriptions live on the + // NATS connection that closes there. + a.stopFailoverDistributed() a.distributed.Shutdown() if a.modelLoader != nil { err = a.modelLoader.StopAllGRPC() @@ -593,6 +617,12 @@ func (a *Application) start() error { assistantClient.RouterEmbedder = a.Embedder assistantClient.RouterEmbedderFingerprint = a.EmbedderFingerprint assistantClient.RouterVectorStore = a.VectorStore + // Failover chains: failoverManager does not exist yet at this point + // in startup (New() in startup.go builds it after start() returns), + // so it can't be wired here like the fields above. New() sets + // assistantClient.Failover directly once the manager is built; + // stash the client so it can reach back into it. + a.assistantClient = assistantClient if err := holder.Initialize(a.applicationConfig.Context, assistantClient, localaitools.Options{}); err != nil { // Why log+continue instead of fail: the assistant is an optional // feature; a failure here must not take down the whole server. diff --git a/core/application/distributed.go b/core/application/distributed.go index bcafcb26e..867743a33 100644 --- a/core/application/distributed.go +++ b/core/application/distributed.go @@ -77,7 +77,9 @@ func (ds *DistributedServices) Shutdown() { // Returns nil if distributed mode is not enabled. // configLoader is used by the SmartRouter to compute concurrency-group // anti-affinity at placement time (#9659); it may be nil in tests. -func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader) (*DistributedServices, error) { +// pinned, when set, replaces configLoader as the source of models the router +// and reconciler must keep loaded (it adds warm failover targets). +func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader, pinned nodes.PinnedModelResolver) (*DistributedServices, error) { if !cfg.Distributed.Enabled { return nil, nil } @@ -387,6 +389,9 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade modelFiles = declaredModelFiles(configLoader, cfg.SystemState.Model.ModelsPath) } } + if pinned != nil { + pinnedResolver = pinned + } modelCleanup := nodes.NewModelCleanupService(registry, remoteUnloader) router := nodes.NewSmartRouter(registry, nodes.SmartRouterOptions{ Unloader: remoteUnloader, diff --git a/core/application/failover.go b/core/application/failover.go new file mode 100644 index 000000000..ffe5a8e51 --- /dev/null +++ b/core/application/failover.go @@ -0,0 +1,59 @@ +package application + +import ( + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/xlog" +) + +// preloadModelByName is a seam over backend.PreloadModelByName so tests can +// substitute a controllable loader instead of touching real models/disk. +var preloadModelByName = backend.PreloadModelByName + +// applyFailoverWarmTargets pins warm failover targets in the watchdog and +// loads them, so a switch does not wait for a cold load. +// +// SyncPinnedModelsToWatchdog runs synchronously: it is a cheap in-memory +// update, and the pin must land before the watchdog can evict a target that +// is about to become (or stay) a chain's active path. Preloading is not +// cheap — it can download or load a multi-GB model — and this callback runs +// on the failover manager's single scheduler goroutine (Sync, called from +// Tick, called from Run), before that tick's probes fire. A slow or hung +// load here would freeze probing and fail-back for every chain, so it runs +// in its own goroutine instead of blocking the scheduler loop. +// +// In distributed mode only the probe leader preloads: every frontend would +// otherwise ask the workers for the same load. Followers still pin, so their +// own watchdog never evicts a warm target they happen to hold. +func (a *Application) applyFailoverWarmTargets(warm []string) { + a.SyncPinnedModelsToWatchdog() + if a.IsDistributed() && !a.failoverManager.IsLeader() { + return + } + go func() { + for _, name := range warm { + if _, err := preloadModelByName(a.ApplicationConfig().Context, a.ModelConfigLoader(), a.ModelLoader(), a.ApplicationConfig(), name); err != nil { + xlog.Warn("failover: could not preload warm target", "model", name, "error", err) + } + } + }() +} + +// failoverLoadedBackend gives the failover prober the running backend of a +// local target without ever loading it. CheckIsLoaded may run the loader's +// own health check and drop a dead process; the target is then "not loaded" +// and the next real request loads and judges it. +func failoverLoadedBackend(ml *model.ModelLoader) failover.LoadedFunc { + return func(cfg config.ModelConfig) grpc.Backend { + m := ml.CheckIsLoaded(cfg.ModelID()) + if m == nil { + return nil + } + // Load always enables parallel requests; match it in case this is + // the first client built for the model. + return m.GRPC(true, ml.GetWatchDog()) + } +} diff --git a/core/application/failover_distributed.go b/core/application/failover_distributed.go new file mode 100644 index 000000000..87adbdeb5 --- /dev/null +++ b/core/application/failover_distributed.go @@ -0,0 +1,94 @@ +package application + +import ( + "context" + + "github.com/mudler/LocalAI/core/services/advisorylock" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/core/services/failover/distsync" + "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/xlog" +) + +// failoverLeaderGate elects the frontend that probes failover targets and +// decides chains, so N frontends do not probe every target N times or +// disagree on which target is active. +// +// Leadership is sticky: the leader keeps the lock across ticks and loses it +// only when its database session dies or it shuts down. A lock taken per tick +// would pass between frontends on almost every tick, and each change of +// leader redelivers the warm set and republishes all state. A lock error +// counts as "not leader": skipping one tick is safe, two leaders are not. +func failoverLeaderGate(l *advisorylock.HeldLock) failover.LeaderGate { + return func(ctx context.Context, fn func()) bool { + if !l.Held() || !l.Verify(ctx) { + ok, err := l.TryAcquire(ctx) + if err != nil { + xlog.Warn("failover: could not take the prober leader lock", "error", err) + return false + } + if !ok { + return false + } + } + fn() + return true + } +} + +// failoverPinnedResolver adds warm failover targets to the config-pinned +// models, so the SmartRouter and ReplicaReconciler keep them loaded on the +// workers the same way the local watchdog keeps them loaded in standalone mode. +type failoverPinnedResolver struct { + base nodes.PinnedModelResolver + fm *failover.Manager +} + +func (r *failoverPinnedResolver) GetPinnedModelNames() []string { + var pinned []string + if r.base != nil { + pinned = r.base.GetPinnedModelNames() + } + return failover.MergePinned(pinned, r.fm.WarmTargets()) +} + +// startFailoverDistributed makes the failover manager cluster-aware: one +// frontend (the advisory-lock holder) probes and decides, and pins, target +// health and chain state are shared over NATS. A sync failure is logged and +// the manager keeps probing on its own, as in standalone mode. +func (a *Application) startFailoverDistributed(ctx context.Context) { + db := a.distributedDB() + pins, err := distsync.NewPinStore(db) + if err != nil { + xlog.Error("failover: pins will not persist, could not prepare the pin store", "error", err) + pins = nil // distsync.New treats a nil store as "no durable pins" + } + s, err := distsync.New(ctx, a.distributed.Nats, pins, a.failoverManager) + if err != nil { + xlog.Error("failover: state will not be shared between frontends", "error", err) + return + } + a.failoverSync = s + // Gate only once state is shared: a follower learns target health and + // chain decisions solely from the leader's publishes, so gating without + // the sync would leave every follower's chains frozen. + a.failoverLock = advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + a.failoverManager.SetLeaderGate(failoverLeaderGate(a.failoverLock)) +} + +// stopFailoverDistributed detaches the manager from the sync before closing +// it, so nothing publishes into closed maps, and gives up leadership at once +// instead of making another frontend wait for this session to time out. +func (a *Application) stopFailoverDistributed() { + if a.failoverSync != nil { + a.failoverManager.SetStateSync(nil) + if err := a.failoverSync.Close(); err != nil { + xlog.Warn("failover: closing state sync", "error", err) + } + } + if a.failoverLock != nil { + // Close, not Release: Run may still be ticking and would take the + // lock straight back. + a.failoverLock.Close() + } +} diff --git a/core/application/failover_distributed_test.go b/core/application/failover_distributed_test.go new file mode 100644 index 000000000..1021d0105 --- /dev/null +++ b/core/application/failover_distributed_test.go @@ -0,0 +1,134 @@ +package application + +import ( + "context" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/advisorylock" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +// staticPins is a nodes.PinnedModelResolver with a fixed list. +type staticPins []string + +func (s staticPins) GetPinnedModelNames() []string { return s } + +// failoverSource is a failover.ConfigSource over a fixed set of configs. +type failoverSource map[string]config.ModelConfig + +func (s failoverSource) GetModelConfig(name string) (config.ModelConfig, bool) { + c, ok := s[name] + return c, ok +} + +func (s failoverSource) GetAllModelsConfigs() []config.ModelConfig { + out := make([]config.ModelConfig, 0, len(s)) + for _, c := range s { + out = append(out, c) + } + return out +} + +// warmChainSource has one chain whose two local targets are warm. +func warmChainSource() failoverSource { + return failoverSource{ + "b": {Name: "b", Backend: "llama-cpp"}, + "c": {Name: "c", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{ + {Model: "b", Warm: true}, {Model: "c", Warm: true}, + }}}, + } +} + +var _ = Describe("failoverPinnedResolver", func() { + It("merges config pins and warm failover targets without duplicates", func() { + fm := failover.New(warmChainSource()) + fm.Sync() + Expect(fm.WarmTargets()).To(Equal([]string{"b", "c"})) + + r := &failoverPinnedResolver{base: staticPins{"pinned-a", "b"}, fm: fm} + Expect(r.GetPinnedModelNames()).To(Equal([]string{"pinned-a", "b", "c"})) + }) +}) + +var _ = Describe("failoverLeaderGate", func() { + It("keeps leadership with the first frontend until it releases", func() { + // Not PostgreSQL, so advisorylock falls back to its in-process lock, + // which has the same try-lock semantics as pg_try_advisory_lock. + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + firstLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + secondLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + DeferCleanup(firstLock.Release) + DeferCleanup(secondLock.Release) + first, second := failoverLeaderGate(firstLock), failoverLeaderGate(secondLock) + ctx := context.Background() + + var firstRuns, secondRuns int + for range 5 { + Expect(first(ctx, func() { firstRuns++ })).To(BeTrue()) + Expect(second(ctx, func() { secondRuns++ })).To(BeFalse(), "leadership must not flip between ticks") + } + Expect(firstRuns).To(Equal(5)) + Expect(secondRuns).To(BeZero()) + + firstLock.Release() + Expect(second(ctx, func() { secondRuns++ })).To(BeTrue(), "a released leadership passes to the next frontend") + Expect(secondRuns).To(Equal(1)) + Expect(first(ctx, func() { firstRuns++ })).To(BeFalse()) + }) +}) + +var _ = Describe("applyFailoverWarmTargets in distributed mode", func() { + It("pins but does not preload on a frontend that is not the probe leader", func() { + preloaded := make(chan string, 1) + orig := preloadModelByName + preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) { + preloaded <- name + return nil, nil + } + DeferCleanup(func() { preloadModelByName = orig }) + + // A gate that never grants leadership: this frontend is a follower. + fm := failover.New(warmChainSource(), + failover.WithLeaderGate(func(context.Context, func()) bool { return false })) + app := &Application{ + applicationConfig: &config.ApplicationConfig{Context: context.Background()}, + distributed: &DistributedServices{}, + failoverManager: fm, + } + + app.applyFailoverWarmTargets([]string{"b"}) + Consistently(preloaded, 200*time.Millisecond).ShouldNot(Receive(), "only the probe leader preloads warm targets") + }) + + It("preloads warm targets on the probe leader", func() { + preloaded := make(chan string, 4) + orig := preloadModelByName + preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) { + preloaded <- name + return nil, nil + } + DeferCleanup(func() { preloadModelByName = orig }) + + app := &Application{ + applicationConfig: &config.ApplicationConfig{Context: context.Background()}, + distributed: &DistributedServices{}, + } + // A gate that always grants: this frontend is the leader. The warm + // set is delivered on the first tick that wins it. + app.failoverManager = failover.New(warmChainSource(), + failover.WithLeaderGate(func(_ context.Context, fn func()) bool { fn(); return true }), + failover.WithOnWarmChanged(app.applyFailoverWarmTargets)) + + app.failoverManager.Tick(context.Background()) + Eventually(preloaded, time.Second).Should(Receive(Equal("b"))) + Eventually(preloaded, time.Second).Should(Receive(Equal("c"))) + }) +}) diff --git a/core/application/failover_test.go b/core/application/failover_test.go new file mode 100644 index 000000000..14eaef241 --- /dev/null +++ b/core/application/failover_test.go @@ -0,0 +1,63 @@ +package application + +import ( + "context" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("applyFailoverWarmTargets", func() { + It("returns promptly even when the preload call blocks", func() { + // Guards against a regression to a synchronous preload loop: onWarm + // runs on the failover manager's single scheduler goroutine, so a + // blocking loader here must not block the caller. + started := make(chan struct{}) + release := make(chan struct{}) + orig := preloadModelByName + preloadModelByName = func(ctx context.Context, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, name string) ([]string, error) { + close(started) + <-release // never released within the test's timeout + return nil, nil + } + DeferCleanup(func() { preloadModelByName = orig }) + DeferCleanup(func() { close(release) }) + + app := &Application{applicationConfig: &config.ApplicationConfig{Context: context.Background()}} + + callReturned := make(chan struct{}) + go func() { + defer GinkgoRecover() + app.applyFailoverWarmTargets([]string{"warm-a"}) + close(callReturned) + }() + + Eventually(callReturned, time.Second).Should(BeClosed(), "applyFailoverWarmTargets must not wait on the preload goroutine") + Eventually(started, time.Second).Should(BeClosed(), "the preload goroutine should still run in the background") + }) +}) + +type healthyBackend struct{ grpc.Backend } + +func (healthyBackend) HealthCheck(context.Context) (bool, error) { return true, nil } + +var _ = Describe("failoverLoadedBackend", func() { + It("returns the running backend and never loads a model that is not loaded", func() { + ml := model.NewModelLoader(&system.SystemState{Model: system.Model{ModelsPath: GinkgoT().TempDir()}}) + store := model.NewInMemoryModelStore() + ml.SetModelStore(store) + loaded := failoverLoadedBackend(ml) + + Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(BeNil()) + Expect(ml.ListLoadedModels()).To(BeEmpty(), "the lookup must not start a load") + + client := healthyBackend{} + store.Set("gemma", model.NewModelWithClient("gemma", "127.0.0.1:0", client)) + Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(Equal(client)) + }) +}) diff --git a/core/application/startup.go b/core/application/startup.go index abc2f4a17..fe8ee1394 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -12,6 +12,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/jobs" "github.com/mudler/LocalAI/core/services/messaging" @@ -250,6 +251,20 @@ func New(opts ...config.AppOption) (*Application, error) { // the embedding-cache stats endpoint sees a single source of truth. application.routerRegistry = router.NewRegistry() + // Failover chains: probe targets and track which one is active per + // chain. WithOnWarmChanged pins and preloads warm local targets so a + // switch to them does not wait for a cold load. + application.failoverManager = failover.New(application.ModelConfigLoader(), + failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()), options.ProxyAPIKeyEnvLookup)), + failover.WithOnWarmChanged(application.applyFailoverWarmTargets), + ) + // The assistant client was built in start() (above), before this + // manager existed; wire it now so list_failover_chains / + // pin_failover_target / unpin_failover_target see real chains. + if application.assistantClient != nil { + application.assistantClient.Failover = application.failoverManager + } + // Subsystem 5: admission control. Limiter is always wired so a // model that gains a limits: block via gallery install or YAML // edit takes effect on the next restart without conditional plumbing. @@ -271,12 +286,16 @@ func New(opts ...config.AppOption) (*Application, error) { // the model configs are loaded, so it is declared out here. var revisionStore modeladmin.RevisionStore - distSvc, err := initDistributed(options, application.authDB, application.ModelConfigLoader()) + distSvc, err := initDistributed(options, application.authDB, application.ModelConfigLoader(), + &failoverPinnedResolver{base: application.ModelConfigLoader(), fm: application.failoverManager}) if err != nil { return nil, fmt.Errorf("distributed mode initialization failed: %w", err) } if distSvc != nil { application.distributed = distSvc + // Before failoverManager.Run starts below: the gate and sync must be + // in place for its first tick. + application.startFailoverDistributed(options.Context) // Wire remote model unloader so ShutdownModel works for remote nodes // Uses NATS to tell serve-backend nodes to Free + kill their backend process application.modelLoader.SetRemoteUnloader(distSvc.Unloader) @@ -548,6 +567,12 @@ func New(opts ...config.AppOption) (*Application, error) { } } + // Start the failover scheduler: it syncs chains from config, runs + // liveness/recovery probes and dwell-based fail-back. Run is the only + // caller of Sync in production so onWarm callbacks stay ordered. + failover.RegisterMetrics(application.failoverManager) + go application.failoverManager.Run(options.Context) + // Watch the configuration directory startWatcher(options) diff --git a/core/application/watchdog.go b/core/application/watchdog.go index 330c95353..4674cb470 100644 --- a/core/application/watchdog.go +++ b/core/application/watchdog.go @@ -2,6 +2,7 @@ package application import ( "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/xlog" ) @@ -23,6 +24,9 @@ func (a *Application) SyncPinnedModelsToWatchdog() { pinned = append(pinned, cfg.Name) } } + if a.failoverManager != nil { + pinned = failover.MergePinned(pinned, a.failoverManager.WarmTargets()) + } wd.SetPinnedModels(pinned) xlog.Debug("Synced pinned models to watchdog", "count", len(pinned)) } diff --git a/core/backend/options.go b/core/backend/options.go index 93700f93d..b2b652ac8 100644 --- a/core/backend/options.go +++ b/core/backend/options.go @@ -542,7 +542,7 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions { Tokenizer: c.Tokenizer, } - if c.Backend == "cloud-proxy" { + if c.IsRemoteProxy() { opts.Proxy = &pb.ProxyOptions{ UpstreamUrl: c.Proxy.UpstreamURL, Mode: c.Proxy.Mode, @@ -553,6 +553,14 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions { RequestTimeoutSeconds: int32(c.Proxy.RequestTimeoutSeconds), CachePrompt: c.Proxy.CachePrompt, } + // localai-proxy calls a LocalAI that knows the model by name, so an + // unset upstream_model means this config's name, the same derivation + // failover.UpstreamModel uses. Not for cloud-proxy: its translate mode + // falls back to parameters.model and passthrough keeps the client's + // model when upstream_model is empty. + if c.Backend == model.LocalAIProxyBackend && opts.Proxy.UpstreamModel == "" { + opts.Proxy.UpstreamModel = c.Name + } } if c.MMProj != "" { diff --git a/core/backend/options_internal_test.go b/core/backend/options_internal_test.go index 18e081fb1..13ad13b57 100644 --- a/core/backend/options_internal_test.go +++ b/core/backend/options_internal_test.go @@ -44,6 +44,58 @@ var _ = Describe("grpcModelOpts EngineArgs", func() { }) }) +var _ = Describe("grpcModelOpts Proxy options", func() { + It("builds Proxy options for localai-proxy, the same as cloud-proxy", func() { + threads := 1 + cfg := config.ModelConfig{ + Threads: &threads, + Backend: "localai-proxy", + Proxy: config.ProxyConfig{ + UpstreamURL: "http://127.0.0.1:8081", + Mode: config.ProxyModePassthrough, + }, + } + + opts := grpcModelOpts(cfg, "/tmp/models") + Expect(opts.Proxy).NotTo(BeNil()) + Expect(opts.Proxy.UpstreamUrl).To(Equal("http://127.0.0.1:8081")) + Expect(opts.Proxy.Mode).To(Equal(config.ProxyModePassthrough)) + }) + + It("sends the config name as the localai-proxy upstream model when none is set", func() { + threads := 1 + cfg := config.ModelConfig{ + Name: "argus-whisper", + Threads: &threads, + Backend: "localai-proxy", + Proxy: config.ProxyConfig{UpstreamURL: "http://127.0.0.1:8081"}, + } + cfg.Model = "some-file" + + Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("argus-whisper")) + + cfg.Proxy.UpstreamModel = "whisper-large" + Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("whisper-large")) + }) + + It("leaves the cloud-proxy upstream model unset so translate mode keeps its fallback", func() { + threads := 1 + cfg := config.ModelConfig{ + Name: "claude-strict", + Threads: &threads, + Backend: "cloud-proxy", + Proxy: config.ProxyConfig{UpstreamURL: "https://api.example.com", Mode: config.ProxyModeTranslate}, + } + Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(BeEmpty()) + }) + + It("leaves Proxy nil for a backend that is not a proxy", func() { + threads := 1 + opts := grpcModelOpts(config.ModelConfig{Threads: &threads, Backend: "llama-cpp"}, "/tmp/models") + Expect(opts.Proxy).To(BeNil()) + }) +}) + var _ = Describe("grpcModelOpts Diffusers options", func() { It("forwards original_config_file without rewriting it", func() { threads := 1 diff --git a/core/backend/preload.go b/core/backend/preload.go index 3525da299..ab725bff3 100644 --- a/core/backend/preload.go +++ b/core/backend/preload.go @@ -66,12 +66,12 @@ func pipelineStages(cl *config.ModelConfigLoader, p *config.Pipeline, modelPath } var stages []PreloadStage for _, s := range []struct{ role, name string }{ - {"vad", p.VAD}, - {"transcription", p.Transcription}, - {"llm", p.LLM}, - {"tts", p.TTS}, - {"sound_detection", p.SoundDetection}, - {"voice_recognition", voiceRec}, + {config.PipelineStageVAD, p.VAD}, + {config.PipelineStageTranscription, p.Transcription}, + {config.PipelineStageLLM, p.LLM}, + {config.PipelineStageTTS, p.TTS}, + {config.PipelineStageSoundDetection, p.SoundDetection}, + {config.PipelineStageVoiceRecognition, voiceRec}, } { if s.name == "" { continue diff --git a/core/cli/run.go b/core/cli/run.go index 3dc165ecc..de6621ba9 100644 --- a/core/cli/run.go +++ b/core/cli/run.go @@ -291,6 +291,7 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error { } opts := []config.AppOption{ + config.WithProxyAPIKeyEnvLookup(os.Getenv), config.WithContext(context.Background()), config.WithArtifactDownloadConcurrency(r.ArtifactDownloadConcurrency), config.WithModelArtifactMaterializer(modelartifacts.NewDefaultManager( diff --git a/core/config/application_config.go b/core/config/application_config.go index b83e8fe58..24e1867c4 100644 --- a/core/config/application_config.go +++ b/core/config/application_config.go @@ -14,6 +14,9 @@ import ( ) type ApplicationConfig struct { + // ProxyAPIKeyEnvLookup resolves upstream credentials at the CLI boundary. + ProxyAPIKeyEnvLookup func(string) string `json:"-" yaml:"-"` + Context context.Context ConfigFile string SystemState *system.SystemState @@ -271,6 +274,10 @@ type AgentPoolConfig struct { type AppOption func(*ApplicationConfig) +func WithProxyAPIKeyEnvLookup(lookup func(string) string) AppOption { + return func(o *ApplicationConfig) { o.ProxyAPIKeyEnvLookup = lookup } +} + func NewApplicationConfig(o ...AppOption) *ApplicationConfig { opt := &ApplicationConfig{ Context: context.Background(), diff --git a/core/config/hooks_localai_proxy.go b/core/config/hooks_localai_proxy.go new file mode 100644 index 000000000..5104fa6bc --- /dev/null +++ b/core/config/hooks_localai_proxy.go @@ -0,0 +1,30 @@ +package config + +import "github.com/mudler/LocalAI/pkg/model" + +func init() { + RegisterBackendHook(model.LocalAIProxyBackend, localAIProxyDefaults) +} + +// localAIProxyDefaults makes chat requests reach the upstream as structured +// messages. Without the tokenizer template core renders the prompt itself +// with no template, and the proxy can only send that text to +// /v1/completions, bypassing the upstream model's chat template, tool +// handling and reasoning parsing. +// +// Only configs that declare the chat usecase get it: usecase guessing reads +// the tokenizer template as "this model chats", so setting it on a +// transcription or embedding proxy would offer that model to chat pickers +// and default-model selection. A config that brings its own templates keeps +// them: the operator chose local templating. +func localAIProxyDefaults(cfg *ModelConfig, _ string) { + t := cfg.TemplateConfig + if t.UseTokenizerTemplate || t.Chat != "" || t.ChatMessage != "" || t.Completion != "" || t.Edit != "" { + return + } + declared := GetUsecasesFromYAML(cfg.KnownUsecaseStrings) + if declared == nil || *declared&FLAG_CHAT != FLAG_CHAT { + return + } + cfg.TemplateConfig.UseTokenizerTemplate = true +} diff --git a/core/config/hooks_test.go b/core/config/hooks_test.go index b69bc6989..60cc8fd82 100644 --- a/core/config/hooks_test.go +++ b/core/config/hooks_test.go @@ -127,6 +127,38 @@ var _ = Describe("Backend hooks and parser defaults", func() { }) }) + Context("localai-proxy hook", func() { + It("defaults a chat proxy to the tokenizer template so chat sends messages upstream", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeTrue()) + }) + + It("leaves non-chat and undeclared proxies alone so they are not guessed as chat", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"transcript"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + Expect(cfg.HasUsecases(FLAG_CHAT)).To(BeFalse()) + + bare := &ModelConfig{Backend: "localai-proxy"} + bare.SetDefaults() + Expect(bare.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + + It("keeps a config that brings its own templates", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.TemplateConfig.Chat = "{{.Input}}" + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + + It("does not touch other backends", func() { + cfg := &ModelConfig{Backend: "cloud-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + }) + Context("vllmDefaults hook", func() { It("auto-sets parsers for known model families on vllm backend", func() { cfg := &ModelConfig{ diff --git a/core/config/meta/registry.go b/core/config/meta/registry.go index 3e53a81fd..c98514a78 100644 --- a/core/config/meta/registry.go +++ b/core/config/meta/registry.go @@ -419,6 +419,39 @@ func DefaultRegistry() map[string]FieldMetaOverride { Order: 0, }, + // --- Failover --- + "failover.targets": { + Section: "failover", + Label: "Failover targets", + Description: "Ordered list of models that serve this chain. The first healthy target serves each request; later targets take over when it fails. Mark a local target warm to keep it loaded.", + Component: "failover-targets", + Order: 0, + }, + "failover.probe.interval": { + Section: "failover", Label: "Probe interval", Component: "input", Order: 1, Advanced: true, + Description: "How often an idle target is checked, as a duration (default 15s).", Placeholder: "15s", + }, + "failover.probe.timeout": { + Section: "failover", Label: "Probe timeout", Component: "input", Order: 2, Advanced: true, + Description: "How long one probe may take (default 5s).", Placeholder: "5s", + }, + "failover.trip.errors": { + Section: "failover", Label: "Errors to trip", Component: "number", Order: 3, Advanced: true, + Description: "Failures within the trip window that mark a target down (default 1).", + }, + "failover.trip.window": { + Section: "failover", Label: "Trip window", Component: "input", Order: 4, Advanced: true, + Description: "Window in which failures are counted (default 30s).", Placeholder: "30s", + }, + "failover.recovery.probes": { + Section: "failover", Label: "Recovery probes", Component: "number", Order: 5, Advanced: true, + Description: "Consecutive real test requests a target must pass before it is used again (default 3).", + }, + "failover.recovery.min_dwell": { + Section: "failover", Label: "Minimum time on fallback", Component: "input", Order: 6, Advanced: true, + Description: "Minimum time on a lower target before traffic moves back to a recovered higher one (default 60s).", Placeholder: "60s", + }, + // --- Pipeline --- "pipeline.llm": { Section: "pipeline", diff --git a/core/config/meta/registry_test.go b/core/config/meta/registry_test.go index c3f1a494d..72b2380e5 100644 --- a/core/config/meta/registry_test.go +++ b/core/config/meta/registry_test.go @@ -28,6 +28,18 @@ var _ = Describe("alias field metadata", func() { } Expect(found).To(BeTrue(), "DefaultSections should include an alias section") }) + + It("registers the failover section", func() { + reg := meta.DefaultRegistry() + Expect(reg).To(HaveKey("failover.targets")) + Expect(reg["failover.targets"].Section).To(Equal("failover")) + Expect(reg["failover.targets"].Component).To(Equal("failover-targets")) + var ids []string + for _, s := range meta.DefaultSections() { + ids = append(ids, s.ID) + } + Expect(ids).To(ContainElement("failover")) + }) }) var _ = Describe("MCP field metadata", func() { diff --git a/core/config/meta/types.go b/core/config/meta/types.go index a29e66967..b7774d157 100644 --- a/core/config/meta/types.go +++ b/core/config/meta/types.go @@ -70,6 +70,7 @@ func DefaultSections() []Section { return []Section{ {ID: "general", Label: "General", Icon: "settings", Order: 0}, {ID: "alias", Label: "Alias", Icon: "git-merge", Order: 5}, + {ID: "failover", Label: "Failover", Icon: "git-merge", Order: 6}, {ID: "llm", Label: "LLM", Icon: "cpu", Order: 10}, {ID: "parameters", Label: "Parameters", Icon: "sliders", Order: 20}, {ID: "templates", Label: "Templates", Icon: "file-text", Order: 30}, diff --git a/core/config/model_config.go b/core/config/model_config.go index 9754ef6ac..3c8987920 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -17,6 +17,7 @@ import ( "github.com/mudler/LocalAI/core/services/routing/piipattern" "github.com/mudler/LocalAI/pkg/downloader" "github.com/mudler/LocalAI/pkg/functions" + "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/modelartifacts" "github.com/mudler/LocalAI/pkg/reasoning" "github.com/mudler/cogito" @@ -75,6 +76,10 @@ type ModelConfig struct { // at create/swap time). See docs/content for Model Aliases. Alias string `yaml:"alias,omitempty" json:"alias,omitempty"` + // Failover makes this config a failover chain over other models. Like an + // alias it has no backend of its own. + Failover *FailoverConfig `yaml:"failover,omitempty" json:"failover,omitempty"` + F16 *bool `yaml:"f16,omitempty" json:"f16,omitempty"` Threads *int `yaml:"threads,omitempty" json:"threads,omitempty"` Debug *bool `yaml:"debug,omitempty" json:"debug,omitempty"` @@ -298,12 +303,37 @@ const ( ProxyProviderAnthropic = "anthropic" ) +// ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or +// "" when neither is set. Mirrored (not imported, to keep backends independent +// of core's package layout) by resolveAPIKey in backend/go/cloud-proxy/proxy.go +// — keep the two in sync, empty-value handling included. +func (p ProxyConfig) ResolveAPIKey(envLookup func(string) string) (string, error) { + switch { + case p.APIKeyEnv != "": + var v string + if envLookup != nil { + v = envLookup(p.APIKeyEnv) + } + if v == "" { + return "", fmt.Errorf("proxy api_key_env %q is unset", p.APIKeyEnv) + } + return v, nil + case p.APIKeyFile != "": + b, err := os.ReadFile(p.APIKeyFile) + if err != nil { + return "", fmt.Errorf("proxy api_key_file %q: %w", p.APIKeyFile, err) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} + // IsCloudProxyBackendPassthrough reports whether this model uses the // cloud-proxy gRPC backend in passthrough mode. Empty Mode counts as // passthrough (SetDefaults normalises it, but Validate accepts empty // too — handlers should not rely on a particular call order). func (c *ModelConfig) IsCloudProxyBackendPassthrough() bool { - if c.Backend != "cloud-proxy" { + if c.Backend != model.CloudProxyBackend { return false } return c.Proxy.Mode == "" || c.Proxy.Mode == ProxyModePassthrough @@ -641,7 +671,7 @@ func (c *ModelConfig) PIIIsEnabled() bool { if c.PII.Enabled != nil { return *c.PII.Enabled } - return c.Backend == "cloud-proxy" + return c.Backend == model.CloudProxyBackend } // PIIDetectors returns the names of the token-classification models that @@ -672,7 +702,7 @@ var piiCoverableUsecases = []ModelConfigUsecase{FLAG_CHAT, FLAG_COMPLETION, FLAG // false naturally: HasUsecases short-circuits to false for any usecase a // declared score/token_classify model did not itself declare. func (c *ModelConfig) PIIFilterApplies() bool { - if c.Backend == "cloud-proxy" { + if c.Backend == model.CloudProxyBackend { return true } return slices.ContainsFunc(piiCoverableUsecases, c.HasUsecases) @@ -770,6 +800,18 @@ type MCPSTDIOServer struct { Command string `json:"command,omitempty"` } +// Pipeline stage names. They match the Pipeline yaml keys and are the stage +// identifiers the realtime endpoint routes by (failover chains per stage, +// model_failover events, preload roles), so every user shares one spelling. +const ( + PipelineStageVAD = "vad" + PipelineStageTranscription = "transcription" + PipelineStageLLM = "llm" + PipelineStageTTS = "tts" + PipelineStageSoundDetection = "sound_detection" + PipelineStageVoiceRecognition = "voice_recognition" +) + // @Description Pipeline defines other models to use for audio-to-audio type Pipeline struct { TTS string `yaml:"tts,omitempty" json:"tts,omitempty"` @@ -1644,6 +1686,13 @@ func (c *ModelConfig) Validate() (bool, error) { return false, fmt.Errorf("a config with artifacts must declare exactly one %q target, found %d", modelartifacts.TargetModel, primaries) } + if c.IsFailover() { + if err := c.validateFailover(); err != nil { + return false, err + } + return true, nil + } + // An alias is a pure redirect: validate only its own shape here. Target // existence and the no-chain rule need the full config set, so the loader // (load-time) and the create/swap endpoints enforce those. diff --git a/core/config/model_config_failover.go b/core/config/model_config_failover.go new file mode 100644 index 000000000..b4829d83e --- /dev/null +++ b/core/config/model_config_failover.go @@ -0,0 +1,163 @@ +package config + +import ( + "fmt" + "time" + + "github.com/mudler/LocalAI/pkg/model" +) + +// FailoverConfig turns a model config into a failover chain: requests for the +// chain name are served by its highest-priority healthy target. See +// core/services/failover for the runtime side. +type FailoverConfig struct { + Targets []FailoverTarget `yaml:"targets" json:"targets"` + Probe FailoverProbe `yaml:"probe,omitempty" json:"probe,omitempty"` + Trip FailoverTrip `yaml:"trip,omitempty" json:"trip,omitempty"` + Recovery FailoverRecovery `yaml:"recovery,omitempty" json:"recovery,omitempty"` +} + +type FailoverTarget struct { + Model string `yaml:"model" json:"model"` + // Warm keeps a local target loaded and exempt from eviction, so a switch + // does not wait for a cold load. + Warm bool `yaml:"warm,omitempty" json:"warm,omitempty"` +} + +type FailoverProbe struct { + Interval string `yaml:"interval,omitempty" json:"interval,omitempty"` + Timeout string `yaml:"timeout,omitempty" json:"timeout,omitempty"` +} + +type FailoverTrip struct { + Errors int `yaml:"errors,omitempty" json:"errors,omitempty"` + Window string `yaml:"window,omitempty" json:"window,omitempty"` +} + +type FailoverRecovery struct { + Probes int `yaml:"probes,omitempty" json:"probes,omitempty"` + MinDwell string `yaml:"min_dwell,omitempty" json:"min_dwell,omitempty"` +} + +const ( + DefaultFailoverProbeInterval = 15 * time.Second + DefaultFailoverProbeTimeout = 5 * time.Second + DefaultFailoverTripErrors = 1 + DefaultFailoverTripWindow = 30 * time.Second + DefaultFailoverRecoveryProbes = 3 + DefaultFailoverMinDwell = 60 * time.Second +) + +// IsFailover reports whether this config is a failover chain. +func (c ModelConfig) IsFailover() bool { return c.Failover != nil } + +// IsRemoteProxy reports whether the model is served by a remote upstream +// through a proxy backend. Failover probes such a target over HTTP, and +// `warm` has no effect on it. +func (c ModelConfig) IsRemoteProxy() bool { + switch c.Backend { + case model.CloudProxyBackend, model.LocalAIProxyBackend: + return true + } + return false +} + +func (f FailoverConfig) ProbeInterval() time.Duration { + return durationOr(f.Probe.Interval, DefaultFailoverProbeInterval) +} +func (f FailoverConfig) ProbeTimeout() time.Duration { + return durationOr(f.Probe.Timeout, DefaultFailoverProbeTimeout) +} +func (f FailoverConfig) TripWindow() time.Duration { + return durationOr(f.Trip.Window, DefaultFailoverTripWindow) +} +func (f FailoverConfig) MinDwell() time.Duration { + return durationOr(f.Recovery.MinDwell, DefaultFailoverMinDwell) +} +func (f FailoverConfig) TripErrors() int { + if f.Trip.Errors <= 0 { + return DefaultFailoverTripErrors + } + return f.Trip.Errors +} +func (f FailoverConfig) RecoveryProbes() int { + if f.Recovery.Probes <= 0 { + return DefaultFailoverRecoveryProbes + } + return f.Recovery.Probes +} + +// WarmFailoverTargets returns the targets marked warm, in chain order. +func (c ModelConfig) WarmFailoverTargets() []string { + if c.Failover == nil { + return nil + } + var out []string + for _, t := range c.Failover.Targets { + if t.Warm { + out = append(out, t.Model) + } + } + return out +} + +func durationOr(s string, def time.Duration) time.Duration { + if s == "" { + return def + } + d, err := time.ParseDuration(s) + if err != nil || d <= 0 { + return def + } + return d +} + +// validateFailover checks what a chain can check without other configs. +// Target existence is checked by ModelConfigLoader.ValidateFailoverTargets. +func (c *ModelConfig) validateFailover() error { + if c.Name == "" { + return fmt.Errorf("failover config requires a name") + } + if c.IsAlias() { + return fmt.Errorf("model %q cannot set both alias and failover", c.Name) + } + if c.Backend != "" || c.Model != "" { + return fmt.Errorf("failover config %q must not set backend or parameters.model: a chain is a pure redirect", c.Name) + } + f := c.Failover + if len(f.Targets) < 2 { + return fmt.Errorf("failover chain %q needs at least 2 targets", c.Name) + } + seen := map[string]bool{} + for _, t := range f.Targets { + switch { + case t.Model == "": + return fmt.Errorf("failover chain %q has a target with no model", c.Name) + case t.Model == c.Name: + return fmt.Errorf("failover chain %q cannot list itself", c.Name) + case seen[t.Model]: + return fmt.Errorf("failover chain %q lists %q twice", c.Name, t.Model) + } + seen[t.Model] = true + } + for key, v := range map[string]string{ + "probe.interval": f.Probe.Interval, + "probe.timeout": f.Probe.Timeout, + "trip.window": f.Trip.Window, + "recovery.min_dwell": f.Recovery.MinDwell, + } { + if v == "" { + continue + } + if d, err := time.ParseDuration(v); err != nil || d <= 0 { + return fmt.Errorf("failover chain %q: invalid %s %q", c.Name, key, v) + } + } + if f.Trip.Errors < 0 { + return fmt.Errorf("failover chain %q: trip.errors must not be negative", c.Name) + } + if f.Recovery.Probes < 0 { + return fmt.Errorf("failover chain %q: recovery.probes must not be negative", c.Name) + } + return nil +} diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go new file mode 100644 index 000000000..6a276f056 --- /dev/null +++ b/core/config/model_config_failover_test.go @@ -0,0 +1,101 @@ +package config + +import ( + "os" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gopkg.in/yaml.v3" +) + +var _ = Describe("ModelConfig failover", func() { + chain := func(targets ...string) ModelConfig { + c := ModelConfig{Name: "chain", Failover: &FailoverConfig{}} + for _, t := range targets { + c.Failover.Targets = append(c.Failover.Targets, FailoverTarget{Model: t}) + } + return c + } + + It("parses the YAML block and applies defaults", func() { + var c ModelConfig + Expect(yaml.Unmarshal([]byte(` +name: assistant-llm +failover: + targets: + - model: argus-llm + - model: gemma-local + warm: true + recovery: + probes: 5 +`), &c)).To(Succeed()) + Expect(c.IsFailover()).To(BeTrue()) + Expect(c.Failover.Targets).To(Equal([]FailoverTarget{{Model: "argus-llm"}, {Model: "gemma-local", Warm: true}})) + Expect(c.Failover.ProbeInterval()).To(Equal(15 * time.Second)) + Expect(c.Failover.ProbeTimeout()).To(Equal(5 * time.Second)) + Expect(c.Failover.TripErrors()).To(Equal(1)) + Expect(c.Failover.TripWindow()).To(Equal(30 * time.Second)) + Expect(c.Failover.RecoveryProbes()).To(Equal(5)) + Expect(c.Failover.MinDwell()).To(Equal(60 * time.Second)) + Expect(c.WarmFailoverTargets()).To(Equal([]string{"gemma-local"})) + }) + + It("accepts a valid chain", func() { + c := chain("a", "b") + ok, err := c.Validate() + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + }) + + DescribeTable("rejects invalid chains", + func(mutate func(*ModelConfig), want string) { + c := chain("a", "b") + mutate(&c) + ok, err := c.Validate() + Expect(ok).To(BeFalse()) + Expect(err).To(MatchError(ContainSubstring(want))) + }, + Entry("alias and failover", func(c *ModelConfig) { c.Alias = "x" }, "both alias and failover"), + Entry("backend set", func(c *ModelConfig) { c.Backend = "llama-cpp" }, "must not set backend"), + Entry("one target", func(c *ModelConfig) { c.Failover.Targets = c.Failover.Targets[:1] }, "at least 2 targets"), + Entry("empty target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "" }, "no model"), + Entry("self target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "chain" }, "cannot list itself"), + Entry("duplicate target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "a" }, "twice"), + Entry("bad duration", func(c *ModelConfig) { c.Failover.Probe.Interval = "soon" }, "invalid probe.interval"), + Entry("negative errors", func(c *ModelConfig) { c.Failover.Trip.Errors = -1 }, "trip.errors"), + Entry("no name", func(c *ModelConfig) { c.Name = "" }, "requires a name"), + ) +}) + +var _ = Describe("ProxyConfig.ResolveAPIKey", func() { + It("reads the key through the application lookup", func() { + app := NewApplicationConfig(WithProxyAPIKeyEnvLookup(func(name string) string { + Expect(name).To(Equal("FAILOVER_TEST_KEY")) + return "k1" + })) + Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(app.ProxyAPIKeyEnvLookup)).To(Equal("k1")) + }) + It("fails without an environment lookup", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(nil) + Expect(err).To(HaveOccurred()) + }) + It("fails on an unset env var", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey(os.Getenv) + Expect(err).To(HaveOccurred()) + }) + It("fails on a set-but-empty env var", func() { + GinkgoT().Setenv("FAILOVER_TEST_EMPTY_KEY", "") + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_EMPTY_KEY"}.ResolveAPIKey(os.Getenv) + Expect(err).To(HaveOccurred()) + }) + It("reads and trims the key file", func() { + f := filepath.Join(GinkgoT().TempDir(), "key") + Expect(os.WriteFile(f, []byte(" k2\n"), 0o600)).To(Succeed()) + Expect(ProxyConfig{APIKeyFile: f}.ResolveAPIKey(os.Getenv)).To(Equal("k2")) + }) + It("returns empty when nothing is set", func() { + Expect(ProxyConfig{}.ResolveAPIKey(os.Getenv)).To(Equal("")) + }) +}) diff --git a/core/config/model_config_loader.go b/core/config/model_config_loader.go index 7733a7392..7c1157ace 100644 --- a/core/config/model_config_loader.go +++ b/core/config/model_config_loader.go @@ -17,6 +17,7 @@ import ( "github.com/charmbracelet/glamour" "github.com/mudler/LocalAI/core/schema" "github.com/mudler/LocalAI/pkg/downloader" + "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/modelartifacts" "github.com/mudler/LocalAI/pkg/safefile" "github.com/mudler/LocalAI/pkg/utils" @@ -501,6 +502,102 @@ func (bcl *ModelConfigLoader) ValidateAliasTarget(cfg *ModelConfig) error { return nil } +// failoverUsecases are the single usecases a chain can share. Checking one +// flag at a time avoids treating "chat+tts" and "tts" as unrelated. +var failoverUsecases = []ModelConfigUsecase{ + FLAG_CHAT, FLAG_COMPLETION, FLAG_EMBEDDINGS, FLAG_RERANK, FLAG_IMAGE, + FLAG_TRANSCRIPT, FLAG_TTS, FLAG_SOUND_GENERATION, FLAG_VAD, FLAG_VIDEO, + FLAG_SOUND_CLASSIFICATION, +} + +// ValidateFailoverTargets checks that every target of a chain exists and is +// not itself a chain. Alias targets are allowed and resolve one hop. +func (bcl *ModelConfigLoader) ValidateFailoverTargets(cfg *ModelConfig) error { + return validateFailoverTargets(cfg, bcl.GetModelConfig) +} + +// FailoverTargetsShareUsecase reports whether all targets of a chain have at +// least one usecase in common. A false result is only a warning: usecases are +// often inferred. +func (bcl *ModelConfigLoader) FailoverTargetsShareUsecase(cfg *ModelConfig) bool { + return failoverTargetsShareUsecase(cfg, bcl.GetModelConfig) +} + +func validateFailoverTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) error { + if cfg == nil || !cfg.IsFailover() { + return nil + } + for _, t := range cfg.Failover.Targets { + target, ok := lookup(t.Model) + if !ok { + return fmt.Errorf("failover chain %q: target %q does not exist", cfg.Name, t.Model) + } + if target.IsAlias() { + if resolved, ok := lookup(target.Alias); ok { + target = resolved + } + } + if target.IsFailover() { + return fmt.Errorf("failover chain %q: target %q is a chain (chains do not nest)", cfg.Name, t.Model) + } + } + return nil +} + +// failoverWarmRemoteTargets lists the chain's targets marked warm that are +// remote. Warm only keeps a local model loaded, so the flag does nothing there. +func failoverWarmRemoteTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) []string { + if cfg == nil || !cfg.IsFailover() { + return nil + } + var out []string + for _, t := range cfg.Failover.Targets { + if !t.Warm { + continue + } + target, ok := lookup(t.Model) + if ok && target.IsAlias() { + target, ok = lookup(target.Alias) + } + if ok && target.IsRemoteProxy() { + out = append(out, t.Model) + } + } + return out +} + +func failoverTargetsShareUsecase(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) bool { + if cfg == nil || !cfg.IsFailover() { + return true + } + var targets []ModelConfig + for _, t := range cfg.Failover.Targets { + target, ok := lookup(t.Model) + if !ok { + return true // missing targets are reported by validateFailoverTargets + } + if target.IsAlias() { + if resolved, ok := lookup(target.Alias); ok { + target = resolved + } + } + targets = append(targets, target) + } + for _, u := range failoverUsecases { + all := true + for i := range targets { + if !targets[i].HasUsecases(u) { + all = false + break + } + } + if all { + return true + } + } + return false +} + type preloadWork struct { key string config ModelConfig @@ -837,6 +934,47 @@ func (bcl *ModelConfigLoader) loadModelConfigsFromPath(path string, strict bool, } } + // Reject failover chains whose targets are missing or are themselves + // chains. bcl.Lock() is held here, so look up configs directly rather + // than through GetModelConfig, which would deadlock on the same mutex. + lookup := func(n string) (ModelConfig, bool) { c, ok := bcl.configs[n]; return c, ok } + for name, cfg := range bcl.configs { + if !cfg.IsFailover() { + continue + } + c := cfg + if err := validateFailoverTargets(&c, lookup); err != nil { + if strict { + return fmt.Errorf("invalid model config %q: %w", name, err) + } + xlog.Error("skipping invalid failover chain", "model", name, "error", err) + delete(bcl.configs, name) + continue + } + if !failoverTargetsShareUsecase(&c, lookup) { + xlog.Warn("failover chain targets share no known usecase", "model", name) + } + if remote := failoverWarmRemoteTargets(&c, lookup); len(remote) > 0 { + xlog.Warn("failover chain: warm has no effect on remote targets", "model", name, "targets", remote) + } + } + + // localai-proxy proxies straight through to another LocalAI instance, so + // it never translates the wire protocol the way cloud-proxy does; warn + // when a config carries settings that only make sense there, or lacks + // the usecases failover's own usecase-sharing check depends on. + for name, cfg := range bcl.configs { + if cfg.Backend != model.LocalAIProxyBackend { + continue + } + if cfg.Proxy.Mode == ProxyModeTranslate || cfg.Proxy.Provider != "" { + xlog.Warn("localai-proxy backend proxies to another LocalAI instance and ignores proxy.mode/proxy.provider", "model", name) + } + if len(cfg.KnownUsecaseStrings) == 0 { + xlog.Warn("localai-proxy config has no known_usecases; failover usecase matching will skip it", "model", name) + } + } + return nil } diff --git a/core/config/model_config_loader_test.go b/core/config/model_config_loader_test.go index 1a3e9b03a..e99a9136d 100644 --- a/core/config/model_config_loader_test.go +++ b/core/config/model_config_loader_test.go @@ -1,7 +1,9 @@ package config import ( + "bytes" "context" + "log/slog" "os" "path/filepath" @@ -9,6 +11,7 @@ import ( . "github.com/onsi/gomega" "github.com/mudler/LocalAI/pkg/modelartifacts" + "github.com/mudler/xlog" ) type preloadArtifactMaterializer struct { @@ -402,3 +405,127 @@ var _ = Describe("ModelConfigLoader ResolveAliasName", func() { Expect(target).To(BeEmpty()) }) }) + +var _ = Describe("ModelConfigLoader failover validation", func() { + var loader *ModelConfigLoader + chain := func(targets ...string) *ModelConfig { + c := &ModelConfig{Name: "chain", Failover: &FailoverConfig{}} + for _, t := range targets { + c.Failover.Targets = append(c.Failover.Targets, FailoverTarget{Model: t}) + } + return c + } + + BeforeEach(func() { + loader = NewModelConfigLoader("") + loader.configs["a"] = ModelConfig{Name: "a", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + loader.configs["b"] = ModelConfig{Name: "b", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + loader.configs["tts"] = ModelConfig{Name: "tts", Backend: "piper", KnownUsecaseStrings: []string{"tts"}} + loader.configs["alias-b"] = ModelConfig{Name: "alias-b", Alias: "b"} + loader.configs["other-chain"] = *chain("a", "b") + loader.configs["alias-chain"] = ModelConfig{Name: "alias-chain", Alias: "other-chain"} + for k, c := range loader.configs { + c.KnownUsecases = GetUsecasesFromYAML(c.KnownUsecaseStrings) + loader.configs[k] = c + } + }) + + It("accepts existing targets and alias targets", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "alias-b"))).To(Succeed()) + }) + It("rejects a missing target", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "nope"))).To(MatchError(ContainSubstring("does not exist"))) + }) + It("rejects a nested chain, directly or through an alias", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "other-chain"))).To(MatchError(ContainSubstring("chains do not nest"))) + Expect(loader.ValidateFailoverTargets(chain("a", "alias-chain"))).To(MatchError(ContainSubstring("chains do not nest"))) + }) + It("reports whether targets share a usecase", func() { + Expect(loader.FailoverTargetsShareUsecase(chain("a", "b"))).To(BeTrue()) + Expect(loader.FailoverTargetsShareUsecase(chain("a", "tts"))).To(BeFalse()) + }) + It("finds warm targets that are remote, where warm has no effect", func() { + loader.configs["remote"] = ModelConfig{Name: "remote", Backend: "cloud-proxy"} + loader.configs["alias-remote"] = ModelConfig{Name: "alias-remote", Alias: "remote"} + c := chain("remote", "alias-remote", "a") + for i := range c.Failover.Targets { + c.Failover.Targets[i].Warm = true + } + Expect(failoverWarmRemoteTargets(c, loader.GetModelConfig)).To(Equal([]string{"remote", "alias-remote"})) + Expect(failoverWarmRemoteTargets(chain("remote", "a"), loader.GetModelConfig)).To(BeEmpty()) + }) +}) + +var _ = Describe("ModelConfigLoader localai-proxy load-time warnings", func() { + var captured *bytes.Buffer + + BeforeEach(func() { + captured = &bytes.Buffer{} + handler := slog.NewTextHandler(captured, &slog.HandlerOptions{Level: slog.LevelWarn}) + xlog.SetLogger(xlog.NewLoggerWithHandler(handler, xlog.LogLevelWarn)) + }) + + AfterEach(func() { + // xlog exposes no getter for the package logger, so restore the same + // default the suite entrypoint installs rather than the prior value. + xlog.SetLogger(xlog.NewLogger(xlog.LogLevel("info"), "text")) + }) + + It("warns and still loads when proxy.mode/proxy.provider are set (ignored by localai-proxy)", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied +backend: localai-proxy +known_usecases: [chat] +proxy: + mode: translate + provider: openai + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + _, ok := loader.GetModelConfig("proxied") + Expect(ok).To(BeTrue()) + Expect(captured.String()).To(ContainSubstring("proxy.mode/proxy.provider")) + Expect(captured.String()).To(ContainSubstring("proxied")) + }) + + It("warns and still loads when known_usecases is empty", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied-no-usecase +backend: localai-proxy +proxy: + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + _, ok := loader.GetModelConfig("proxied-no-usecase") + Expect(ok).To(BeTrue()) + Expect(captured.String()).To(ContainSubstring("known_usecases")) + Expect(captured.String()).To(ContainSubstring("proxied-no-usecase")) + }) + + It("does not warn when localai-proxy uses only passthrough and known_usecases", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied-clean +backend: localai-proxy +known_usecases: [chat] +proxy: + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + Expect(captured.String()).To(BeEmpty()) + }) +}) diff --git a/core/http/app.go b/core/http/app.go index c94b323e4..4543a0048 100644 --- a/core/http/app.go +++ b/core/http/app.go @@ -475,6 +475,7 @@ func API(application *application.Application) (*echo.Echo, error) { mcpJobsMw := auth.RequireFeature(application.AuthDB(), auth.FeatureMCPJobs) requestExtractor := httpMiddleware.NewRequestExtractor(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()) + requestExtractor.SetFailoverManager(application.FailoverManager()) // Register auth routes (login, callback, API keys, user management) routes.RegisterAuthRoutes(e, application) diff --git a/core/http/auth/helpers_test.go b/core/http/auth/helpers_test.go index 1e31ac27f..1b52d4500 100644 --- a/core/http/auth/helpers_test.go +++ b/core/http/auth/helpers_test.go @@ -78,6 +78,9 @@ func newAuthTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo e.GET("/api/settings", ok) e.POST("/api/settings", ok) + // Failover chain reads and the event stream: standard auth, no admin gate. + e.GET("/api/failover", ok) + // Auth routes (exempt) e.GET("/api/auth/status", ok) e.GET("/api/auth/github/login", ok) @@ -139,6 +142,10 @@ func newAdminTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Ech e.POST("/backends/apply", ok, adminMw) e.GET("/api/agents", ok, adminMw) + // Failover chain pin/unpin (admin only) + e.POST("/api/failover/:chain/pin", ok, adminMw) + e.DELETE("/api/failover/:chain/pin", ok, adminMw) + // Trace/log endpoints (admin only) e.GET("/api/traces", ok, adminMw) e.POST("/api/traces/clear", ok, adminMw) diff --git a/core/http/auth/middleware_test.go b/core/http/auth/middleware_test.go index bdfadaafe..bcbba0096 100644 --- a/core/http/auth/middleware_test.go +++ b/core/http/auth/middleware_test.go @@ -223,6 +223,17 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusOK)) }) + It("allows requests to the failover status endpoint with a valid session", func() { + sessionID := createTestSession(db, user.ID) + rec := doRequest(app, http.MethodGet, "/api/failover", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("returns 401 for the failover status endpoint without credentials", func() { + rec := doRequest(app, http.MethodGet, "/api/failover") + Expect(rec.Code).To(Equal(http.StatusUnauthorized)) + }) + It("allows authenticated users to call moderation by default", func() { sessionID := createTestSession(db, user.ID) rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID)) @@ -526,6 +537,30 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusForbidden)) }) + It("allows admin to pin and unpin failover chains", func() { + admin := createTestUser(db, "admin5@example.com", auth.RoleAdmin, auth.ProviderGitHub) + sessionID := createTestSession(db, admin.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("blocks non-admin from pinning or unpinning failover chains", func() { + user := createTestUser(db, "user5@example.com", auth.RoleUser, auth.ProviderGitHub) + sessionID := createTestSession(db, user.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + }) + It("allows user to access regular inference endpoints", func() { user := createTestUser(db, "user@example.com", auth.RoleUser, auth.ProviderGitHub) sessionID := createTestSession(db, user.ID) diff --git a/core/http/endpoints/localai/api_instructions.go b/core/http/endpoints/localai/api_instructions.go index f8af73de2..8d0ea6d2f 100644 --- a/core/http/endpoints/localai/api_instructions.go +++ b/core/http/endpoints/localai/api_instructions.go @@ -129,6 +129,12 @@ var instructionDefs = []instructionDef{ Tags: []string{"middleware", "pii", "router"}, Intro: "GET /api/middleware/status is the single round-trip the /app/middleware admin page reads to render the current state: every model's resolved PII enabled state and the NER detector models it references, recent event count, and the active routing models with their classifier configurations. Admin-only (the synthetic local user is admin in no-auth mode). PII detection policy is edited on each detector model's `pii_detection:` block via the model-config tools/UI — there is no global pattern set to mutate. GET /api/router/decisions returns the routing decision log filtered by correlation_id / user_id / router_model. The same surface is exposed as MCP tools (`get_middleware_status`, `get_pii_events`, `get_router_decisions`) for agent-driven inspection.", }, + { + Name: "failover", + Description: "Model failover chains: target health, pinning and switch events", + Tags: []string{"failover"}, + Intro: "A failover chain is a model config with a failover block. Requests for the chain name are served by its highest-priority healthy target; the X-LocalAI-Served-Model response header names it. Subscribe to GET /api/failover/events (SSE) to follow switches.", + }, { Name: "intelligent-routing", Description: "Per-model `router:` configuration that classifies requests and rewrites the served model", diff --git a/core/http/endpoints/localai/api_instructions_test.go b/core/http/endpoints/localai/api_instructions_test.go index 710d4d982..f42e1c92d 100644 --- a/core/http/endpoints/localai/api_instructions_test.go +++ b/core/http/endpoints/localai/api_instructions_test.go @@ -39,7 +39,7 @@ var _ = Describe("API Instructions Endpoints", func() { instructions, ok := resp["instructions"].([]any) Expect(ok).To(BeTrue()) - Expect(instructions).To(HaveLen(19)) + Expect(instructions).To(HaveLen(20)) // Verify each instruction has required fields and correct URL format for _, s := range instructions { @@ -81,6 +81,7 @@ var _ = Describe("API Instructions Endpoints", func() { "intelligent-routing", "voice-library", "3d", + "failover", )) }) }) diff --git a/core/http/endpoints/localai/failover.go b/core/http/endpoints/localai/failover.go new file mode 100644 index 000000000..608fe69f2 --- /dev/null +++ b/core/http/endpoints/localai/failover.go @@ -0,0 +1,169 @@ +package localai + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" +) + +type FailoverChainsResponse struct { + Chains []failover.ChainStatus `json:"chains"` +} + +type FailoverPinRequest struct { + Target string `json:"target"` +} + +func failoverError(c echo.Context, code int, msg string) error { + return c.JSON(code, schema.ErrorResponse{Error: &schema.APIError{Message: msg, Code: code, Type: "failover_error"}}) +} + +// ListFailoverChainsEndpoint lists failover chains and the health of their targets +// +// @Summary List failover chains and the health of their targets +// @Tags failover +// @Produce json +// @Success 200 {object} FailoverChainsResponse +// @Router /api/failover [get] +func ListFailoverChainsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + return c.JSON(http.StatusOK, FailoverChainsResponse{Chains: fm.Status()}) + } +} + +// GetFailoverChainEndpoint returns one failover chain +// +// @Summary Get one failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain} [get] +func GetFailoverChainEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + st, ok := fm.ChainStatus(c.Param("chain")) + if !ok { + return failoverError(c, http.StatusNotFound, fmt.Sprintf("failover chain %q not found", c.Param("chain"))) + } + return c.JSON(http.StatusOK, st) + } +} + +// PinFailoverTargetEndpoint forces a chain to one target +// +// @Summary Pin a failover chain to one target +// @Tags failover +// @Accept json +// @Produce json +// @Param chain path string true "Chain name" +// @Param request body FailoverPinRequest true "Target to pin" +// @Success 200 {object} failover.ChainStatus +// @Failure 400 {object} schema.ErrorResponse +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [post] +func PinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + var req FailoverPinRequest + if err := c.Bind(&req); err != nil || req.Target == "" { + return failoverError(c, http.StatusBadRequest, "request body must set \"target\"") + } + chain := c.Param("chain") + if err := fm.Pin(chain, req.Target); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +// UnpinFailoverTargetEndpoint removes a pin +// +// @Summary Remove the pin from a failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [delete] +func UnpinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + chain := c.Param("chain") + if err := fm.Unpin(chain); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +func pinError(c echo.Context, err error) error { + switch { + case errors.Is(err, failover.ErrChainNotFound): + return failoverError(c, http.StatusNotFound, err.Error()) + case errors.Is(err, failover.ErrTargetNotInChain): + return failoverError(c, http.StatusBadRequest, err.Error()) + } + return failoverError(c, http.StatusInternalServerError, err.Error()) +} + +// FailoverEventsEndpoint streams failover events +// +// @Summary Stream failover events (server-sent events) +// @Description The first event is "snapshot" with the full state, then "chain.switched" and "target.state" events. +// @Tags failover +// @Produce text/event-stream +// @Success 200 +// @Router /api/failover/events [get] +func FailoverEventsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + // Subscribe before the snapshot so no event falls between the two. + events, cancel := fm.Subscribe(64) + defer cancel() + w := c.Response() + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + send := func(name string, v any) error { + data, err := json.Marshal(v) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", name, data); err != nil { + return err + } + w.Flush() + return nil + } + if err := send("snapshot", FailoverChainsResponse{Chains: fm.Status()}); err != nil { + return nil + } + keepalive := time.NewTicker(15 * time.Second) + defer keepalive.Stop() + for { + select { + case <-c.Request().Context().Done(): + return nil + case <-keepalive.C: + if _, err := fmt.Fprint(w, ": keepalive\n\n"); err != nil { + return nil + } + w.Flush() + case ev, ok := <-events: + if !ok { + return nil + } + if err := send(string(ev.Type), ev); err != nil { + return nil + } + } + } + } +} diff --git a/core/http/endpoints/localai/failover_gap_internal_test.go b/core/http/endpoints/localai/failover_gap_internal_test.go new file mode 100644 index 000000000..553bd2fc5 --- /dev/null +++ b/core/http/endpoints/localai/failover_gap_internal_test.go @@ -0,0 +1,70 @@ +package localai + +import ( + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +// The non-OpenAI endpoints (depth, detection, face_*, voice_*, images, video, +// 3d) return mapBackendError's echo 501 for a backend without the method, not +// the raw gRPC status. The failover retry must still see a capability gap. +var _ = Describe("mapBackendError under a failover chain", func() { + It("spills a backend's Unimplemented to the next target without tripping it", func() { + dir := GinkgoT().TempDir() + write := func(name, body string) { + Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed()) + } + write("a", "name: a\nbackend: fake-a\n") + write("b", "name: b\nbackend: fake-b\n") + write("chain", "name: chain\nfailover:\n targets:\n - model: a\n - model: b\n") + + ss := &system.SystemState{Model: system.Model{ModelsPath: dir}} + appConfig := config.NewApplicationConfig() + appConfig.SystemState = ss + mcl := config.NewModelConfigLoader(dir) + Expect(mcl.LoadModelConfigsFromPath(dir)).To(Succeed()) + re := middleware.NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig) + fm := failover.New(mcl) + fm.Sync() + re.SetFailoverManager(fm) + + var calls []string + handler := func(c echo.Context) error { + cfg := c.Get(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + calls = append(calls, cfg.Name) + if cfg.Name == "a" { + return mapBackendError(grpcstatus.Error(codes.Unimplemented, "unimplemented: Detect")) + } + return c.JSON(http.StatusOK, map[string]string{"served": cfg.Name}) + } + app := echo.New() + app.POST("/v1/detection", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.DetectionRequest) })) + + req := httptest.NewRequest(http.MethodPost, "/v1/detection", strings.NewReader(`{"model":"chain","image":"x"}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) +}) diff --git a/core/http/endpoints/localai/failover_test.go b/core/http/endpoints/localai/failover_test.go new file mode 100644 index 000000000..77fb48678 --- /dev/null +++ b/core/http/endpoints/localai/failover_test.go @@ -0,0 +1,110 @@ +package localai + +import ( + "bufio" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type mapSource map[string]config.ModelConfig + +func (s mapSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok } +func (s mapSource) GetAllModelsConfigs() []config.ModelConfig { + var out []config.ModelConfig + for _, c := range s { + out = append(out, c) + } + return out +} + +var _ = Describe("failover endpoints", func() { + var ( + e *echo.Echo + fm *failover.Manager + ) + + BeforeEach(func() { + src := mapSource{ + "a": {Name: "a", Backend: "cloud-proxy"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}}, + } + fm = failover.New(src) + e = echo.New() + e.GET("/api/failover", ListFailoverChainsEndpoint(fm)) + e.GET("/api/failover/events", FailoverEventsEndpoint(fm)) + e.GET("/api/failover/:chain", GetFailoverChainEndpoint(fm)) + e.POST("/api/failover/:chain/pin", PinFailoverTargetEndpoint(fm)) + e.DELETE("/api/failover/:chain/pin", UnpinFailoverTargetEndpoint(fm)) + }) + + do := func(method, path, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(method, path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + return rec + } + + It("lists chains", func() { + rec := do(http.MethodGet, "/api/failover", "") + Expect(rec.Code).To(Equal(http.StatusOK)) + var out FailoverChainsResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &out)).To(Succeed()) + Expect(out.Chains).To(HaveLen(1)) + Expect(out.Chains[0].Active).To(Equal("a")) + }) + + It("gets one chain or 404", func() { + Expect(do(http.MethodGet, "/api/failover/chain", "").Code).To(Equal(http.StatusOK)) + Expect(do(http.MethodGet, "/api/failover/nope", "").Code).To(Equal(http.StatusNotFound)) + }) + + It("pins and unpins", func() { + rec := do(http.MethodPost, "/api/failover/chain/pin", `{"target":"b"}`) + Expect(rec.Code).To(Equal(http.StatusOK)) + st, _ := fm.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{"target":"zzz"}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/nope/pin", `{"target":"a"}`).Code).To(Equal(http.StatusNotFound)) + Expect(do(http.MethodDelete, "/api/failover/chain/pin", "").Code).To(Equal(http.StatusOK)) + st, _ = fm.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + + It("streams a snapshot, then switch events", func() { + srv := httptest.NewServer(e) + defer srv.Close() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil) + resp, err := http.DefaultClient.Do(req) + Expect(err).ToNot(HaveOccurred()) + defer func() { _ = resp.Body.Close() }() + Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream")) + r := bufio.NewReader(resp.Body) + next := func() string { + for { + line, err := r.ReadString('\n') + Expect(err).ToNot(HaveOccurred()) + if strings.HasPrefix(line, "event: ") { + return strings.TrimSpace(strings.TrimPrefix(line, "event: ")) + } + } + } + Expect(next()).To(Equal("snapshot")) + Expect(fm.Pin("chain", "b")).To(Succeed()) + Expect(next()).To(Equal("chain.switched")) + }) +}) diff --git a/core/http/endpoints/localai/import_model.go b/core/http/endpoints/localai/import_model.go index b3cf491eb..73448430a 100644 --- a/core/http/endpoints/localai/import_model.go +++ b/core/http/endpoints/localai/import_model.go @@ -187,6 +187,12 @@ func ImportModelEndpoint(cl *config.ModelConfigLoader, gs *galleryop.GalleryServ return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) } + // Reject failover chains whose targets are missing or are themselves + // chains, for the same reason. + if err := cl.ValidateFailoverTargets(&modelConfig); err != nil { + return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) + } + // Create the configuration file configPath := filepath.Join(appConfig.SystemState.Model.ModelsPath, modelConfig.Name+".yaml") if err := utils.VerifyPath(modelConfig.Name+".yaml", appConfig.SystemState.Model.ModelsPath); err != nil { diff --git a/core/http/endpoints/mcp/localai_assistant_test.go b/core/http/endpoints/mcp/localai_assistant_test.go index 71ce38a7a..8629bd9e1 100644 --- a/core/http/endpoints/mcp/localai_assistant_test.go +++ b/core/http/endpoints/mcp/localai_assistant_test.go @@ -214,3 +214,11 @@ func (stubClient) SeedRouterCorpus(_ context.Context, req localaitools.RouterCor func (stubClient) ClearRouterCorpus(_ context.Context, routerModel string) (*localaitools.RouterCorpusClearResult, error) { return &localaitools.RouterCorpusClearResult{Router: routerModel}, nil } + +func (stubClient) ListFailoverChains(_ context.Context) ([]localaitools.FailoverChainInfo, error) { + return []localaitools.FailoverChainInfo{}, nil +} + +func (stubClient) PinFailoverTarget(_ context.Context, _, _ string) error { return nil } + +func (stubClient) UnpinFailoverTarget(_ context.Context, _ string) error { return nil } diff --git a/core/http/endpoints/openai/audio_upload_test.go b/core/http/endpoints/openai/audio_upload_test.go new file mode 100644 index 000000000..ee788a4fd --- /dev/null +++ b/core/http/endpoints/openai/audio_upload_test.go @@ -0,0 +1,78 @@ +package openai + +import ( + "errors" + "net/http" + "net/http/httptest" + "strings" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/schema" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// A broken or missing upload is the caller's fault: the audio endpoints must +// answer 400, not surface the multipart parser error as a 500. +var _ = Describe("audio upload endpoints reject bad uploads as client errors", func() { + type endpointCase struct { + path string + handler func() echo.HandlerFunc + } + + cases := map[string]endpointCase{ + "transcription": {"/v1/audio/transcriptions", func() echo.HandlerFunc { + return TranscriptEndpoint(nil, nil, config.NewApplicationConfig()) + }}, + "diarization": {"/v1/audio/diarization", func() echo.HandlerFunc { + return DiarizationEndpoint(nil, nil, config.NewApplicationConfig()) + }}, + "sound classification": {"/v1/audio/classifications", func() echo.HandlerFunc { + return SoundClassificationEndpoint(nil, nil, config.NewApplicationConfig()) + }}, + } + + run := func(ec endpointCase, contentType, body string) error { + req := httptest.NewRequest(http.MethodPost, ec.path, strings.NewReader(body)) + req.Header.Set("Content-Type", contentType) + c := echo.New().NewContext(req, httptest.NewRecorder()) + c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.OpenAIRequest{PredictionOptions: schema.PredictionOptions{BasicModelRequest: schema.BasicModelRequest{Model: "m"}}}) + c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{}) + return ec.handler()(c) + } + + expectBadRequest := func(err error) { + Expect(err).To(HaveOccurred()) + var he *echo.HTTPError + Expect(errors.As(err, &he)).To(BeTrue(), "expected *echo.HTTPError, got %T: %v", err, err) + Expect(he.Code).To(Equal(http.StatusBadRequest)) + } + + // An unknown response_format is the caller's fault too, and must be + // rejected before the backend runs: a failover chain would otherwise + // transcribe on every target and count the error against each of them. + for name, ec := range map[string]endpointCase{ + "transcription": cases["transcription"], + "diarization": cases["diarization"], + } { + It(name+": unknown response_format", func() { + body := "--xyz\r\nContent-Disposition: form-data; name=\"response_format\"\r\n\r\nbogus\r\n" + + "--xyz\r\nContent-Disposition: form-data; name=\"file\"; filename=\"a.wav\"\r\n\r\nRIFF\r\n--xyz--\r\n" + var err error + Expect(func() { err = run(ec, "multipart/form-data; boundary=xyz", body) }).NotTo(Panic(), "the backend must not be reached") + expectBadRequest(err) + }) + } + + for name, ec := range cases { + It(name+": multipart content type without a boundary", func() { + expectBadRequest(run(ec, "multipart/form-data", "")) + }) + It(name+": multipart body without the file field", func() { + body := "--xyz\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\nm\r\n--xyz--\r\n" + expectBadRequest(run(ec, "multipart/form-data; boundary=xyz", body)) + }) + } +}) diff --git a/core/http/endpoints/openai/diarization.go b/core/http/endpoints/openai/diarization.go index 75e9715db..b39935e8b 100644 --- a/core/http/endpoints/openai/diarization.go +++ b/core/http/endpoints/openai/diarization.go @@ -1,7 +1,6 @@ package openai import ( - "errors" "fmt" "io" "net/http" @@ -75,8 +74,15 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap if responseFormat == "" { responseFormat = schema.DiarizationResponseFormatJson } + switch responseFormat { + case schema.DiarizationResponseFormatJson, schema.DiarizationResponseFormatJsonVerbose, schema.DiarizationResponseFormatRTTM: + default: + // Checked before the backend runs, for the same reason as in + // TranscriptEndpoint. + return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format (expected: json, verbose_json, rttm)") + } - file, err := c.FormFile("file") + file, err := uploadedFile(c, "file") if err != nil { return err } @@ -126,7 +132,7 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap case schema.DiarizationResponseFormatJsonVerbose: return c.JSON(http.StatusOK, result) default: - return errors.New("invalid response_format (expected: json, verbose_json, rttm)") + return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format (expected: json, verbose_json, rttm)") } } } diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index 9db5f899f..11d1653cb 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -32,6 +32,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/turncoord" "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" "github.com/mudler/LocalAI/core/templates" @@ -638,6 +639,7 @@ func runRealtimeSession(application *application.Application, t Transport, model application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), + application.FailoverManager(), ) } else { m, err = newModel( @@ -709,7 +711,7 @@ func runRealtimeSession(application *application.Application, t Transport, model var gateErr error if session.voiceGate != nil { _, gateErr = backend.PreloadStages(context.Background(), application.ModelLoader(), application.ApplicationConfig(), []backend.PreloadStage{ - {Role: "voice_recognition", Cfg: session.voiceGate.recCfg}, + {Role: config.PipelineStageVoiceRecognition, Cfg: session.voiceGate.recCfg}, }) } if err := errors.Join(<-warmErr, gateErr); err != nil { @@ -754,6 +756,13 @@ func runRealtimeSession(application *application.Application, t Transport, model Session: session.ToServer(), }) + // Sent after session.created, which clients expect as the first event. + // This function runs until the connection closes, so the defer stops the + // events at session end. A transcription session.update that swaps the + // model restarts them for the new model's chains. + stopFailoverEvents := startModelFailoverEvents(t, m) + defer func() { stopFailoverEvents() }() + var ( msg []byte wg sync.WaitGroup @@ -824,12 +833,14 @@ func runRealtimeSession(application *application.Application, t Transport, model // Handle transcription session update if e.Session.Transcription != nil { + prevModel := session.ModelInterface if err := updateTransSession( session, &e.Session, application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), + application.FailoverManager(), ); err != nil { xlog.Error("failed to update session", "error", err) // The cause is validation feedback on the client's own @@ -846,6 +857,10 @@ func runRealtimeSession(application *application.Application, t Transport, model }, Session: session.ToServer(), }) + if session.ModelInterface != prevModel { + stopFailoverEvents() + stopFailoverEvents = startModelFailoverEvents(t, session.ModelInterface) + } } // Handle realtime session update @@ -1143,7 +1158,7 @@ func sendTestTone(t Transport) { } } -func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) error { +func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) error { sessionLock.Lock() defer sessionLock.Unlock() @@ -1166,7 +1181,7 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config return fmt.Errorf("model is not a valid pipeline model: %s", trUpd.Model) } - m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig) + m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig, fm) if err != nil { return err } diff --git a/core/http/endpoints/openai/realtime_classifier_test.go b/core/http/endpoints/openai/realtime_classifier_test.go index b8630ca81..224808350 100644 --- a/core/http/endpoints/openai/realtime_classifier_test.go +++ b/core/http/endpoints/openai/realtime_classifier_test.go @@ -45,7 +45,7 @@ var classifierTestHistory = schema.Messages{ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { var out []types.ClassifierResultEvent - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ClassifierResultEvent); ok { out = append(out, ev) } @@ -57,7 +57,7 @@ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { // item — what a classifier response actually "spoke". func replyTexts(t *fakeTransport) []string { var out []string - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ResponseOutputTextDoneEvent); ok { out = append(out, ev.Text) } @@ -277,7 +277,7 @@ var _ = Describe("classifierRespond", func() { Expect(t.countEvents(types.ServerEventTypeResponseOutputTextDone)).To(Equal(1)) Expect(t.countEvents(types.ServerEventTypeResponseFunctionCallArgumentsDone)).To(Equal(1)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -656,7 +656,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(results[0].Arguments).To(MatchJSON(`{"direction":"up","distance":3,"units":"meters"}`)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -695,7 +695,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(handled).To(BeTrue()) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } diff --git a/core/http/endpoints/openai/realtime_doubles_test.go b/core/http/endpoints/openai/realtime_doubles_test.go index a2c104b3c..e82a4a130 100644 --- a/core/http/endpoints/openai/realtime_doubles_test.go +++ b/core/http/endpoints/openai/realtime_doubles_test.go @@ -17,8 +17,10 @@ import ( // so streaming behaviour can be asserted without a real WebSocket/WebRTC peer. // It is not a *WebRTCTransport, so handler code takes the WebSocket path. type fakeTransport struct { - events []types.ServerEvent - audio []fakeAudioChunk + // mu guards sent: some specs send from a background goroutine. + mu sync.Mutex + sent []types.ServerEvent + audio []fakeAudioChunk } type fakeAudioChunk struct { @@ -27,10 +29,19 @@ type fakeAudioChunk struct { } func (f *fakeTransport) SendEvent(e types.ServerEvent) error { - f.events = append(f.events, e) + f.mu.Lock() + defer f.mu.Unlock() + f.sent = append(f.sent, e) return nil } +// events returns a copy of the server events sent so far. +func (f *fakeTransport) events() []types.ServerEvent { + f.mu.Lock() + defer f.mu.Unlock() + return append([]types.ServerEvent(nil), f.sent...) +} + func (f *fakeTransport) ReadEvent() ([]byte, error) { return nil, nil } func (f *fakeTransport) SendAudio(_ context.Context, pcm []byte, sampleRate int) error { @@ -43,7 +54,7 @@ func (f *fakeTransport) Close() error { return nil } // countEvents returns how many recorded events have the given type. func (f *fakeTransport) countEvents(et types.ServerEventType) int { n := 0 - for _, e := range f.events { + for _, e := range f.events() { if e.ServerEventType() == et { n++ } @@ -55,7 +66,7 @@ func (f *fakeTransport) countEvents(et types.ServerEventType) int { // delta event — i.e. the text streamed to the client as it is generated. func (f *fakeTransport) transcriptDeltaText() string { var b strings.Builder - for _, e := range f.events { + for _, e := range f.events() { if d, ok := e.(types.ResponseOutputAudioTranscriptDeltaEvent); ok { b.WriteString(d.Delta) } diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go new file mode 100644 index 000000000..c998e05c2 --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover.go @@ -0,0 +1,180 @@ +package openai + +import ( + "context" + "errors" + "fmt" + "sort" + "sync" + + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" +) + +// stageRouter routes realtime pipeline stages that name a failover chain. +// Every realtime model kind (full pipeline, transcription-only, sound-only) +// embeds one, so a chain resolves the same way whatever the session does. +type stageRouter struct { + // failover and stageChains route pipeline stages that name a failover + // chain; stageChains maps a stage (config.PipelineStage*) to its chain. + // The model's *Config fields then hold the target that was active at + // session start, for the checks that run once (voice, templates). + failover *failover.Manager + stageChains map[string]string + stageTargetConfig func(name string) (*config.ModelConfig, error) + appTracing bool +} + +func newStageRouter(fm *failover.Manager, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) stageRouter { + return stageRouter{ + failover: fm, + stageChains: map[string]string{}, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { + cfg, err := cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err != nil { + return nil, err + } + failover.PrepareTarget(cfg) + return cfg, nil + }, + appTracing: appConfig.EnableTracing, + } +} + +// resolveStage records a stage that names a chain and returns the chain's +// active target, so everything that inspects stage configs at session start +// sees a real model. Any other config is returned as is. A chain config +// reaching the model loader would have no backend and trigger backend +// auto-detection. +func (r *stageRouter) resolveStage(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { + if cfg == nil || !cfg.IsFailover() { + return cfg, nil + } + if r.failover == nil { + return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) + } + st, ok := r.failover.ChainStatus(cfg.Name) + if !ok { + return nil, fmt.Errorf("failover chain %q not found", cfg.Name) + } + r.stageChains[stage] = cfg.Name + return r.stageTargetConfig(st.Active) +} + +// isChainStage reports whether stage names a failover chain. +func (r *stageRouter) isChainStage(stage string) bool { + _, ok := r.stageChains[stage] + return ok && r.failover != nil +} + +// hasChainStages reports whether any stage names a chain. +func (r *stageRouter) hasChainStages() bool { + return r.failover != nil && len(r.stageChains) > 0 +} + +func (r *stageRouter) router() *stageRouter { return r } + +// stageRouted is implemented by every realtime model that embeds a +// stageRouter; the session uses it to send failover events. +type stageRouted interface{ router() *stageRouter } + +// stageCall runs fn with the config that should serve stage now. A plain +// stage uses base. A chain stage goes through the failover plan on every +// call, so a switch takes effect on the next call without rebuilding the +// session, and fn is retried on the next target until it calls commit. +func (r *stageRouter) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { + if !r.isChainStage(stage) { + return fn(base, func() {}) + } + chain := r.stageChains[stage] + return r.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { + cfg, err := r.stageTargetConfig(target) + if err != nil { + return err + } + err = fn(cfg, commit) + if err != nil { + failover.RecordAttemptTrace(r.appTracing, chain, target, err) + } + return err + }) +} + +// warmStages preloads the stages. A chain stage warms through its failover +// plan: a target that fails to load moves the stage to the next one instead +// of failing the session. +func (r *stageRouter) warmStages(ctx context.Context, ml *model.ModelLoader, appConfig *config.ApplicationConfig, stages []backend.PreloadStage) error { + var ( + plain []backend.PreloadStage + wg sync.WaitGroup + mu sync.Mutex + errs []error + ) + for _, s := range stages { + if !r.isChainStage(s.Role) { + plain = append(plain, s) + continue + } + wg.Go(func() { + err := r.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error { + _, err := backend.PreloadStages(ctx, ml, appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}}) + return err + }) + mu.Lock() + errs = append(errs, err) + mu.Unlock() + }) + } + _, err := backend.PreloadStages(ctx, ml, appConfig, plain) + wg.Wait() + return errors.Join(append(errs, err)...) +} + +// startModelFailoverEvents starts failover events for m when it has chain +// stages. The returned func stops them and is never nil. +func startModelFailoverEvents(t Transport, m Model) func() { + sr, ok := m.(stageRouted) + if !ok { + return func() {} + } + r := sr.router() + if !r.hasChainStages() { + return func() {} + } + return startFailoverEvents(t, r.failover, r.stageChains) +} + +// startFailoverEvents tells the client which target serves each chain stage +// now, and again whenever a chain switches. The returned func stops it. +func startFailoverEvents(t Transport, fm *failover.Manager, stageChains map[string]string) func() { + // Subscribe before reading the status, so a switch that lands in + // between is still delivered. + events, cancel := fm.Subscribe(16) + stages := make([]string, 0, len(stageChains)) + for s := range stageChains { + stages = append(stages, s) + } + sort.Strings(stages) + for _, stage := range stages { + chain := stageChains[stage] + if st, ok := fm.ChainStatus(chain); ok { + sendEvent(t, types.ModelFailoverEvent{Chain: chain, Stage: stage, To: st.Active, State: string(st.State), Reason: string(failover.ReasonInitial)}) + } + } + go func() { + for ev := range events { + if ev.Type != failover.EventChainSwitched { + continue + } + for _, stage := range stages { + if stageChains[stage] == ev.Chain { + sendEvent(t, types.ModelFailoverEvent{Chain: ev.Chain, Stage: stage, From: ev.From, To: ev.To, State: ev.State, Reason: string(ev.Reason)}) + } + } + } + }() + return cancel +} diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go new file mode 100644 index 000000000..06dfda2b3 --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -0,0 +1,206 @@ +package openai + +import ( + "context" + "errors" + "os" + "path/filepath" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type rtSource map[string]config.ModelConfig + +func (s rtSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok } +func (s rtSource) GetAllModelsConfigs() []config.ModelConfig { + var out []config.ModelConfig + for _, c := range s { + out = append(out, c) + } + return out +} + +var _ = Describe("realtime failover", func() { + var fm *failover.Manager + + BeforeEach(func() { + fm = failover.New(rtSource{ + "a": {Name: "a", Backend: "cloud-proxy"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}}, + }) + }) + + chainModel := func() *wrappedModel { + return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{config.PipelineStageTTS: "chain"}, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }}} + } + + It("routes a chain stage through the plan and retries before commit", func() { + m := chainModel() + var tried []string + err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, _ func()) error { + tried = append(tried, cfg.Name) + if cfg.Name == "a" { + return errors.New("dial tcp: refused") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + }) + + It("does not retry a chain stage once output was committed", func() { + m := chainModel() + var tried []string + err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, commit func()) error { + tried = append(tried, cfg.Name) + commit() + return errors.New("dial tcp: refused") + }) + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"a"})) + }) + + It("calls a plain stage once with its own config", func() { + m := &wrappedModel{} + base := &config.ModelConfig{Name: "plain"} + calls := 0 + err := m.stageCall(context.Background(), config.PipelineStageTTS, base, func(cfg *config.ModelConfig, _ func()) error { + calls++ + Expect(cfg).To(BeIdenticalTo(base)) + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(calls).To(Equal(1)) + }) + + It("sends initial events, then switch events, and stops on cancel", func() { + t := &fakeTransport{} + failoverEvents := func() []types.ModelFailoverEvent { + var out []types.ModelFailoverEvent + for _, e := range t.events() { + if fe, ok := e.(types.ModelFailoverEvent); ok { + out = append(out, fe) + } + } + return out + } + stop := startFailoverEvents(t, fm, map[string]string{config.PipelineStageLLM: "chain"}) + Eventually(failoverEvents).Should(ContainElement(And( + HaveField("Reason", "initial"), HaveField("To", "a"), HaveField("Stage", "llm")))) + fm.ReportFailure("a", errors.New("dial tcp: refused")) + Eventually(failoverEvents).Should(ContainElement(And( + HaveField("Reason", "trip"), HaveField("From", "a"), HaveField("To", "b"), HaveField("Chain", "chain")))) + stop() + }) +}) + +var _ = Describe("realtime failover in transcription-only and sound-only sessions", func() { + var ( + cl *config.ModelConfigLoader + ml *model.ModelLoader + appConfig *config.ApplicationConfig + fm *failover.Manager + ) + + BeforeEach(func() { + dir := GinkgoT().TempDir() + write := func(name, body string) { + Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed()) + } + write("vad", "name: vad\nbackend: silero-vad\n") + write("stt-a", "name: stt-a\nbackend: fake-stt-a\n") + write("stt-b", "name: stt-b\nbackend: fake-stt-b\n") + write("stt-chain", "name: stt-chain\nfailover:\n targets:\n - model: stt-a\n - model: stt-b\n") + write("sound-a", "name: sound-a\nbackend: fake-sound-a\n") + write("sound-b", "name: sound-b\nbackend: fake-sound-b\n") + write("sound-chain", "name: sound-chain\nfailover:\n targets:\n - model: sound-a\n - model: sound-b\n") + ss := &system.SystemState{Model: system.Model{ModelsPath: dir}} + appConfig = config.NewApplicationConfig() + appConfig.SystemState = ss + cl = config.NewModelConfigLoader(dir) + Expect(cl.LoadModelConfigsFromPath(dir)).To(Succeed()) + ml = model.NewModelLoader(ss) + fm = failover.New(cl) + }) + + // failoverEvents collects the failover events a transport received. + failoverEvents := func(t *fakeTransport) func() []types.ModelFailoverEvent { + return func() []types.ModelFailoverEvent { + var out []types.ModelFailoverEvent + for _, e := range t.events() { + if fe, ok := e.(types.ModelFailoverEvent); ok { + out = append(out, fe) + } + } + return out + } + } + + It("resolves a sound_detection chain and routes each call through it", func() { + m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + tm := m.(*transcriptOnlyModel) + Expect(tm.SoundDetectionConfig.Name).To(Equal("sound-a")) + Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageSoundDetection: "sound-chain"})) + + var tried []string + tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { + tried = append(tried, name) + return nil, errors.New("dial tcp: refused") + } + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _, err = tm.SoundDetection(ctx, "a.wav", 3, 0.1) + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"sound-a", "sound-b"})) + + t := &fakeTransport{} + stop := startModelFailoverEvents(t, m) + defer stop() + Eventually(failoverEvents(t)).Should(ContainElement(And( + HaveField("Stage", "sound_detection"), HaveField("Chain", "sound-chain"), HaveField("Reason", "initial")))) + }) + + It("resolves a transcription chain in a transcription-only session", func() { + m, cfg, err := newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + Expect(cfg.Name).To(Equal("stt-a")) + tm := m.(*transcriptOnlyModel) + Expect(tm.VADConfig.Name).To(Equal("vad")) + Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageTranscription: "stt-chain"})) + + var tried []string + tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { + tried = append(tried, name) + return nil, errors.New("dial tcp: refused") + } + _, err = tm.Transcribe(context.Background(), "a.wav", "", false, false, "") + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"stt-a", "stt-b"})) + }) + + It("fails before touching a backend when failover is not running", func() { + _, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + _, _, err = newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + Expect(ml.ListLoadedModels()).To(BeEmpty()) + }) + + It("sends no failover events for a session without chains", func() { + m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-a"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + t := &fakeTransport{} + startModelFailoverEvents(t, m)() + Consistently(failoverEvents(t), 200*time.Millisecond).Should(BeEmpty()) + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index a02696a1f..678a57a0d 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -19,6 +19,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/http/middleware" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" "github.com/mudler/LocalAI/core/templates" @@ -84,6 +85,11 @@ type wrappedModel struct { routerStore router.DecisionStore routerSessionID string routerUserID string + + stageRouter + // tuneLLM applies the pipeline's LLM overrides (reasoning effort, + // disable_thinking) to a chain target loaded per call. + tuneLLM func(cfg *config.ModelConfig) } // anyToAnyModel represent a model which supports Any-to-Any operations @@ -106,18 +112,38 @@ type transcriptOnlyModel struct { appConfig *config.ApplicationConfig modelLoader *model.ModelLoader confLoader *config.ModelConfigLoader + + stageRouter } func (m *transcriptOnlyModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { - return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig) + var res *schema.VADResponse + err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg) + return err + }) + return res, err } func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) { - return backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *cfg, m.appConfig) + return err + }) + return res, err } func (m *transcriptOnlyModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) { - return modelSoundDetection(ctx, m.modelLoader, m.appConfig, m.SoundDetectionConfig, audio, topK, threshold) + var res *schema.SoundClassificationResult + err := m.stageCall(ctx, config.PipelineStageSoundDetection, m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = modelSoundDetection(ctx, m.modelLoader, m.appConfig, cfg, audio, topK, threshold) + return err + }) + return res, err } func (m *transcriptOnlyModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) { @@ -144,11 +170,27 @@ func (m *transcriptOnlyModel) TTSStream(ctx context.Context, text, voice, langua } func (m *transcriptOnlyModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) { - return transcribeStream(ctx, m.modelLoader, *m.TranscriptionConfig, m.appConfig, audio, language, translate, diarize, prompt, onDelta) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error { + var err error + res, err = transcribeStream(ctx, m.modelLoader, *cfg, m.appConfig, audio, language, translate, diarize, prompt, func(s string) { + commit() + onDelta(s) + }) + return err + }) + return res, err } func (m *transcriptOnlyModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) { - return backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *m.TranscriptionConfig, m.appConfig, onEvent) + var live backend.LiveTranscriptionSession + // Only opening the live session can move to the next target. + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent) + return err + }) + return live, err } func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { @@ -156,24 +198,41 @@ func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { } func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ - {Role: "vad", Cfg: m.VADConfig}, - {Role: "transcription", Cfg: m.TranscriptionConfig}, - {Role: "sound_detection", Cfg: m.SoundDetectionConfig}, + return m.warmStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ + {Role: config.PipelineStageVAD, Cfg: m.VADConfig}, + {Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig}, + {Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig}, }) - return err } func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { - return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig) + var res *schema.VADResponse + err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg) + return err + }) + return res, err } func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) { - return backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *cfg, m.appConfig) + return err + }) + return res, err } func (m *wrappedModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) { - return modelSoundDetection(ctx, m.modelLoader, m.appConfig, m.SoundDetectionConfig, audio, topK, threshold) + var res *schema.SoundClassificationResult + err := m.stageCall(ctx, config.PipelineStageSoundDetection, m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = modelSoundDetection(ctx, m.modelLoader, m.appConfig, cfg, audio, topK, threshold) + return err + }) + return res, err } func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) { @@ -181,11 +240,22 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im Messages: messages, } + toolsJSON, toolChoiceJSON := realtimeToolsJSON(tools, toolChoice) + + // infer renders the prompt for cfg and starts inference on it. Everything + // that reads the LLM config lives here, so a chain stage can run it again + // against the next target. + infer := func(cfg *config.ModelConfig, cb func(string, backend.TokenUsage) bool) (func() (backend.LLMResponse, error), error) { + predInput := m.renderPredictPrompt(input, cfg, tools, toolChoice) + return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, cfg, m.confLoader, m.appConfig, cb, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil) + } + // Per-turn routing: when the session's LLMConfig is a router, swap // to the candidate the classifier picks for this turn's prompt. // LLMConfig itself is held by value (we never mutate it) — turnCfg // is the config we dispatch against. turnCfg := m.LLMConfig + routed := false if m.LLMConfig.HasRouter() && m.routerDeps != nil { chosen, err := m.routeTurn(ctx, &input) if err != nil { @@ -193,9 +263,47 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im "router_model", m.LLMConfig.Name, "error", err) } else if chosen != nil { turnCfg = chosen + routed = true } } + // A routed turn dispatches to the router's pick: chains as router + // candidates are not resolved here. + if routed || !m.isChainStage(config.PipelineStageLLM) { + return infer(turnCfg, tokenCallback) + } + + return func() (backend.LLMResponse, error) { + var resp backend.LLMResponse + err := m.stageCall(ctx, config.PipelineStageLLM, turnCfg, func(cfg *config.ModelConfig, commit func()) error { + if m.tuneLLM != nil { + m.tuneLLM(cfg) + } + // Without a callback nothing reaches the client before the + // reply is complete, so every failure can still be retried. + var cb func(string, backend.TokenUsage) bool + if tokenCallback != nil { + cb = func(s string, u backend.TokenUsage) bool { + commit() + return tokenCallback(s, u) + } + } + predict, err := infer(cfg, cb) + if err != nil { + return err + } + resp, err = predict() + return err + }) + return resp, err + }, nil +} + +// renderPredictPrompt templates the turn's prompt for cfg. It also applies +// the turn's tool choice and function-calling grammar to cfg, which the +// backend reads when inference starts. The prompt is empty for models that +// use the tokenizer's template. +func (m *wrappedModel) renderPredictPrompt(input schema.OpenAIRequest, turnCfg *config.ModelConfig, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) string { // Surface the resolved reasoning effort to the Go-side template path too // (jinja models get it via backend metadata in gRPCPredictOpts; Go-templated // models like gpt-oss read it from the template's .ReasoningEffort). @@ -303,6 +411,12 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im } } + return predInput +} + +// realtimeToolsJSON serializes the turn's tools and tool choice the way the +// backends expect them. Neither depends on the LLM config. +func realtimeToolsJSON(tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) (string, string) { var toolsJSON string if len(tools) > 0 { // Convert tools to OpenAI Chat Completions format (nested) @@ -348,7 +462,7 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im toolChoiceJSON = string(b) } - return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, turnCfg, m.confLoader, m.appConfig, tokenCallback, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil) + return toolsJSON, toolChoiceJSON } // routeTurn classifies this turn's prompt against the session's router @@ -395,7 +509,16 @@ func newRealtimeDecisionID() string { } func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) { - return backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *m.TTSConfig) + var ( + out string + res *proto.Result + ) + err := m.stageCall(ctx, config.PipelineStageTTS, m.TTSConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + out, res, err = backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *cfg) + return err + }) + return out, res, err } func (m *wrappedModel) setTTSParams(params map[string]string) { @@ -403,7 +526,14 @@ func (m *wrappedModel) setTTSParams(params map[string]string) { } func (m *wrappedModel) TTSStream(ctx context.Context, text, voice, language string, onAudio func(pcm []byte, sampleRate int) error) error { - return ttsStream(ctx, m.modelLoader, m.appConfig, *m.TTSConfig, text, voice, language, maps.Clone(m.ttsParams), onAudio) + return m.stageCall(ctx, config.PipelineStageTTS, m.TTSConfig, func(cfg *config.ModelConfig, commit func()) error { + // Audio that reached the client cannot be taken back, so the first + // chunk ends the retries. + return ttsStream(ctx, m.modelLoader, m.appConfig, *cfg, text, voice, language, maps.Clone(m.ttsParams), func(pcm []byte, sr int) error { + commit() + return onAudio(pcm, sr) + }) + }) } func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig *config.ModelConfig, profiles *voiceprofile.Store) (string, map[string]string, func(), error) { @@ -431,11 +561,28 @@ func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig } func (m *wrappedModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) { - return transcribeStream(ctx, m.modelLoader, *m.TranscriptionConfig, m.appConfig, audio, language, translate, diarize, prompt, onDelta) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error { + var err error + res, err = transcribeStream(ctx, m.modelLoader, *cfg, m.appConfig, audio, language, translate, diarize, prompt, func(s string) { + commit() + onDelta(s) + }) + return err + }) + return res, err } func (m *wrappedModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) { - return backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *m.TranscriptionConfig, m.appConfig, onEvent) + var live backend.LiveTranscriptionSession + // Only opening the live session can move to the next target: once it is + // open, events flow to the client for the rest of the utterance. + err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent) + return err + }) + return live, err } func (m *wrappedModel) PredictConfig() *config.ModelConfig { @@ -683,18 +830,17 @@ func (m *wrappedModel) FillToolArguments(ctx context.Context, messages schema.Me func (m *wrappedModel) Warmup(ctx context.Context) error { stages := []backend.PreloadStage{ - {Role: "vad", Cfg: m.VADConfig}, - {Role: "transcription", Cfg: m.TranscriptionConfig}, - {Role: "llm", Cfg: m.LLMConfig}, - {Role: "tts", Cfg: m.TTSConfig}, - {Role: "sound_detection", Cfg: m.SoundDetectionConfig}, + {Role: config.PipelineStageVAD, Cfg: m.VADConfig}, + {Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig}, + {Role: config.PipelineStageLLM, Cfg: m.LLMConfig}, + {Role: config.PipelineStageTTS, Cfg: m.TTSConfig}, + {Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig}, } // The scoring model is a separate stage only when it isn't the LLM. if m.ScoreConfig != nil && m.ScoreConfig != m.LLMConfig { stages = append(stages, backend.PreloadStage{Role: "classifier", Cfg: m.ScoreConfig}) } - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, stages) - return err + return m.warmStages(ctx, m.modelLoader, m.appConfig, stages) } // wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream @@ -787,8 +933,12 @@ func loadSoundDetectionConfig(pipeline *config.Pipeline, cl *config.ModelConfigL return cfg, nil } -func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, *config.ModelConfig, error) { +func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, *config.ModelConfig, error) { + sr := newStageRouter(fm, cl, ml, appConfig) cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgVAD, err = sr.resolveStage(config.PipelineStageVAD, cfgVAD) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -799,6 +949,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig } cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgSST, err = sr.resolveStage(config.PipelineStageTranscription, cfgSST) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -809,6 +962,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig } cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound) + } if err != nil { return nil, nil, err } @@ -821,6 +977,7 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, cfgSST, nil } @@ -829,8 +986,12 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig // a sound-detection-only realtime session, which activates on sounds (not // speech) and is driven by client-side windowing (turn_detection none + // input_audio_buffer.commit) rather than the voice VAD loop. -func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, error) { +func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, error) { + sr := newStageRouter(fm, cl, ml, appConfig) cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound) + } if err != nil { return nil, err } @@ -842,6 +1003,7 @@ func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfi confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, nil } @@ -854,6 +1016,8 @@ type RealtimeRoutingContext struct { Store router.DecisionStore SessionID string UserID string + // Failover resolves pipeline stages that name a failover chain. + Failover *failover.Manager } // buildRealtimeRoutingContext assembles the routing dependencies the @@ -875,6 +1039,7 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * Store: a.RouterDecisions(), SessionID: sessionID, UserID: userID, + Failover: a.FailoverManager(), } } @@ -882,7 +1047,20 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, evaluator *templates.Evaluator, routing *RealtimeRoutingContext) (Model, error) { xlog.Debug("Creating new model pipeline model", "pipeline", pipeline) + // A stage that names a failover chain is resolved on every call. Here it + // takes the chain's active target, so everything that inspects stage + // configs at session start (voice, reasoning, templates) sees a real model. + var fm *failover.Manager + if routing != nil { + fm = routing.Failover + } + sr := newStageRouter(fm, cl, ml, appConfig) + resolveStage := sr.resolveStage + cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgVAD, err = resolveStage(config.PipelineStageVAD, cfgVAD) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -894,6 +1072,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // TODO: Do we always need a transcription model? It can be disabled. Note that any-to-any instruction following models don't transcribe as such, so if transcription is required it is a separate process cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgSST, err = resolveStage(config.PipelineStageTranscription, cfgSST) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -926,6 +1107,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // Otherwise we want to return a wrapped model, which is a "virtual" model that re-uses other models to perform operations cfgLLM, err := cl.LoadResolvedModelConfig(pipeline.LLM, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgLLM, err = resolveStage(config.PipelineStageLLM, cfgLLM) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -937,10 +1121,17 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // Let the pipeline set the LLM's reasoning effort and force thinking off // (cfgLLM is a per-session copy). disable_thinking applies after the effort. - applyPipelineReasoning(cfgLLM, *pipeline) - applyPipelineThinking(cfgLLM, *pipeline) + pipelineCopy := *pipeline + tuneLLM := func(cfg *config.ModelConfig) { + applyPipelineReasoning(cfg, pipelineCopy) + applyPipelineThinking(cfg, pipelineCopy) + } + tuneLLM(cfgLLM) cfgTTS, err := cl.LoadResolvedModelConfig(pipeline.TTS, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgTTS, err = resolveStage(config.PipelineStageTTS, cfgTTS) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -951,6 +1142,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model } cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = resolveStage(config.PipelineStageSoundDetection, cfgSound) + } if err != nil { return nil, err } @@ -1000,6 +1194,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model modelLoader: ml, appConfig: appConfig, evaluator: evaluator, + + stageRouter: sr, + tuneLLM: tuneLLM, } if routing != nil { wm.routerDeps = routing.Deps diff --git a/core/http/endpoints/openai/realtime_semantic_vad_test.go b/core/http/endpoints/openai/realtime_semantic_vad_test.go index c36e13563..b1107c1f6 100644 --- a/core/http/endpoints/openai/realtime_semantic_vad_test.go +++ b/core/http/endpoints/openai/realtime_semantic_vad_test.go @@ -269,7 +269,7 @@ var _ = Describe("liveTurnState", func() { lts.drainEvents(1.0) var got []types.ConversationItemInputAudioTranscriptionDeltaEvent - for _, e := range ftr.events { + for _, e := range ftr.events() { if d, ok := e.(types.ConversationItemInputAudioTranscriptionDeltaEvent); ok { got = append(got, d) } @@ -335,7 +335,7 @@ var _ = Describe("commitUtteranceWithTranscript", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) var completed types.ConversationItemInputAudioTranscriptionCompletedEvent - for _, e := range tr.events { + for _, e := range tr.events() { if c, ok := e.(types.ConversationItemInputAudioTranscriptionCompletedEvent); ok { completed = c } @@ -407,7 +407,7 @@ var _ = Describe("emitPrecomputedTranscription", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionDelta)).To(Equal(2), "empty deltas skipped") Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) - for _, e := range tr.events { + for _, e := range tr.events() { switch ev := e.(type) { case types.ConversationItemInputAudioTranscriptionDeltaEvent: Expect(ev.ItemID).To(Equal("item42")) diff --git a/core/http/endpoints/openai/realtime_sound_detection_test.go b/core/http/endpoints/openai/realtime_sound_detection_test.go index e440e80c3..058c74076 100644 --- a/core/http/endpoints/openai/realtime_sound_detection_test.go +++ b/core/http/endpoints/openai/realtime_sound_detection_test.go @@ -38,7 +38,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.ItemID).To(Equal("item1")) Expect(ev.ContentIndex).To(Equal(0)) @@ -62,7 +62,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.Detections).To(BeEmpty()) }) diff --git a/core/http/endpoints/openai/realtime_stream_test.go b/core/http/endpoints/openai/realtime_stream_test.go index 2d5d7d7a1..5f160027e 100644 --- a/core/http/endpoints/openai/realtime_stream_test.go +++ b/core/http/endpoints/openai/realtime_stream_test.go @@ -250,8 +250,8 @@ var _ = Describe("triggerResponse", func() { // The single terminal carries the produced output item and the usage — // both empty in the legacy code. var done *types.ResponseDoneEvent - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { done = &d } } @@ -287,8 +287,8 @@ var _ = Describe("triggerResponse", func() { var created *types.ResponseCreatedEvent var done *types.ResponseDoneEvent - for i := range t.events { - switch e := t.events[i].(type) { + for i := range t.events() { + switch e := t.events()[i].(type) { case types.ResponseCreatedEvent: created = &e case types.ResponseDoneEvent: @@ -317,8 +317,8 @@ var _ = Describe("triggerResponse", func() { triggerResponse(context.Background(), session, &Conversation{}, t, nil) - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { Expect(d.Response.Metadata).To(BeEmpty()) } } diff --git a/core/http/endpoints/openai/realtime_voicegate_integration_test.go b/core/http/endpoints/openai/realtime_voicegate_integration_test.go index b0f7f0b49..4da774c76 100644 --- a/core/http/endpoints/openai/realtime_voicegate_integration_test.go +++ b/core/http/endpoints/openai/realtime_voicegate_integration_test.go @@ -67,7 +67,7 @@ func itSession(gate *voiceGate) (*Session, *fakeModel) { // hasSpeakerNotAuthorized reports whether a speaker_not_authorized error event // was emitted to the client. func hasSpeakerNotAuthorized(tr *fakeTransport) bool { - for _, e := range tr.events { + for _, e := range tr.events() { if ev, ok := e.(types.ErrorEvent); ok && ev.Error.Code == "speaker_not_authorized" { return true } diff --git a/core/http/endpoints/openai/sound_classification.go b/core/http/endpoints/openai/sound_classification.go index b7e23f1b1..878311d3a 100644 --- a/core/http/endpoints/openai/sound_classification.go +++ b/core/http/endpoints/openai/sound_classification.go @@ -48,7 +48,7 @@ func SoundClassificationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLo Threshold: float32(parseFormFloat(c, "threshold", 0)), } - file, err := c.FormFile("file") + file, err := uploadedFile(c, "file") if err != nil { return err } diff --git a/core/http/endpoints/openai/transcription.go b/core/http/endpoints/openai/transcription.go index 920094f84..6c99d7afc 100644 --- a/core/http/endpoints/openai/transcription.go +++ b/core/http/endpoints/openai/transcription.go @@ -2,9 +2,9 @@ package openai import ( "encoding/json" - "errors" "fmt" "io" + "mime/multipart" "net/http" "os" "path" @@ -56,6 +56,18 @@ func resolveTranscriptionTranslate(formTranslate string, configTranslate bool) b return configTranslate } +// uploadedFile reads a required multipart file field. Any failure here (no +// multipart boundary, malformed body, missing field) is caused by the request, +// so it maps to 400 instead of leaking the parser error as a 500. +func uploadedFile(c echo.Context, field string) (*multipart.FileHeader, error) { + file, err := c.FormFile(field) + if err != nil { + return nil, echo.NewHTTPError(http.StatusBadRequest, + fmt.Sprintf("missing or invalid %q file upload: %v", field, err)) + } + return file, nil +} + func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) echo.HandlerFunc { return func(c echo.Context) error { input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) @@ -104,8 +116,15 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app } } + // Reject an unknown format before the backend runs: the backend work + // would be wasted, and a failover chain would count the error + // against every target. + if !stream && !validTranscriptionResponseFormat(responseFormat) { + return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format") + } + // retrieve the file data from the request - file, err := c.FormFile("file") + file, err := uploadedFile(c, "file") if err != nil { return err } @@ -217,11 +236,21 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app } return c.JSON(http.StatusOK, trs) default: - return errors.New("invalid response_format") + return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format") } } } +func validTranscriptionResponseFormat(f schema.TranscriptionResponseFormatType) bool { + switch f { + case "", schema.TranscriptionResponseFormatLrc, schema.TranscriptionResponseFormatText, + schema.TranscriptionResponseFormatSrt, schema.TranscriptionResponseFormatVtt, + schema.TranscriptionResponseFormatJson, schema.TranscriptionResponseFormatJsonVerbose: + return true + } + return false +} + // streamTranscription emits OpenAI-format SSE events for a transcription // request: one `transcript.text.delta` per backend chunk, a final // `transcript.text.done` with the assembled text, and `[DONE]`. Backends that diff --git a/core/http/endpoints/openai/types/failover.go b/core/http/endpoints/openai/types/failover.go new file mode 100644 index 000000000..69c048207 --- /dev/null +++ b/core/http/endpoints/openai/types/failover.go @@ -0,0 +1,46 @@ +package types + +import "encoding/json" + +// ModelFailoverEvent is a LocalAI extension server event +// (localai.model.failover). It tells a client which target serves a +// pipeline stage that names a failover chain: once per chain stage at +// session start (reason "initial"), then on every switch of that chain. +type ModelFailoverEvent struct { + ServerEventBase + + // The failover chain the stage names. + Chain string `json:"chain"` + + // The pipeline stage: vad, transcription, llm, tts or sound_detection. + Stage string `json:"stage"` + + // The target that served the stage before the switch; "" at session start. + From string `json:"from"` + + // The target that serves the stage now. + To string `json:"to"` + + // The chain state: primary, fallback or degraded. + State string `json:"state"` + + // Why the chain switched, or "initial" at session start. + Reason string `json:"reason"` +} + +func (m ModelFailoverEvent) ServerEventType() ServerEventType { + return ServerEventTypeModelFailover +} + +func (m ModelFailoverEvent) MarshalJSON() ([]byte, error) { + type typeAlias ModelFailoverEvent + type typeWrapper struct { + typeAlias + Type ServerEventType `json:"type"` + } + shadow := typeWrapper{ + typeAlias: typeAlias(m), + Type: m.ServerEventType(), + } + return json.Marshal(shadow) +} diff --git a/core/http/endpoints/openai/types/server_events.go b/core/http/endpoints/openai/types/server_events.go index b847a35a7..114a7065a 100644 --- a/core/http/endpoints/openai/types/server_events.go +++ b/core/http/endpoints/openai/types/server_events.go @@ -27,7 +27,11 @@ const ( // ServerEventTypeClassifierResult is a LocalAI extension: it carries the // classifier-mode score distribution and decision for a response. OpenAI // clients ignore it. - ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + // ServerEventTypeModelFailover is a LocalAI extension: it names the target + // that serves a pipeline stage backed by a failover chain, at session start + // and on every chain switch. OpenAI clients ignore it. + ServerEventTypeModelFailover ServerEventType = "localai.model.failover" ServerEventTypeInputAudioBufferCommitted ServerEventType = "input_audio_buffer.committed" ServerEventTypeInputAudioBufferCleared ServerEventType = "input_audio_buffer.cleared" ServerEventTypeInputAudioBufferSpeechStarted ServerEventType = "input_audio_buffer.speech_started" diff --git a/core/http/middleware/admission.go b/core/http/middleware/admission.go index d6134b026..f69dc3206 100644 --- a/core/http/middleware/admission.go +++ b/core/http/middleware/admission.go @@ -38,6 +38,7 @@ func AdmissionControl(limiter *admission.Limiter, events pii.EventStore) echo.Mi if !ok { retryAfter := admission.RetryAfter(cfg.Limits.RetryAfterSeconds) recordAdmissionRejection(events, cfg.Name, retryAfter) + c.Set(ContextKeyAdmissionRejected, true) c.Response().Header().Set("Retry-After", strconv.Itoa(int(retryAfter.Seconds()))) return c.JSON(http.StatusTooManyRequests, map[string]any{ "error": map[string]any{ diff --git a/core/http/middleware/context_keys.go b/core/http/middleware/context_keys.go index d1983c882..1e1f3886c 100644 --- a/core/http/middleware/context_keys.go +++ b/core/http/middleware/context_keys.go @@ -47,4 +47,13 @@ const ( // router nor the body-parse path has produced one. Distinct from // ContextKeyServedModel, which is the router's resolved choice. ContextKeyResponseModel = "routing.response_model" + + // ContextKeyFailoverAttempt holds the *failoverState of a request whose + // model is a failover chain. + ContextKeyFailoverAttempt = "failover.attempt" + + // ContextKeyAdmissionRejected is set to true by AdmissionControl when it + // turns a request away because the model is at capacity. The failover + // retry sends such a request to the next target without tripping this one. + ContextKeyAdmissionRejected = "admission.rejected" ) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go new file mode 100644 index 000000000..e2df6285f --- /dev/null +++ b/core/http/middleware/failover.go @@ -0,0 +1,335 @@ +package middleware + +import ( + "bufio" + "bytes" + "fmt" + "io" + "net" + "net/http" + "slices" + "strings" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" +) + +const ( + HeaderServedModel = "X-LocalAI-Served-Model" + HeaderFailover = "X-LocalAI-Failover" +) + +// MaxFailoverReplayBody caps the request body kept for a retry. A larger body +// is still served, by one target only. +const MaxFailoverReplayBody = 32 << 20 + +type failoverState struct { + attempt *failover.Attempt +} + +// SetFailoverManager enables failover chains. Without it, a request for a +// chain fails with 503. +func (re *RequestExtractor) SetFailoverManager(m *failover.Manager) { re.failover = m } + +// resolveFailover returns the config of the target that should serve this +// attempt. The first attempt plans the chain; retries reuse the plan. +func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, chain *config.ModelConfig) (*config.ModelConfig, error) { + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + if st == nil || st.attempt.Chain() != chain.Name { + if re.failover == nil { + return nil, fmt.Errorf("model %q is a failover chain, but failover is not running", chain.Name) + } + att, err := re.failover.Plan(chain.Name) + if err != nil { + return nil, err + } + st = &failoverState{attempt: att} + c.Set(ContextKeyFailoverAttempt, st) + } + for { + cfg, err := re.loadFailoverTarget(st.attempt.Target()) + // A target that became a chain after its chain was saved has no + // backend of its own; like a disabled one it is skipped, not tripped. + if err == nil && (cfg.IsDisabled() || cfg.IsFailover()) { + // Disabled on purpose, not broken: move on without a trip. + if st.attempt.Skip() { + continue + } + c.Set(ContextKeyFailoverAttempt, nil) + return nil, fmt.Errorf("failover chain %q: target %q is disabled or is itself a chain", chain.Name, cfg.Name) + } + if err == nil { + failover.PrepareTarget(cfg) // cfg is a copy + c.Set(ContextKeyRequestedModel, requested) + c.Set(ContextKeyServedModel, cfg.Name) + setFailoverHeaders(c.Response().Header(), st.attempt) + return cfg, nil + } + if !st.attempt.Fail(err) { + // Clear the state so the retry wrapper sends this 503 as is. + c.Set(ContextKeyFailoverAttempt, nil) + return nil, err + } + } +} + +func (re *RequestExtractor) loadFailoverTarget(name string) (*config.ModelConfig, error) { + cfg, err := re.modelConfigLoader.LoadModelConfigFileByNameDefaultOptions(name, re.applicationConfig) + if err != nil { + return nil, err + } + resolved, _, err := re.modelConfigLoader.ResolveAlias(cfg) + return resolved, err +} + +func setFailoverHeaders(h http.Header, att *failover.Attempt) { + h.Set(HeaderServedModel, att.Target()) + switch { + case att.Degraded(): + h.Set(HeaderFailover, "degraded") + case att.Target() != att.Primary(): + h.Set(HeaderFailover, "fallback") + default: + h.Del(HeaderFailover) + } +} + +// failoverRetry runs h again on the next target while the response is not +// committed. h is SetModelAndConfig's body plus the rest of the chain, so +// every attempt binds the request again from the replayed body. +func (re *RequestExtractor) failoverRetry(h echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + // Without chains there is nothing to retry, so plain installations + // pay neither for the body copy nor for the writer. + if !re.failover.HasChains() { + return h(c) + } + req := c.Request() + src := req.Body + if src == nil { + src = http.NoBody + } + rec := &replayBody{src: src, limit: MaxFailoverReplayBody} + req.Body = rec + // A default-model middleware may have parsed a multipart form before + // this point, draining the body. That form stays valid for every + // attempt; a form parsed during an attempt is dropped and parsed again + // from the replayed body. + entryMultipart, entryForm, entryPostForm := req.MultipartForm, req.Form, req.PostForm + resp := c.Response() + orig := resp.Writer + baseHeader := resp.Header().Clone() + defer func() { resp.Writer = orig }() + active := func() bool { + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + return st != nil + } + tracing := re.applicationConfig != nil && re.applicationConfig.EnableTracing + for { + w := &failoverWriter{ResponseWriter: orig, active: active} + resp.Writer = w + c.Set(ContextKeyAdmissionRejected, nil) + err := h(c) + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + if st == nil { + w.release() + return err + } + att := st.attempt + status := w.held + if err == nil && status == 0 { + // A 4xx says nothing about the target's health. + if w.status < http.StatusBadRequest { + att.Succeed() + } + return nil + } + rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool) + // A handler that wrote its 501 itself instead of returning it + // reports the same gap. + gap := (failover.IsCapabilityGap(err) || status == http.StatusNotImplemented) && !w.committed + if (rejected && !w.committed) || gap { + // The target is at capacity, or cannot serve this kind of + // request at all: spill to the next target without counting + // a failure. Neither says anything about the target's health. + if !rec.replayable() || !att.Skip() { + w.release() + return err + } + } else { + cause := attemptError(err, status, w.body.Bytes()) + retryable := req.Context().Err() == nil && failover.IsRetryable(err, status) + if !retryable || w.committed || !rec.replayable() { + if retryable { + att.Report(cause) + } + w.release() + return err + } + failover.RecordAttemptTrace(tracing, att.Chain(), att.Target(), cause) + if !att.Fail(cause) { + w.release() + return err + } + } + // Temp files of a form parsed during this attempt would otherwise + // outlive the request: the server only cleans up the last form. + if mf := c.Request().MultipartForm; mf != nil && mf != entryMultipart { + _ = mf.RemoveAll() + } + req.Body = rec.replay() + req.MultipartForm, req.Form, req.PostForm = entryMultipart, entryForm, entryPostForm + // Middleware after this one may have replaced the request; the + // next attempt starts again from the request as it arrived here. + c.SetRequest(req) + resetResponse(resp, baseHeader) + } + } +} + +// stopFailoverRecording releases the recorded body of a request whose model +// turned out not to be a chain: it will never be replayed. +func stopFailoverRecording(c echo.Context) { + if rb, ok := c.Request().Body.(*replayBody); ok { + rb.stop() + } +} + +func attemptError(err error, status int, body []byte) error { + if err != nil { + return err + } + msg := strings.TrimSpace(string(body)) + if len(msg) > 200 { + msg = msg[:200] + } + return fmt.Errorf("HTTP %d: %s", status, msg) +} + +func resetResponse(resp *echo.Response, base http.Header) { + h := resp.Header() + for k := range h { + delete(h, k) + } + for k, v := range base { + h[k] = slices.Clone(v) + } + resp.Committed = false + resp.Status = http.StatusOK + resp.Size = 0 +} + +// failoverWriter holds back an error response (status >= 500) of a chain +// request until the handler returns, so the retry can drop it. +type failoverWriter struct { + http.ResponseWriter + active func() bool + held int + body bytes.Buffer + committed bool + // status is the code sent to the client, 0 until one is sent. + status int +} + +func (w *failoverWriter) WriteHeader(code int) { + if w.held != 0 { + return + } + if !w.committed && code >= 500 && w.active() { + w.held = code + return + } + w.committed = true + w.status = code + w.ResponseWriter.WriteHeader(code) +} + +func (w *failoverWriter) Write(b []byte) (int, error) { + if w.held != 0 { + return w.body.Write(b) + } + w.committed = true + return w.ResponseWriter.Write(b) +} + +// FlushError and Hijack go through a ResponseController so the capabilities of +// wrapped writers further down stay reachable, as they were before this writer +// was inserted. +func (w *failoverWriter) FlushError() error { + w.release() + w.committed = true + return http.NewResponseController(w.ResponseWriter).Flush() +} + +func (w *failoverWriter) Flush() { _ = w.FlushError() } + +func (w *failoverWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + w.committed = true + return http.NewResponseController(w.ResponseWriter).Hijack() +} + +func (w *failoverWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } + +// release sends a held error response to the client. +func (w *failoverWriter) release() { + if w.held == 0 { + return + } + code := w.held + w.held = 0 + w.committed = true + w.ResponseWriter.WriteHeader(code) + _, _ = w.ResponseWriter.Write(w.body.Bytes()) + w.body.Reset() +} + +// replayBody records what the handler reads, up to limit, so the body can be +// sent again to the next target. Decoders often stop at the end of the value +// without reading to EOF, so a replay is the recorded bytes followed by +// whatever the previous attempt left unread. +type replayBody struct { + src io.ReadCloser + buf bytes.Buffer + limit int + overflow bool + stopped bool +} + +func (r *replayBody) Read(p []byte) (int, error) { + n, err := r.src.Read(p) + if n > 0 && !r.overflow && !r.stopped { + if r.buf.Len()+n > r.limit { + r.overflow = true + // A new buffer, not Reset: Reset keeps the memory. + r.buf = bytes.Buffer{} + } else { + r.buf.Write(p[:n]) + } + } + return n, err +} + +func (r *replayBody) Close() error { return r.src.Close() } + +// replayable reports whether everything read so far was kept. +func (r *replayBody) replayable() bool { return !r.overflow && !r.stopped } + +// stop ends recording and frees what was kept. +func (r *replayBody) stop() { + r.stopped = true + r.buf = bytes.Buffer{} +} + +// replay rewinds to the start of the body and keeps recording, so a third +// attempt can replay too. +func (r *replayBody) replay() io.ReadCloser { + data := bytes.Clone(r.buf.Bytes()) + rest := r.src + r.src = struct { + io.Reader + io.Closer + }{io.MultiReader(bytes.NewReader(data), rest), rest} + r.buf.Reset() + return r +} diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go new file mode 100644 index 000000000..17625621f --- /dev/null +++ b/core/http/middleware/failover_test.go @@ -0,0 +1,381 @@ +package middleware + +import ( + "bytes" + "context" + "errors" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + "testing/iotest" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/core/services/routing/admission" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +var _ = Describe("failover chains in the request pipeline", func() { + var ( + app *echo.Echo + fm *failover.Manager + re *RequestExtractor + limiter *admission.Limiter + mu sync.Mutex + calls []string + behavior map[string]func(c echo.Context) error + ) + + served := func(c echo.Context) error { + cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + return c.JSON(http.StatusOK, map[string]string{"served": cfg.Name}) + } + + handler := func(c echo.Context) error { + cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + mu.Lock() + calls = append(calls, cfg.Name) + b := behavior[cfg.Name] + mu.Unlock() + if b == nil { + return served(c) + } + return b(c) + } + + post := func(path, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + return rec + } + chat := func(model string) *httptest.ResponseRecorder { + return post("/v1/chat/completions", `{"model":"`+model+`","messages":[{"role":"user","content":"hi"}]}`) + } + + BeforeEach(func() { + calls = nil + behavior = map[string]func(c echo.Context) error{} + dir := GinkgoT().TempDir() + write := func(name, body string) { + Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed()) + } + write("a", "name: a\nbackend: fake-a\n") + write("b", "name: b\nbackend: fake-b\n") + write("plain", "name: plain\nbackend: fake-p\n") + write("chain", "name: chain\nfailover:\n targets:\n - model: a\n - model: b\n") + write("capped", "name: capped\nbackend: fake-c\nlimits:\n max_concurrent: 1\n") + write("off", "name: off\nbackend: fake-o\ndisabled: true\n") + write("chain-capped", "name: chain-capped\nfailover:\n targets:\n - model: capped\n - model: b\n") + write("chain-off", "name: chain-off\nfailover:\n targets:\n - model: off\n - model: b\n") + write("remote", "name: remote\nbackend: cloud-proxy\nproxy:\n mode: passthrough\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n") + write("remote-mapped", "name: remote-mapped\nbackend: cloud-proxy\nproxy:\n mode: translate\n provider: openai\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n upstream_model: big-llm\n") + write("chain-remote", "name: chain-remote\nfailover:\n targets:\n - model: remote\n - model: remote-mapped\n") + + ss := &system.SystemState{Model: system.Model{ModelsPath: dir}} + appConfig := config.NewApplicationConfig() + appConfig.SystemState = ss + mcl := config.NewModelConfigLoader(dir) + Expect(mcl.LoadModelConfigsFromPath(dir)).To(Succeed()) + re = NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig) + fm = failover.New(mcl) + // The scheduler's first tick syncs in the application; HasChains + // answers from the last sync. + fm.Sync() + re.SetFailoverManager(fm) + + app = echo.New() + // echo's default handler hides internal error messages; the specs + // below check which target's error reached the client. + app.HTTPErrorHandler = func(err error, c echo.Context) { + code := http.StatusInternalServerError + var he *echo.HTTPError + if errors.As(err, &he) { + code = he.Code + } + _ = c.JSON(code, map[string]string{"error": err.Error()}) + } + app.POST("/v1/chat/completions", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + app.POST("/v1/audio/transcriptions", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + limiter = admission.New() + app.POST("/v1/chat/admitted", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }), + AdmissionControl(limiter, nil)) + // The real transcription route resolves a default model first, which + // parses the multipart form before SetModelAndConfig runs. + app.POST("/v1/audio/transcriptions-default", handler, + re.BuildFilteredFirstAvailableDefaultModel(config.NoFilterFn), + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + }) + + It("serves from the next target when the first fails before responding", func() { + behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: connection refused") } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(rec.Header().Get(HeaderServedModel)).To(Equal("b")) + Expect(rec.Header().Get(HeaderFailover)).To(Equal("fallback")) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + }) + + It("sends a remote target its own upstream model, not the chain name", func() { + upstream := map[string]string{} + record := func(c echo.Context) error { + cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + mu.Lock() + upstream[cfg.Name] = cfg.Proxy.UpstreamModel + mu.Unlock() + return nil + } + behavior["remote"] = func(c echo.Context) error { + _ = record(c) + return errors.New("dial tcp: connection refused") + } + behavior["remote-mapped"] = func(c echo.Context) error { _ = record(c); return served(c) } + Expect(chat("chain-remote").Code).To(Equal(http.StatusOK)) + // The probe checks the same name, so a request cannot fail on a model + // the liveness probe just found. + Expect(upstream).To(Equal(map[string]string{"remote": "remote", "remote-mapped": "big-llm"})) + for _, name := range []string{"remote", "remote-mapped"} { + cfg, ok := re.modelConfigLoader.GetModelConfig(name) + Expect(ok).To(BeTrue()) + Expect(upstream[name]).To(Equal(failover.UpstreamModel(cfg))) + } + stored, _ := re.modelConfigLoader.GetModelConfig("remote") + Expect(stored.Proxy.UpstreamModel).To(BeEmpty(), "the shared config must not change") + }) + + It("serves the primary without the failover header", func() { + rec := chat("chain") + Expect(rec.Header().Get(HeaderServedModel)).To(Equal("a")) + Expect(rec.Header().Get(HeaderFailover)).To(BeEmpty()) + }) + + It("drops a buffered 5xx response and retries", func() { + behavior["a"] = func(c echo.Context) error { + return c.JSON(http.StatusServiceUnavailable, map[string]string{"error": "no healthy nodes"}) + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).ToNot(ContainSubstring("no healthy nodes")) + }) + + It("does not retry after streaming started, and trips the target", func() { + behavior["a"] = func(c echo.Context) error { + c.Response().Header().Set("Content-Type", "text/event-stream") + _, _ = c.Response().Write([]byte("data: x\n\n")) + c.Response().Flush() + return errors.New("connection reset by peer") + } + rec := chat("chain") + Expect(rec.Body.String()).To(HavePrefix("data: x")) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) + }) + + It("does not retry or trip on 4xx", func() { + behavior["a"] = func(echo.Context) error { return echo.NewHTTPError(http.StatusBadRequest, "bad") } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusBadRequest)) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("does not retry when the client cancelled", func() { + ctx, cancel := context.WithCancel(context.Background()) + behavior["a"] = func(echo.Context) error { cancel(); return context.Canceled } + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", + strings.NewReader(`{"model":"chain","messages":[{"role":"user","content":"hi"}]}`)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + app.ServeHTTP(httptest.NewRecorder(), req) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("gives each attempt a fresh request", func() { + behavior["a"] = func(c echo.Context) error { + in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + in.Messages = nil + return errors.New("dial tcp: connection refused") + } + behavior["b"] = func(c echo.Context) error { + in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + Expect(in.Messages).To(HaveLen(1)) + return served(c) + } + Expect(chat("chain").Code).To(Equal(http.StatusOK)) + }) + + It("degraded: tries every target in priority order and returns the last error", func() { + behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: a down") } + behavior["b"] = func(echo.Context) error { return errors.New("dial tcp: b down") } + chat("chain") // trips both + calls = nil + rec := chat("chain") + Expect(calls).To(Equal([]string{"a", "b"})) + Expect(rec.Code).To(Equal(http.StatusInternalServerError)) + Expect(rec.Body.String()).To(ContainSubstring("b down")) + Expect(rec.Header().Get(HeaderFailover)).To(Equal("degraded")) + }) + + DescribeTable("replays a multipart body for the next target", func(path string) { + var body bytes.Buffer + mw := multipart.NewWriter(&body) + _ = mw.WriteField("model", "chain") + fw, _ := mw.CreateFormFile("file", "a.wav") + _, _ = fw.Write(bytes.Repeat([]byte{7}, 4096)) + Expect(mw.Close()).To(Succeed()) + size := func(c echo.Context) int64 { + fh, err := c.FormFile("file") + Expect(err).ToNot(HaveOccurred()) + f, _ := fh.Open() + n, _ := io.Copy(io.Discard, f) + return n + } + behavior["a"] = func(c echo.Context) error { size(c); return errors.New("dial tcp: refused") } + behavior["b"] = func(c echo.Context) error { + Expect(size(c)).To(Equal(int64(4096))) + return served(c) + } + req := httptest.NewRequest(http.MethodPost, path, &body) + req.Header.Set("Content-Type", mw.FormDataContentType()) + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(calls).To(Equal([]string{"a", "b"})) + }, + Entry("read by the handler", "/v1/audio/transcriptions"), + Entry("parsed before SetModelAndConfig", "/v1/audio/transcriptions-default"), + ) + + It("spills an admission rejection to the next target without tripping it", func() { + release, ok := limiter.Acquire("capped", 1) + Expect(ok).To(BeTrue()) + defer release() + rec := post("/v1/chat/admitted", `{"model":"chain-capped","messages":[{"role":"user","content":"hi"}]}`) + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(rec.Header().Get("Retry-After")).To(BeEmpty()) + st, _ := fm.ChainStatus("chain-capped") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + Expect(st.Active).To(Equal("capped")) + }) + + It("spills a capability gap (gRPC Unimplemented) to the next target without tripping it", func() { + behavior["a"] = func(echo.Context) error { + return grpcstatus.Error(codes.Unimplemented, "localai-proxy: Rerank has no upstream counterpart") + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("spills a written 501 to the next target without tripping it", func() { + behavior["a"] = func(c echo.Context) error { + return c.JSON(http.StatusNotImplemented, map[string]string{"error": "not supported"}) + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("fails a rate-limited target (gRPC ResourceExhausted) over to the next target and trips it", func() { + behavior["a"] = func(echo.Context) error { + return grpcstatus.Error(codes.ResourceExhausted, "localai-proxy: upstream /v1/chat/completions returned 429: slow down") + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) + }) + + It("skips a disabled target without tripping it", func() { + rec := chat("chain-off") + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"b"})) + st, _ := fm.ChainStatus("chain-off") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("counts a 4xx response as neither success nor failure", func() { + behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: a down") } + behavior["b"] = func(echo.Context) error { return errors.New("dial tcp: b down") } + chat("chain") // trips both + behavior["a"] = func(c echo.Context) error { + return c.JSON(http.StatusBadRequest, map[string]string{"error": "bad"}) + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusBadRequest)) + st, _ := fm.ChainStatus("chain") + // a is a cold local target: a success would have recovered it. + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) + }) + + It("stops recording the body once the model is known not to be a chain", func() { + behavior["plain"] = func(c echo.Context) error { + rb, ok := c.Request().Body.(*replayBody) + Expect(ok).To(BeTrue()) + Expect(rb.replayable()).To(BeFalse()) + Expect(rb.buf.Cap()).To(BeZero()) + return served(c) + } + Expect(chat("plain").Code).To(Equal(http.StatusOK)) + }) + + It("does not record the body without a failover manager", func() { + re.SetFailoverManager(nil) + behavior["plain"] = func(c echo.Context) error { + _, ok := c.Request().Body.(*replayBody) + Expect(ok).To(BeFalse()) + return served(c) + } + Expect(chat("plain").Code).To(Equal(http.StatusOK)) + }) + + It("releases the recorded bytes on overflow", func() { + // One byte per read, so some bytes are recorded before the limit is hit. + rb := &replayBody{src: io.NopCloser(iotest.OneByteReader(bytes.NewReader(make([]byte, 64)))), limit: 16} + _, _ = io.ReadAll(rb) + Expect(rb.replayable()).To(BeFalse()) + Expect(rb.buf.Cap()).To(BeZero()) + }) + + It("leaves plain models untouched", func() { + behavior["plain"] = func(c echo.Context) error { + return c.JSON(http.StatusInternalServerError, map[string]string{"error": "boom"}) + } + rec := chat("plain") + Expect(rec.Code).To(Equal(http.StatusInternalServerError)) + Expect(rec.Body.String()).To(ContainSubstring("boom")) + Expect(rec.Header().Get(HeaderServedModel)).To(BeEmpty()) + }) +}) diff --git a/core/http/middleware/request.go b/core/http/middleware/request.go index 080a0b73c..7bc702e20 100644 --- a/core/http/middleware/request.go +++ b/core/http/middleware/request.go @@ -12,6 +12,7 @@ import ( "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/templates" "github.com/mudler/LocalAI/pkg/distributedhdr" @@ -29,6 +30,7 @@ type RequestExtractor struct { modelConfigLoader *config.ModelConfigLoader modelLoader *model.ModelLoader applicationConfig *config.ApplicationConfig + failover *failover.Manager } func NewRequestExtractor(modelConfigLoader *config.ModelConfigLoader, modelLoader *model.ModelLoader, applicationConfig *config.ApplicationConfig) *RequestExtractor { @@ -122,7 +124,7 @@ func (re *RequestExtractor) BuildFilteredFirstAvailableDefaultModel(filterFn con // Otherwise, it's in its own method below for now func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIRequest) echo.MiddlewareFunc { return func(next echo.HandlerFunc) echo.HandlerFunc { - return func(c echo.Context) error { + return re.failoverRetry(func(c echo.Context) error { input := initializer() if input == nil { return echo.NewHTTPError(http.StatusBadRequest, "unable to initialize body") @@ -194,6 +196,24 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR cfg = resolved } + // A failover chain resolves to one of its targets, like an alias. + // failoverRetry re-runs this middleware for the next target. + if cfg != nil && cfg.IsFailover() { + resolved, fErr := re.resolveFailover(c, modelName, cfg) + if fErr != nil { + return c.JSON(http.StatusServiceUnavailable, schema.ErrorResponse{ + Error: &schema.APIError{ + Message: fErr.Error(), + Code: http.StatusServiceUnavailable, + Type: "failover_unavailable", + }, + }) + } + cfg = resolved + } else { + stopFailoverRecording(c) + } + // Check if the model is disabled if cfg != nil && cfg.IsDisabled() { return c.JSON(http.StatusForbidden, schema.ErrorResponse{ @@ -209,7 +229,7 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg) return next(c) - } + }) } } diff --git a/core/http/react-ui/e2e/failover-editor.spec.js b/core/http/react-ui/e2e/failover-editor.spec.js new file mode 100644 index 000000000..da66227eb --- /dev/null +++ b/core/http/react-ui/e2e/failover-editor.spec.js @@ -0,0 +1,175 @@ +import { test, expect } from './coverage-fixtures.js' + +// Failover Chain template + FailoverTargetsEditor regression tests. +// +// A failover chain is a model config whose `failover.targets` field lists, +// in order, the downstream models that answer for one name — the first +// healthy target serves each request, the rest take over when it fails. +// This covers: +// - the create-flow template gallery exposes a "Failover Chain" card that +// seeds a minimal name + two empty targets +// - the dedicated FailoverTargetsEditor renders {model, warm} rows with +// add/remove/move controls +// - inline warnings for a single-target chain, a duplicated target model, +// and a target that names the chain itself +// - saving an edited chain sends the updated failover.targets array in +// the PATCH body + +const FAILOVER_METADATA = { + sections: [ + { id: 'general', label: 'General', icon: 'settings', order: 0 }, + { id: 'failover', label: 'Failover', icon: 'shuffle', order: 90 }, + ], + fields: [ + { + path: 'name', yaml_key: 'name', go_type: 'string', ui_type: 'string', + section: 'general', label: 'Model Name', component: 'input', order: 0, + }, + { + path: 'failover.targets', yaml_key: 'targets', go_type: '[]FailoverTarget', ui_type: 'object', + section: 'failover', label: 'Failover targets', component: 'failover-targets', + description: 'Ordered list of models that serve this chain.', + order: 1, + }, + ], +} + +async function mockCommon(page) { + await page.route('**/api/auth/status', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ authEnabled: false, staticApiKeyRequired: false, providers: [] }) })) + await page.route('**/api/models/config-metadata*', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(FAILOVER_METADATA) })) + await page.route('**/api/models/config-metadata/autocomplete/**', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ values: [] }) })) + + page.on('pageerror', (err) => { + throw new Error(`uncaught page error: ${err.message}`) + }) +} + +test.describe('Failover Chain template - create flow', () => { + test.beforeEach(async ({ page }) => { + await mockCommon(page) + }) + + test('template gallery exposes the Failover Chain card', async ({ page }) => { + await page.goto('/app/model-editor') + await expect(page.getByRole('button', { name: /Failover Chain/i })).toBeVisible({ timeout: 10_000 }) + }) + + test('failover template loads the editor with two target rows', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + await expect(page.getByText(/Unexpected Application Error/i)).toHaveCount(0) + await expect(page.locator('h1.page-title')).toBeVisible({ timeout: 10_000 }) + await expect(page.getByText('Failover targets').first()).toBeVisible() + + // Two empty {model: ''} rows seeded by the template. + await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(2) + await expect(page.getByRole('button', { name: /Add target/i }).first()).toBeVisible() + }) + + test('Add target adds a row', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + await page.getByRole('button', { name: /Add target/i }).first().click() + await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(3) + }) + + test('removing a row removes it', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + await page.locator('button[title="Remove target"]').first().click() + await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(1) + }) + + test('move down/up reorders the target rows', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + const rows = page.locator('input[placeholder="target model..."]') + await rows.nth(0).fill('model-a') + await rows.nth(1).fill('model-b') + + // Move the first row down — it should now be the second row's value. + await page.locator('button[title="Move down"]').first().click() + await expect(rows.nth(0)).toHaveValue('model-b') + await expect(rows.nth(1)).toHaveValue('model-a') + + // Move it back up with the second row's "Move up" control. + await page.locator('button[title="Move up (tried earlier)"]').nth(1).click() + await expect(rows.nth(0)).toHaveValue('model-a') + await expect(rows.nth(1)).toHaveValue('model-b') + }) + + test('a single remaining target shows a too-few warning', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + await page.locator('button[title="Remove target"]').first().click() + await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(1) + await expect(page.getByText(/nothing to fail over to/i)).toBeVisible() + }) + + test('duplicate target models flag both rows', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + const rows = page.locator('input[placeholder="target model..."]') + await rows.nth(0).fill('same-model') + await rows.nth(1).fill('same-model') + await expect(page.getByText(/Duplicate target/i)).toHaveCount(2) + }) + + test('a target naming the chain itself is flagged', async ({ page }) => { + await page.goto('/app/model-editor?template=failover') + // Create mode renders the model name through a dedicated input (not the + // generic field renderer), whose placeholder comes from the modelEditor + // i18n namespace rather than the field registry. + await page.locator('input[placeholder="my-model-name"]').fill('my-chain') + const rows = page.locator('input[placeholder="target model..."]') + await rows.nth(0).fill('my-chain') + await expect(page.getByText(/chain's own name/i)).toBeVisible() + }) +}) + +test.describe('Failover Chain - saving an edited chain', () => { + const MOCK_YAML = 'name: my-chain\nfailover:\n targets:\n - model: model-a\n warm: true\n - model: model-b\n warm: false\n' + + test.beforeEach(async ({ page }) => { + await mockCommon(page) + await page.route('**/api/models/edit/my-chain', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: MOCK_YAML, name: 'my-chain' }) })) + // ModelFailoverStatus mounts unconditionally in edit mode; keep its + // fetch + SSE subscription harmless for a chain the test doesn't care + // about the live status of. + await page.route('**/api/failover', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [] }) })) + await page.route('**/api/failover/events', (route) => + route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: '' })) + }) + + test('saving sends the updated failover.targets array in the PATCH body', async ({ page }) => { + let patchBody = null + await page.route('**/api/models/config-json/my-chain', (route) => { + if (route.request().method() === 'PATCH') { + patchBody = route.request().postDataJSON() + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ success: true, message: "Model 'my-chain' updated successfully" }) }) + } else { + route.fulfill({ contentType: 'application/json', body: '{}' }) + } + }) + + await page.goto('/app/model-editor/my-chain') + await expect(page.locator('h1', { hasText: 'Model Editor' })).toBeVisible({ timeout: 10_000 }) + + // Existing targets loaded from YAML. + const rows = page.locator('input[placeholder="target model..."]') + await expect(rows).toHaveCount(2) + + // Add a third target, then save. + await page.getByRole('button', { name: /Add target/i }).first().click() + await rows.nth(2).fill('model-c') + + await page.locator('button', { hasText: 'Save Changes' }).click() + await expect(page.locator('text=Configuration saved')).toBeVisible({ timeout: 5_000 }) + + expect(patchBody).toBeTruthy() + expect(patchBody.failover.targets).toEqual([ + { model: 'model-a', warm: true }, + { model: 'model-b', warm: false }, + { model: 'model-c', warm: false }, + ]) + }) +}) diff --git a/core/http/react-ui/e2e/failover-health.spec.js b/core/http/react-ui/e2e/failover-health.spec.js new file mode 100644 index 000000000..bc2b3070d --- /dev/null +++ b/core/http/react-ui/e2e/failover-health.spec.js @@ -0,0 +1,124 @@ +import { test, expect } from './coverage-fixtures.js' + +// Live failover chain health in the Model Editor. A model that is a failover +// chain shows a strip under the editor header with the chain state, the +// active target, and a per-target health table fed by GET /api/failover plus +// the /api/failover/events SSE stream. Pin controls are admin-only. + +const MOCK_METADATA = { + sections: [{ id: 'general', label: 'General', icon: 'settings', order: 0 }], + fields: [ + { path: 'name', yaml_key: 'name', go_type: 'string', ui_type: 'string', section: 'general', label: 'Model Name', description: 'id', component: 'input', order: 0 }, + ], +} +const MOCK_YAML = 'name: chain\nfailover:\n targets: [a, b]\n' + +const CHAIN = { + name: 'chain', + state: 'primary', + active: 'a', + active_since: '2026-09-26T09:00:00Z', + pinned: null, + targets: [ + { model: 'a', kind: 'local', warm: true, state: 'healthy', consecutive_ok: 5, last_probe: '2026-09-26T09:59:00Z' }, + { model: 'b', kind: 'remote', warm: false, state: 'healthy', consecutive_ok: 3, last_error: 'dial tcp: connection refused while probing upstream' }, + ], +} + +// The stream replays a snapshot (primary on a) then a switch to b. The body +// ends after the two frames; EventSource reconnects and replays them, so the +// settled state stays fallback on b. +const SSE_BODY = + `event: snapshot\ndata: ${JSON.stringify({ chains: [CHAIN] })}\n\n` + + 'event: chain.switched\ndata: {"type":"chain.switched","chain":"chain","from":"a","to":"b","state":"fallback","reason":"trip","at":"2026-09-26T10:00:00Z"}\n\n' + +async function mockEditor(page, authStatus) { + await page.route('**/api/auth/status', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(authStatus) })) + await page.route('**/api/models/config-metadata*', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK_METADATA) })) + await page.route('**/api/models/edit/**', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: MOCK_YAML, name: 'chain' }) })) + await page.route('**/api/models/config-json/**', (route) => + route.fulfill({ contentType: 'application/json', body: '{}' })) + await page.route('**/api/failover', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [CHAIN] }) })) + await page.route('**/api/failover/chain', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(CHAIN) })) + await page.route('**/api/failover/events', (route) => + route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: SSE_BODY })) +} + +const NO_AUTH = { authEnabled: false, staticApiKeyRequired: false, providers: [] } +const NON_ADMIN = { + authEnabled: true, + staticApiKeyRequired: false, + providers: ['local'], + user: { id: 'user-uuid', name: 'User', role: 'user', provider: 'local' }, +} + +test.describe('Model Editor — failover chain health', () => { + test('shows the chain strip and follows chain.switched to fallback', async ({ page }) => { + await mockEditor(page, NO_AUTH) + await page.goto('/app/model-editor/chain') + + const strip = page.locator('.failover-status') + await expect(strip).toBeVisible({ timeout: 10_000 }) + const chainPill = strip.locator('.failover-status__state .status-pill') + await expect(chainPill).toHaveText(/fallback/i) + await expect(chainPill).toHaveClass(/status-pill--warning/) + await expect(strip.locator('.failover-status__active')).toHaveText('b') + + // Both targets are listed with their health. + const rows = strip.locator('tbody tr') + await expect(rows).toHaveCount(2) + await expect(rows.nth(0)).toContainText('a') + await expect(rows.nth(1).locator('.status-pill')).toHaveClass(/status-pill--success/) + // Long errors are truncated but kept whole in the title. + await expect(rows.nth(1).locator('.failover-status__error')).toHaveAttribute('title', /connection refused/) + }) + + test('admin can pin a target after confirming', async ({ page }) => { + await mockEditor(page, NO_AUTH) + let pinBody = null + await page.route('**/api/failover/chain/pin', async (route) => { + if (route.request().method() === 'POST') pinBody = route.request().postDataJSON() + await route.fulfill({ contentType: 'application/json', body: JSON.stringify({ ...CHAIN, pinned: 'b' }) }) + }) + await page.goto('/app/model-editor/chain') + + const strip = page.locator('.failover-status') + await expect(strip).toBeVisible({ timeout: 10_000 }) + const pinB = strip.locator('tbody tr').nth(1).getByRole('button', { name: /pin/i }) + await expect(pinB).toBeVisible() + await pinB.click() + + const dialog = page.getByRole('alertdialog') + await expect(dialog).toBeVisible() + await dialog.getByRole('button', { name: /^pin$/i }).click() + await expect.poll(() => pinBody).toEqual({ target: 'b' }) + }) + + // The editor route itself is admin-only (RequireAdmin), so a non-admin never + // reaches the strip or its pin controls; the component additionally gates + // pinning on isAdmin for surfaces that are not admin-gated. + test('non-admin users get no pin controls', async ({ page }) => { + await mockEditor(page, NON_ADMIN) + await page.route('**/api/auth/me', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(NON_ADMIN.user) })) + await page.goto('/app/model-editor/chain') + + await page.waitForURL(/\/app(?!\/model-editor)/, { timeout: 5000 }) + await expect(page.locator('.failover-status')).toHaveCount(0) + await expect(page.getByRole('button', { name: /^pin/i })).toHaveCount(0) + }) + + test('models that are not failover chains show no strip', async ({ page }) => { + await mockEditor(page, NO_AUTH) + await page.route('**/api/models/edit/**', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: 'name: plain\n', name: 'plain' }) })) + await page.goto('/app/model-editor/plain') + await expect(page.locator('h1.page-title')).toBeVisible({ timeout: 10_000 }) + await expect(page.locator('.failover-status')).toHaveCount(0) + }) +}) diff --git a/core/http/react-ui/e2e/failover-overview.spec.js b/core/http/react-ui/e2e/failover-overview.spec.js new file mode 100644 index 000000000..d9c1688c0 --- /dev/null +++ b/core/http/react-ui/e2e/failover-overview.spec.js @@ -0,0 +1,134 @@ +import { test, expect } from './coverage-fixtures.js' + +// Failover overview page (Operate -> Runtime -> Failover) and the matching +// "chain" badge on an installed model that is itself a failover chain. +// +// The overview is a dense table over GET /api/failover, kept live by the +// /api/failover/events SSE stream (same hook the Model Editor's chain strip +// uses), plus an empty state that points at the Failover Chain template. + +const NO_AUTH = { authEnabled: false, staticApiKeyRequired: false, providers: [] } + +const CHAIN_A = { + name: 'chain-a', + state: 'primary', + active: 'a', + active_since: '2026-09-26T09:00:00Z', + pinned: null, + targets: [ + { model: 'a', kind: 'local', warm: true, state: 'healthy', last_probe: '2026-09-26T09:59:00Z' }, + { model: 'b', kind: 'remote', warm: false, state: 'healthy' }, + ], +} + +const CHAIN_B = { + name: 'chain-b', + state: 'degraded', + active: 'x', + active_since: '2026-09-26T08:00:00Z', + pinned: null, + targets: [ + { model: 'x', kind: 'local', warm: true, state: 'down' }, + { model: 'y', kind: 'local', warm: false, state: 'healthy' }, + ], +} + +async function mockAuth(page, authStatus = NO_AUTH) { + await page.route('**/api/auth/status', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(authStatus) })) +} + +async function mockFailover(page, chains) { + await page.route('**/api/failover', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains }) })) + await page.route('**/api/failover/events', (route) => + route.fulfill({ + status: 200, + headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, + body: `event: snapshot\ndata: ${JSON.stringify({ chains })}\n\n`, + })) +} + +test.describe('Failover overview', () => { + test.beforeEach(async ({ page }) => { + page.on('pageerror', (err) => { + throw new Error(`uncaught page error: ${err.message}`) + }) + }) + + test('admin sees the Failover entry in the Operate rail', async ({ page }) => { + await mockAuth(page) + await mockFailover(page, [CHAIN_A]) + await page.goto('/app/operate') + + const rail = page.locator('.console-layout > .console-rail') + const link = rail.locator('a.nav-item[href="/app/failover"]') + await expect(link).toBeVisible({ timeout: 10_000 }) + await expect(link).toContainText('Failover') + }) + + test('lists chains from GET /api/failover with their live health', async ({ page }) => { + await mockAuth(page) + await mockFailover(page, [CHAIN_A, CHAIN_B]) + await page.goto('/app/failover') + + const rows = page.locator('[data-testid="failover-overview"] tbody tr') + await expect(rows).toHaveCount(2) + + const first = rows.nth(0) + await expect(first.getByRole('link', { name: 'chain-a' })).toHaveAttribute('href', '/app/model-editor/chain-a') + await expect(first.locator('.status-pill').first()).toHaveText(/primary/i) + await expect(first).toContainText('a') + await expect(first.locator('.failover-overview__targets .status-pill')).toHaveCount(2) + + const second = rows.nth(1) + await expect(second.getByRole('link', { name: 'chain-b' })).toHaveAttribute('href', '/app/model-editor/chain-b') + await expect(second.locator('.status-pill').first()).toHaveText(/degraded/i) + }) + + test('follows target.state over SSE without a page reload', async ({ page }) => { + await mockAuth(page) + await page.route('**/api/failover', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [CHAIN_A] }) })) + const sseBody = + `event: snapshot\ndata: ${JSON.stringify({ chains: [CHAIN_A] })}\n\n` + + 'event: target.state\ndata: {"type":"target.state","chain":"chain-a","target":"b","from":"healthy","to":"down","error":"dial tcp: refused"}\n\n' + await page.route('**/api/failover/events', (route) => + route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: sseBody })) + await page.goto('/app/failover') + + const row = page.locator('[data-testid="failover-overview"] tbody tr').first() + await expect(row).toBeVisible({ timeout: 10_000 }) + const targetB = row.locator('.failover-overview__targets .status-pill', { hasText: 'b' }) + await expect(targetB).toHaveClass(/status-pill--error/) + }) + + test('shows an empty state pointing at the Failover Chain template', async ({ page }) => { + await mockAuth(page) + await mockFailover(page, []) + await page.goto('/app/failover') + + await expect(page.locator('.empty-state')).toBeVisible({ timeout: 10_000 }) + const cta = page.getByRole('link', { name: /create a failover chain/i }) + await expect(cta).toHaveAttribute('href', '/app/model-editor?template=failover') + }) +}) + +test.describe('Installed Models - chain badge', () => { + test.beforeEach(async ({ page }) => { + await mockAuth(page) + await page.route('**/api/models/capabilities', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ data: [ + { id: 'chain-a', capabilities: ['chat'], backend: 'llama-cpp' }, + ] }) })) + await page.route('**/api/aliases', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify([]) })) + await mockFailover(page, [CHAIN_A]) + }) + + test('renders a read-only chain -> target badge on a chain model', async ({ page }) => { + await page.goto('/app/models?view=installed') + await page.locator('[data-entity="chain-a"]').click() + await expect(page.getByText('chain → a')).toBeVisible({ timeout: 10_000 }) + }) +}) diff --git a/core/http/react-ui/e2e/nodes-fleet-dashboard.spec.js b/core/http/react-ui/e2e/nodes-fleet-dashboard.spec.js index e7b0ef6e0..b572af619 100644 --- a/core/http/react-ui/e2e/nodes-fleet-dashboard.spec.js +++ b/core/http/react-ui/e2e/nodes-fleet-dashboard.spec.js @@ -47,7 +47,7 @@ test.describe('Nodes fleet dashboard', () => { const rail = page.locator('.console-layout > .console-rail') await expect(rail).toBeVisible() - await expect(rail.locator('a.nav-item')).toHaveCount(13) + await expect(rail.locator('a.nav-item')).toHaveCount(14) await expect(rail.locator('a[href="/app/nodes"]')).toHaveClass(/active/) await expect(rail.locator('a[href$="/swagger/index.html"]')).toHaveAttribute('target', '_blank') }) @@ -67,7 +67,7 @@ test.describe('Nodes fleet dashboard', () => { await expect(rail.locator('.console-rail-groups')).toBeHidden() await rail.getByRole('button', { name: 'Expand Operate navigation' }).click() await expect(rail.locator('.console-rail-groups')).toBeVisible() - await expect(rail.locator('a.nav-item')).toHaveCount(13) + await expect(rail.locator('a.nav-item')).toHaveCount(14) }) test('shows aggregate health, capacity, attention filtering, search, sorting, and grouping', async ({ page }) => { diff --git a/core/http/react-ui/public/locales/de/models.json b/core/http/react-ui/public/locales/de/models.json index 779e0e7f6..d88cb8c70 100644 --- a/core/http/react-ui/public/locales/de/models.json +++ b/core/http/react-ui/public/locales/de/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -224,5 +225,67 @@ "pickVision": "Bilder und Dokumente lesen", "pickAudio": "Sprache rein, Sprache raus", "pickVisual": "Bilder und Video erzeugen" + }, + "failover": { + "title": "Failover-Kette", + "servedBy": "Bedient von", + "changed": "geändert {{time}}", + "pinnedTo": "Fixiert auf {{target}}", + "warm": "Vorgeladen", + "never": "Nie", + "states": { + "primary": "Primär", + "fallback": "Fallback", + "degraded": "Beeinträchtigt", + "healthy": "Gesund", + "down": "Ausgefallen", + "recovering": "Erholt sich", + "missing": "Fehlt" + }, + "kinds": { + "local": "Lokal", + "remote": "Remote" + }, + "columns": { + "target": "Ziel", + "kind": "Art", + "warm": "Vorgeladen", + "health": "Zustand", + "lastProbe": "Letzte Prüfung", + "lastError": "Letzter Fehler", + "actions": "Aktionen" + }, + "actions": { + "pin": "Fixieren", + "pinning": "Wird fixiert…", + "unpin": "Lösen", + "unpinning": "Wird gelöst…" + }, + "confirm": { + "pinTitle": "{{target}} fixieren?", + "pinMessage": "Nur {{target}} bedient {{chain}}, bis Sie die Fixierung lösen. Das automatische Failover ist für diese Kette angehalten.", + "unpinTitle": "{{chain}} lösen?", + "unpinMessage": "{{chain}} kehrt zum automatischen Failover zurück." + }, + "errors": { + "pin": "Fixieren fehlgeschlagen: {{error}}", + "unpin": "Lösen fehlgeschlagen: {{error}}" + }, + "overview": { + "title": "Failover-Ketten", + "subtitle": "Live-Zustand und aktives Ziel für jede Failover-Kette.", + "columns": { + "chain": "Kette", + "state": "Status", + "active": "Aktives Ziel", + "targets": "Ziele", + "changed": "Geändert" + }, + "empty": { + "title": "Noch keine Failover-Ketten", + "text": "Fügen Sie eine Failover-Kette hinzu, um fehlerhafte Ziele automatisch zu umgehen.", + "cta": "Failover-Kette erstellen" + } + } } } diff --git a/core/http/react-ui/public/locales/de/nav.json b/core/http/react-ui/public/locales/de/nav.json index e881b3cad..dea1d5f80 100644 --- a/core/http/react-ui/public/locales/de/nav.json +++ b/core/http/react-ui/public/locales/de/nav.json @@ -60,7 +60,8 @@ "middleware": "Middleware", "activity": "Aktivität", "overview": "Übersicht", - "thisMachine": "Dieser Rechner" + "thisMachine": "Dieser Rechner", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/en/models.json b/core/http/react-ui/public/locales/en/models.json index 04e32e4fa..a2150e785 100644 --- a/core/http/react-ui/public/locales/en/models.json +++ b/core/http/react-ui/public/locales/en/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -241,5 +242,67 @@ "pickVision": "Read images and documents", "pickAudio": "Speech in and speech out", "pickVisual": "Generate images and video" + }, + "failover": { + "title": "Failover chain", + "servedBy": "Served by", + "changed": "changed {{time}}", + "pinnedTo": "Pinned to {{target}}", + "warm": "Warm", + "never": "Never", + "states": { + "primary": "Primary", + "fallback": "Fallback", + "degraded": "Degraded", + "healthy": "Healthy", + "down": "Down", + "recovering": "Recovering", + "missing": "Missing" + }, + "kinds": { + "local": "Local", + "remote": "Remote" + }, + "columns": { + "target": "Target", + "kind": "Kind", + "warm": "Warm", + "health": "Health", + "lastProbe": "Last probe", + "lastError": "Last error", + "actions": "Actions" + }, + "actions": { + "pin": "Pin", + "pinning": "Pinning…", + "unpin": "Unpin", + "unpinning": "Unpinning…" + }, + "confirm": { + "pinTitle": "Pin {{target}}?", + "pinMessage": "Only {{target}} serves {{chain}} until you unpin it. Automatic failover stops for this chain.", + "unpinTitle": "Unpin {{chain}}?", + "unpinMessage": "{{chain}} returns to automatic failover." + }, + "errors": { + "pin": "Pin failed: {{error}}", + "unpin": "Unpin failed: {{error}}" + }, + "overview": { + "title": "Failover chains", + "subtitle": "Live health and the active target for every failover chain.", + "columns": { + "chain": "Chain", + "state": "State", + "active": "Active target", + "targets": "Targets", + "changed": "Changed" + }, + "empty": { + "title": "No failover chains yet", + "text": "Add a failover chain to route around unhealthy targets automatically.", + "cta": "Create a failover chain" + } + } } } diff --git a/core/http/react-ui/public/locales/en/nav.json b/core/http/react-ui/public/locales/en/nav.json index ac55dd533..6d4af666a 100644 --- a/core/http/react-ui/public/locales/en/nav.json +++ b/core/http/react-ui/public/locales/en/nav.json @@ -61,7 +61,8 @@ "api": "API", "activity": "Activity", "overview": "Overview", - "thisMachine": "This machine" + "thisMachine": "This machine", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/es/models.json b/core/http/react-ui/public/locales/es/models.json index d833fba55..989189850 100644 --- a/core/http/react-ui/public/locales/es/models.json +++ b/core/http/react-ui/public/locales/es/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -224,5 +225,67 @@ "pickVision": "Leer imágenes y documentos", "pickAudio": "Voz de entrada y de salida", "pickVisual": "Generar imágenes y vídeo" + }, + "failover": { + "title": "Cadena de failover", + "servedBy": "Servido por", + "changed": "cambió {{time}}", + "pinnedTo": "Fijado en {{target}}", + "warm": "Precargado", + "never": "Nunca", + "states": { + "primary": "Primario", + "fallback": "Respaldo", + "degraded": "Degradado", + "healthy": "Correcto", + "down": "Caído", + "recovering": "Recuperándose", + "missing": "Ausente" + }, + "kinds": { + "local": "Local", + "remote": "Remoto" + }, + "columns": { + "target": "Destino", + "kind": "Tipo", + "warm": "Precargado", + "health": "Estado", + "lastProbe": "Última comprobación", + "lastError": "Último error", + "actions": "Acciones" + }, + "actions": { + "pin": "Fijar", + "pinning": "Fijando…", + "unpin": "Liberar", + "unpinning": "Liberando…" + }, + "confirm": { + "pinTitle": "¿Fijar {{target}}?", + "pinMessage": "Solo {{target}} atiende {{chain}} hasta que lo liberes. El failover automático se detiene para esta cadena.", + "unpinTitle": "¿Liberar {{chain}}?", + "unpinMessage": "{{chain}} vuelve al failover automático." + }, + "errors": { + "pin": "No se pudo fijar: {{error}}", + "unpin": "No se pudo liberar: {{error}}" + }, + "overview": { + "title": "Cadenas de failover", + "subtitle": "Estado en vivo y destino activo de cada cadena de failover.", + "columns": { + "chain": "Cadena", + "state": "Estado", + "active": "Destino activo", + "targets": "Destinos", + "changed": "Cambiado" + }, + "empty": { + "title": "Aún no hay cadenas de failover", + "text": "Agrega una cadena de failover para evitar automáticamente los destinos en mal estado.", + "cta": "Crear una cadena de failover" + } + } } } diff --git a/core/http/react-ui/public/locales/es/nav.json b/core/http/react-ui/public/locales/es/nav.json index 8976ff270..333fd4919 100644 --- a/core/http/react-ui/public/locales/es/nav.json +++ b/core/http/react-ui/public/locales/es/nav.json @@ -60,7 +60,8 @@ "middleware": "Middleware", "activity": "Actividad", "overview": "Resumen", - "thisMachine": "Esta máquina" + "thisMachine": "Esta máquina", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/id/models.json b/core/http/react-ui/public/locales/id/models.json index 9088647b9..1dee74031 100644 --- a/core/http/react-ui/public/locales/id/models.json +++ b/core/http/react-ui/public/locales/id/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -237,5 +238,67 @@ "pickVision": "Membaca gambar dan dokumen", "pickAudio": "Suara masuk dan keluar", "pickVisual": "Membuat gambar dan video" + }, + "failover": { + "title": "Rantai failover", + "servedBy": "Dilayani oleh", + "changed": "berubah {{time}}", + "pinnedTo": "Disematkan ke {{target}}", + "warm": "Siap", + "never": "Belum pernah", + "states": { + "primary": "Utama", + "fallback": "Cadangan", + "degraded": "Menurun", + "healthy": "Sehat", + "down": "Mati", + "recovering": "Memulihkan", + "missing": "Tidak ada" + }, + "kinds": { + "local": "Lokal", + "remote": "Jarak jauh" + }, + "columns": { + "target": "Target", + "kind": "Jenis", + "warm": "Siap", + "health": "Kesehatan", + "lastProbe": "Pemeriksaan terakhir", + "lastError": "Kesalahan terakhir", + "actions": "Aksi" + }, + "actions": { + "pin": "Sematkan", + "pinning": "Menyematkan…", + "unpin": "Lepas", + "unpinning": "Melepas…" + }, + "confirm": { + "pinTitle": "Sematkan {{target}}?", + "pinMessage": "Hanya {{target}} yang melayani {{chain}} sampai Anda melepasnya. Failover otomatis berhenti untuk rantai ini.", + "unpinTitle": "Lepas {{chain}}?", + "unpinMessage": "{{chain}} kembali ke failover otomatis." + }, + "errors": { + "pin": "Gagal menyematkan: {{error}}", + "unpin": "Gagal melepas: {{error}}" + }, + "overview": { + "title": "Rantai failover", + "subtitle": "Kesehatan langsung dan target aktif untuk setiap rantai failover.", + "columns": { + "chain": "Rantai", + "state": "Status", + "active": "Target aktif", + "targets": "Target", + "changed": "Berubah" + }, + "empty": { + "title": "Belum ada rantai failover", + "text": "Tambahkan rantai failover untuk melewati target yang tidak sehat secara otomatis.", + "cta": "Buat rantai failover" + } + } } } diff --git a/core/http/react-ui/public/locales/id/nav.json b/core/http/react-ui/public/locales/id/nav.json index ab8f8735d..90d1a2453 100644 --- a/core/http/react-ui/public/locales/id/nav.json +++ b/core/http/react-ui/public/locales/id/nav.json @@ -61,7 +61,8 @@ "api": "API", "activity": "Aktivitas", "overview": "Ikhtisar", - "thisMachine": "Mesin ini" + "thisMachine": "Mesin ini", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/it/models.json b/core/http/react-ui/public/locales/it/models.json index cbc06d22c..edcc1b587 100644 --- a/core/http/react-ui/public/locales/it/models.json +++ b/core/http/react-ui/public/locales/it/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -224,5 +225,67 @@ "pickVision": "Leggere immagini e documenti", "pickAudio": "Voce in ingresso e in uscita", "pickVisual": "Generare immagini e video" + }, + "failover": { + "title": "Catena di failover", + "servedBy": "Servito da", + "changed": "cambiato {{time}}", + "pinnedTo": "Fissato su {{target}}", + "warm": "Pronto", + "never": "Mai", + "states": { + "primary": "Primario", + "fallback": "Fallback", + "degraded": "Degradato", + "healthy": "Integro", + "down": "Non disponibile", + "recovering": "In ripristino", + "missing": "Mancante" + }, + "kinds": { + "local": "Locale", + "remote": "Remoto" + }, + "columns": { + "target": "Destinazione", + "kind": "Tipo", + "warm": "Pronto", + "health": "Stato", + "lastProbe": "Ultimo controllo", + "lastError": "Ultimo errore", + "actions": "Azioni" + }, + "actions": { + "pin": "Fissa", + "pinning": "Fissaggio…", + "unpin": "Sblocca", + "unpinning": "Sblocco…" + }, + "confirm": { + "pinTitle": "Fissare {{target}}?", + "pinMessage": "Solo {{target}} serve {{chain}} finché non lo sblocchi. Il failover automatico si ferma per questa catena.", + "unpinTitle": "Sbloccare {{chain}}?", + "unpinMessage": "{{chain}} torna al failover automatico." + }, + "errors": { + "pin": "Fissaggio non riuscito: {{error}}", + "unpin": "Sblocco non riuscito: {{error}}" + }, + "overview": { + "title": "Catene di failover", + "subtitle": "Stato in tempo reale e destinazione attiva per ogni catena di failover.", + "columns": { + "chain": "Catena", + "state": "Stato", + "active": "Destinazione attiva", + "targets": "Destinazioni", + "changed": "Cambiato" + }, + "empty": { + "title": "Nessuna catena di failover", + "text": "Aggiungi una catena di failover per aggirare automaticamente le destinazioni non integre.", + "cta": "Crea una catena di failover" + } + } } } diff --git a/core/http/react-ui/public/locales/it/nav.json b/core/http/react-ui/public/locales/it/nav.json index 74ff7c99d..0f819b039 100644 --- a/core/http/react-ui/public/locales/it/nav.json +++ b/core/http/react-ui/public/locales/it/nav.json @@ -60,7 +60,8 @@ "middleware": "Middleware", "activity": "Attività", "overview": "Panoramica", - "thisMachine": "Questa macchina" + "thisMachine": "Questa macchina", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/ko/models.json b/core/http/react-ui/public/locales/ko/models.json index 8ed38bf0f..b2a20016e 100644 --- a/core/http/react-ui/public/locales/ko/models.json +++ b/core/http/react-ui/public/locales/ko/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -208,5 +209,67 @@ "pickVision": "이미지와 문서 읽기", "pickAudio": "음성 입력과 출력", "pickVisual": "이미지와 영상 생성" + }, + "failover": { + "title": "페일오버 체인", + "servedBy": "처리 대상", + "changed": "{{time}} 변경됨", + "pinnedTo": "{{target}}에 고정됨", + "warm": "예열됨", + "never": "없음", + "states": { + "primary": "기본", + "fallback": "대체", + "degraded": "성능 저하", + "healthy": "정상", + "down": "중단", + "recovering": "복구 중", + "missing": "없음" + }, + "kinds": { + "local": "로컬", + "remote": "원격" + }, + "columns": { + "target": "대상", + "kind": "종류", + "warm": "예열", + "health": "상태", + "lastProbe": "마지막 확인", + "lastError": "마지막 오류", + "actions": "작업" + }, + "actions": { + "pin": "고정", + "pinning": "고정 중…", + "unpin": "고정 해제", + "unpinning": "고정 해제 중…" + }, + "confirm": { + "pinTitle": "{{target}}을(를) 고정할까요?", + "pinMessage": "고정을 해제할 때까지 {{target}}만 {{chain}}을(를) 처리합니다. 이 체인의 자동 페일오버가 중지됩니다.", + "unpinTitle": "{{chain}} 고정을 해제할까요?", + "unpinMessage": "{{chain}}이(가) 자동 페일오버로 돌아갑니다." + }, + "errors": { + "pin": "고정 실패: {{error}}", + "unpin": "고정 해제 실패: {{error}}" + }, + "overview": { + "title": "페일오버 체인", + "subtitle": "모든 페일오버 체인의 실시간 상태와 활성 대상입니다.", + "columns": { + "chain": "체인", + "state": "상태", + "active": "활성 대상", + "targets": "대상", + "changed": "변경됨" + }, + "empty": { + "title": "아직 페일오버 체인이 없습니다", + "text": "상태가 나빌 대상을 자동으로 우회하려면 페일오버 체인을 추가하세요.", + "cta": "페일오버 체인 만들기" + } + } } } diff --git a/core/http/react-ui/public/locales/ko/nav.json b/core/http/react-ui/public/locales/ko/nav.json index 16d40a299..739a53088 100644 --- a/core/http/react-ui/public/locales/ko/nav.json +++ b/core/http/react-ui/public/locales/ko/nav.json @@ -60,7 +60,8 @@ "api": "API", "activity": "활동", "overview": "개요", - "thisMachine": "이 머신" + "thisMachine": "이 머신", + "failover": "페일오버" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/pt-BR/models.json b/core/http/react-ui/public/locales/pt-BR/models.json index c354e89ae..26e567a44 100644 --- a/core/http/react-ui/public/locales/pt-BR/models.json +++ b/core/http/react-ui/public/locales/pt-BR/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -240,5 +241,67 @@ "pickVision": "Leia imagens e documentos", "pickAudio": "Fala de entrada e saída", "pickVisual": "Gere imagens e vídeos" + }, + "failover": { + "title": "Cadeia de failover", + "servedBy": "Atendido por", + "changed": "alterado {{time}}", + "pinnedTo": "Fixado em {{target}}", + "warm": "Pré-carregado", + "never": "Nunca", + "states": { + "primary": "Primário", + "fallback": "Reserva", + "degraded": "Degradado", + "healthy": "Saudável", + "down": "Fora do ar", + "recovering": "Recuperando", + "missing": "Ausente" + }, + "kinds": { + "local": "Local", + "remote": "Remoto" + }, + "columns": { + "target": "Destino", + "kind": "Tipo", + "warm": "Pré-carregado", + "health": "Saúde", + "lastProbe": "Última verificação", + "lastError": "Último erro", + "actions": "Ações" + }, + "actions": { + "pin": "Fixar", + "pinning": "Fixando…", + "unpin": "Desafixar", + "unpinning": "Desafixando…" + }, + "confirm": { + "pinTitle": "Fixar {{target}}?", + "pinMessage": "Somente {{target}} atende {{chain}} até você desafixar. O failover automático para nesta cadeia.", + "unpinTitle": "Desafixar {{chain}}?", + "unpinMessage": "{{chain}} volta ao failover automático." + }, + "errors": { + "pin": "Falha ao fixar: {{error}}", + "unpin": "Falha ao desafixar: {{error}}" + }, + "overview": { + "title": "Cadeias de failover", + "subtitle": "Saúde em tempo real e destino ativo de cada cadeia de failover.", + "columns": { + "chain": "Cadeia", + "state": "Estado", + "active": "Destino ativo", + "targets": "Destinos", + "changed": "Alterado" + }, + "empty": { + "title": "Ainda não há cadeias de failover", + "text": "Adicione uma cadeia de failover para contornar automaticamente destinos com problemas.", + "cta": "Criar uma cadeia de failover" + } + } } } diff --git a/core/http/react-ui/public/locales/pt-BR/nav.json b/core/http/react-ui/public/locales/pt-BR/nav.json index f0164bf0b..2d53ff717 100644 --- a/core/http/react-ui/public/locales/pt-BR/nav.json +++ b/core/http/react-ui/public/locales/pt-BR/nav.json @@ -61,7 +61,8 @@ "api": "API", "activity": "Atividade", "overview": "Visão geral", - "thisMachine": "Esta máquina" + "thisMachine": "Esta máquina", + "failover": "Failover" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/public/locales/zh-CN/models.json b/core/http/react-ui/public/locales/zh-CN/models.json index 667009901..40f78260e 100644 --- a/core/http/react-ui/public/locales/zh-CN/models.json +++ b/core/http/react-ui/public/locales/zh-CN/models.json @@ -40,7 +40,8 @@ "license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source", "files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files", "adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.", - "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}" + "alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}", + "chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}" }, "open": { "title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS", @@ -224,5 +225,67 @@ "pickVision": "读取图像与文档", "pickAudio": "语音输入与输出", "pickVisual": "生成图像与视频" + }, + "failover": { + "title": "故障转移链", + "servedBy": "当前服务", + "changed": "{{time}}变更", + "pinnedTo": "已固定到 {{target}}", + "warm": "预热", + "never": "从未", + "states": { + "primary": "主用", + "fallback": "备用", + "degraded": "降级", + "healthy": "健康", + "down": "不可用", + "recovering": "恢复中", + "missing": "缺失" + }, + "kinds": { + "local": "本地", + "remote": "远程" + }, + "columns": { + "target": "目标", + "kind": "类型", + "warm": "预热", + "health": "健康状态", + "lastProbe": "上次探测", + "lastError": "上次错误", + "actions": "操作" + }, + "actions": { + "pin": "固定", + "pinning": "正在固定…", + "unpin": "取消固定", + "unpinning": "正在取消固定…" + }, + "confirm": { + "pinTitle": "固定 {{target}}?", + "pinMessage": "在取消固定之前,只有 {{target}} 为 {{chain}} 提供服务。此链的自动故障转移将停止。", + "unpinTitle": "取消固定 {{chain}}?", + "unpinMessage": "{{chain}} 将恢复自动故障转移。" + }, + "errors": { + "pin": "固定失败:{{error}}", + "unpin": "取消固定失败:{{error}}" + }, + "overview": { + "title": "故障转移链", + "subtitle": "每条故障转移链的实时健康状态和当前目标。", + "columns": { + "chain": "链", + "state": "状态", + "active": "当前目标", + "targets": "目标", + "changed": "变更时间" + }, + "empty": { + "title": "尚无故障转移链", + "text": "添加故障转移链以自动绕过不健康的目标。", + "cta": "创建故障转移链" + } + } } } diff --git a/core/http/react-ui/public/locales/zh-CN/nav.json b/core/http/react-ui/public/locales/zh-CN/nav.json index 6e88ead91..d3dd01d04 100644 --- a/core/http/react-ui/public/locales/zh-CN/nav.json +++ b/core/http/react-ui/public/locales/zh-CN/nav.json @@ -60,7 +60,8 @@ "middleware": "Middleware", "activity": "活动", "overview": "概览", - "thisMachine": "本机" + "thisMachine": "本机", + "failover": "故障转移" }, "footer": { "github": "GitHub", diff --git a/core/http/react-ui/src/App.css b/core/http/react-ui/src/App.css index 5bafd472a..3ac4da254 100644 --- a/core/http/react-ui/src/App.css +++ b/core/http/react-ui/src/App.css @@ -12275,6 +12275,15 @@ button.collapsible-header:focus-visible { .cfr-num { width: 120px; font-size: var(--text-sm); } .cfr-select { width: 220px; font-size: var(--text-sm); } .cfr-radio { display: flex; align-items: center; gap: var(--spacing-sm); font-size: var(--text-sm); cursor: pointer; } + +/* Failover targets editor — chain member list with a warm toggle and + duplicate/self-reference warnings. */ +.fte-list { display: flex; flex-direction: column; gap: var(--spacing-sm); width: 100%; } +.fte-empty { font-size: var(--text-sm); color: var(--color-text-muted); padding: var(--spacing-sm) 0; } +.fte-warning { font-size: var(--text-xs); color: var(--color-warning); padding: var(--spacing-sm) 0; } +.fte-row { padding: var(--spacing-sm); display: flex; flex-direction: column; gap: var(--spacing-xs); } +.fte-row--error { border-color: var(--color-error); } +.fte-row__head { font-size: var(--text-xs); color: var(--color-text-muted); } .pill-tiny { padding: 2px 6px; font-size: var(--text-xs); } /* =========================================================================== @@ -12361,6 +12370,48 @@ button.collapsible-header:focus-visible { justify-content: space-between; padding: var(--spacing-lg) var(--spacing-lg) var(--spacing-md); } +/* Failover chain health strip (Model Editor, under the header) */ +.failover-status { + padding: 0 var(--spacing-lg) var(--spacing-md); + border-bottom: 1px solid var(--color-border); +} +.failover-status__summary { + display: flex; + flex-wrap: wrap; + align-items: center; + gap: var(--spacing-xs) var(--spacing-md); + padding-bottom: var(--spacing-sm); + font-size: var(--text-sm); + color: var(--color-text-secondary); +} +.failover-status__label { + font-size: 0.75rem; + font-weight: 500; + letter-spacing: 0.04em; + text-transform: uppercase; + color: var(--color-text-muted); +} +.failover-status__active { color: var(--color-text-primary); font-family: var(--font-mono); } +.failover-status__pinned { color: var(--color-warning); } +.failover-status__unpin { margin-left: auto; } +.failover-status__table td, .failover-status__table th { white-space: nowrap; } +.failover-status__table code { font-family: var(--font-mono); font-size: 0.8125rem; } +.failover-status__row--active td:first-child { box-shadow: inset 2px 0 0 var(--color-success); } +.failover-status__error { + display: inline-block; + max-width: 28ch; + overflow: hidden; + text-overflow: ellipsis; + vertical-align: bottom; + color: var(--color-error); +} +.failover-status__action { text-align: right; } + +/* Failover overview (Operate -> Runtime): one row per chain, with a compact + strip of per-target pills standing in for the detailed table the Model + Editor shows for a single chain. */ +.failover-overview__targets { display: inline-flex; flex-wrap: wrap; gap: 4px; } +.failover-overview__target-pill { font-size: 0.75rem; } .me-tabs { display: flex; gap: 0; padding: 0 var(--spacing-lg); border-bottom: 1px solid var(--color-border); } .me-tab { padding: var(--spacing-sm) var(--spacing-md); diff --git a/core/http/react-ui/src/components/ConfigFieldRenderer.jsx b/core/http/react-ui/src/components/ConfigFieldRenderer.jsx index 4e93e1e5a..26cdd8d01 100644 --- a/core/http/react-ui/src/components/ConfigFieldRenderer.jsx +++ b/core/http/react-ui/src/components/ConfigFieldRenderer.jsx @@ -11,6 +11,7 @@ import PatternListEditor from './PatternListEditor' import ModelMultiSelect from './ModelMultiSelect' import RouterCandidatesEditor from './RouterCandidatesEditor' import RouterPoliciesEditor from './RouterPoliciesEditor' +import FailoverTargetsEditor from './FailoverTargetsEditor' // Map autocomplete provider to SearchableModelSelect capability const PROVIDER_TO_CAPABILITY = { @@ -381,6 +382,23 @@ export default function ConfigFieldRenderer({ field, value, onChange, onRemove, ) } + // Failover targets — ordered member list of a failover chain. Each row + // is {model, warm}; duplicate/self-reference detection reads the edited + // model's own name from FormContext. + if (component === 'failover-targets') { + return ( +
+
+
+
+
{description}
+
+
+ +
+ ) + } + // PII detectors — a capability-filtered multi-select of token_classify // models (the consuming model's pii.detectors list). if (component === 'model-multi-select') { diff --git a/core/http/react-ui/src/components/FailoverChainStatus.jsx b/core/http/react-ui/src/components/FailoverChainStatus.jsx new file mode 100644 index 000000000..be1579218 --- /dev/null +++ b/core/http/react-ui/src/components/FailoverChainStatus.jsx @@ -0,0 +1,164 @@ +import { useState } from 'react' +import { useTranslation } from 'react-i18next' +import StatusPill from './StatusPill' +import ConfirmDialog from './ConfirmDialog' +import useFailoverChains from '../hooks/useFailoverChains' +import { useAuth } from '../context/AuthContext' +import { failoverApi } from '../utils/api' + +const UNITS = [ + ['day', 86_400], + ['hour', 3_600], + ['minute', 60], +] + +// Localized "5 minutes ago" from an RFC 3339 timestamp. Intl keeps the phrase +// in the viewer's language without a translation key per unit. Exported so +// other failover surfaces (the overview table) share the same phrasing. +export function relative(ts, lng) { + const ms = Date.parse(ts) + if (!ms) return null + const seconds = Math.round((ms - Date.now()) / 1000) + const rtf = new Intl.RelativeTimeFormat(lng, { numeric: 'auto' }) + for (const [unit, size] of UNITS) { + if (Math.abs(seconds) >= size) return rtf.format(Math.round(seconds / size), unit) + } + return rtf.format(seconds, 'second') +} + +// FailoverChainStatus renders the live health of one failover chain: its +// state, the target serving it, and a per-target health table. Pin controls +// appear only when canPin, and every pin change needs a confirmation because +// it overrides automatic failover for all callers. +export default function FailoverChainStatus({ chain, onPin, onUnpin, canPin = false }) { + const { t, i18n } = useTranslation('models') + const [confirm, setConfirm] = useState(null) + const [pending, setPending] = useState(false) + + const runConfirmed = async () => { + setPending(true) + try { + if (confirm.kind === 'pin') await onPin?.(confirm.target) + else await onUnpin?.() + } finally { + setPending(false) + setConfirm(null) + } + } + + const since = relative(chain.active_since, i18n.language) + + return ( +
+
+ {t('failover.title')} + + + + + {t('failover.servedBy')} {chain.active} + + {since && ( + + {t('failover.changed', { time: since })} + + )} + {chain.pinned && ( + +