mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
Merge PR #12285: feat(failover): serve a model name from a chain of local and remote targets
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
commit
9c156656bd
186 files changed
+17157
-1085
No files matched your search
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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/
|
||||
|
||||
|
||||
@@ -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).
|
||||
@@ -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
|
||||
########################################################
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,453 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
const transcriptionsPath = "/v1/audio/transcriptions"
|
||||
|
||||
// ttsRequest is the body of LocalAI's /tts (schema.TTSRequest).
|
||||
type ttsRequest struct {
|
||||
Model string `json:"model"`
|
||||
Input string `json:"input"`
|
||||
Voice string `json:"voice,omitempty"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Instructions string `json:"instructions,omitempty"`
|
||||
Params map[string]string `json:"params,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) ttsRequest(req *pb.TTSRequest, stream bool) ttsRequest {
|
||||
return ttsRequest{
|
||||
Model: p.model(""),
|
||||
Input: req.GetText(),
|
||||
Voice: req.GetVoice(),
|
||||
Language: req.GetLanguage(),
|
||||
Instructions: req.GetInstructions(),
|
||||
Params: req.GetParams(),
|
||||
Stream: stream,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) TTS(req *pb.TTSRequest) error {
|
||||
return p.postJSONToFile(context.Background(), "/tts", p.ttsRequest(req, false), req.GetDst())
|
||||
}
|
||||
|
||||
// TTSStream forwards the upstream's chunked audio unchanged. That body is
|
||||
// already what core expects from a streaming backend: a WAV header followed
|
||||
// by PCM. out is closed on every path because the gRPC server drains it until
|
||||
// closed and would otherwise hang.
|
||||
func (p *LocalAIProxy) TTSStream(req *pb.TTSRequest, out chan []byte) error {
|
||||
defer close(out)
|
||||
resp, err := p.postStream(context.Background(), "/tts", p.ttsRequest(req, true))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
n, err := resp.Body.Read(buf)
|
||||
if n > 0 {
|
||||
// The reader reuses buf, so each chunk needs its own copy.
|
||||
out <- append([]byte(nil), buf[:n]...)
|
||||
}
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
// A cut-off stream is a failed synthesis, not a short one.
|
||||
return transportError("/tts", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// soundGenerationRequest is the body of /v1/sound-generation
|
||||
// (schema.ElevenLabsSoundGenerationRequest). Pointers keep "unset" distinct
|
||||
// from zero so the upstream model's defaults apply.
|
||||
type soundGenerationRequest struct {
|
||||
ModelID string `json:"model_id"`
|
||||
Text string `json:"text"`
|
||||
Duration *float32 `json:"duration_seconds,omitempty"`
|
||||
Temperature *float32 `json:"prompt_influence,omitempty"`
|
||||
DoSample *bool `json:"do_sample,omitempty"`
|
||||
Think *bool `json:"think,omitempty"`
|
||||
Caption string `json:"caption,omitempty"`
|
||||
Lyrics string `json:"lyrics,omitempty"`
|
||||
BPM *int32 `json:"bpm,omitempty"`
|
||||
Keyscale string `json:"keyscale,omitempty"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Timesignature string `json:"timesignature,omitempty"`
|
||||
Instrumental *bool `json:"instrumental,omitempty"`
|
||||
}
|
||||
|
||||
// SoundGeneration refuses audio-conditioned requests: the REST endpoint has no
|
||||
// field for a source clip, and dropping it would return unconditioned audio as
|
||||
// if it were the answer.
|
||||
func (p *LocalAIProxy) SoundGeneration(req *pb.SoundGenerationRequest) error {
|
||||
if req.GetSrc() != "" {
|
||||
return unimplemented("SoundGeneration with src")
|
||||
}
|
||||
body := soundGenerationRequest{
|
||||
ModelID: p.model(""),
|
||||
Text: req.GetText(),
|
||||
Duration: req.Duration,
|
||||
Temperature: req.Temperature,
|
||||
DoSample: req.Sample,
|
||||
Think: req.Think,
|
||||
Caption: req.GetCaption(),
|
||||
Lyrics: req.GetLyrics(),
|
||||
BPM: req.Bpm,
|
||||
Keyscale: req.GetKeyscale(),
|
||||
Language: req.GetLanguage(),
|
||||
Timesignature: req.GetTimesignature(),
|
||||
Instrumental: req.Instrumental,
|
||||
}
|
||||
return p.postJSONToFile(context.Background(), "/v1/sound-generation", body, req.GetDst())
|
||||
}
|
||||
|
||||
// transcriptionForm builds the /v1/audio/transcriptions upload. Dst is the
|
||||
// input audio (core names it that way). Language and translate are only sent
|
||||
// when set so the upstream model config's defaults still apply; diarize is
|
||||
// always sent because the upstream treats a missing field as true.
|
||||
func (p *LocalAIProxy) transcriptionForm(req *pb.TranscriptRequest, stream bool) multipartForm {
|
||||
f := url.Values{}
|
||||
f.Set("model", p.model(""))
|
||||
f.Set("diarize", strconv.FormatBool(req.GetDiarize()))
|
||||
// verbose_json keeps segments and words; the default drops nothing today,
|
||||
// but naming it keeps the reply shape pinned.
|
||||
f.Set("response_format", "verbose_json")
|
||||
if v := req.GetLanguage(); v != "" {
|
||||
f.Set("language", v)
|
||||
}
|
||||
if req.GetTranslate() {
|
||||
f.Set("translate", "true")
|
||||
}
|
||||
if v := req.GetPrompt(); v != "" {
|
||||
f.Set("prompt", v)
|
||||
}
|
||||
if v := req.GetTemperature(); v != 0 {
|
||||
f.Set("temperature", formatFloat(v))
|
||||
}
|
||||
for _, g := range req.GetTimestampGranularities() {
|
||||
f.Add("timestamp_granularities[]", g)
|
||||
}
|
||||
if stream {
|
||||
f.Set("stream", "true")
|
||||
}
|
||||
return multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}
|
||||
}
|
||||
|
||||
// transcriptionResult is TranscriptionResultSeconds: the REST API reports
|
||||
// times in seconds, pb in nanoseconds (core reads them as time.Duration).
|
||||
type transcriptionResult struct {
|
||||
Text string `json:"text"`
|
||||
Language string `json:"language"`
|
||||
Duration float64 `json:"duration"`
|
||||
Segments []transcriptionSegment `json:"segments"`
|
||||
}
|
||||
|
||||
type transcriptionSegment struct {
|
||||
ID int32 `json:"id"`
|
||||
Start float64 `json:"start"`
|
||||
End float64 `json:"end"`
|
||||
Text string `json:"text"`
|
||||
Tokens []int32 `json:"tokens"`
|
||||
Speaker string `json:"speaker"`
|
||||
Words []struct {
|
||||
Start float64 `json:"start"`
|
||||
End float64 `json:"end"`
|
||||
Text string `json:"text"`
|
||||
} `json:"words"`
|
||||
}
|
||||
|
||||
// toProto drops the top-level words: core rebuilds them from segment words.
|
||||
func (r transcriptionResult) toProto() *pb.TranscriptResult {
|
||||
out := &pb.TranscriptResult{Text: r.Text, Language: r.Language, Duration: float32(r.Duration)}
|
||||
for _, s := range r.Segments {
|
||||
seg := &pb.TranscriptSegment{
|
||||
Id: s.ID, Start: nanos(s.Start), End: nanos(s.End), Text: s.Text,
|
||||
Tokens: s.Tokens, Speaker: s.Speaker,
|
||||
}
|
||||
for _, w := range s.Words {
|
||||
seg.Words = append(seg.Words, &pb.TranscriptWord{Start: nanos(w.Start), End: nanos(w.End), Text: w.Text})
|
||||
}
|
||||
out.Segments = append(out.Segments, seg)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func nanos(seconds float64) int64 {
|
||||
return int64(seconds * float64(time.Second))
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) AudioTranscription(ctx context.Context, req *pb.TranscriptRequest) (pb.TranscriptResult, error) {
|
||||
var resp transcriptionResult
|
||||
if err := p.postForm(ctx, transcriptionsPath, p.transcriptionForm(req, false), &resp); err != nil {
|
||||
return pb.TranscriptResult{}, err
|
||||
}
|
||||
return *resp.toProto(), nil
|
||||
}
|
||||
|
||||
// transcriptEvent covers every frame of the upstream transcription SSE stream.
|
||||
type transcriptEvent struct {
|
||||
Type string `json:"type"`
|
||||
Delta string `json:"delta"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
transcriptionResult
|
||||
}
|
||||
|
||||
// AudioTranscriptionStream maps upstream SSE frames to stream responses. out
|
||||
// is closed on every path: the gRPC server drains it until closed.
|
||||
func (p *LocalAIProxy) AudioTranscriptionStream(ctx context.Context, req *pb.TranscriptRequest, out chan *pb.TranscriptStreamResponse) error {
|
||||
defer close(out)
|
||||
resp, err := p.postMultipartStream(ctx, transcriptionsPath, p.transcriptionForm(req, true))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
// The done frame carries every segment of the recording.
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 4<<20)
|
||||
for scanner.Scan() {
|
||||
payload, ok := strings.CutPrefix(scanner.Text(), "data:")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
payload = strings.TrimSpace(payload)
|
||||
if payload == "" || payload == "[DONE]" {
|
||||
continue
|
||||
}
|
||||
var ev transcriptEvent
|
||||
if err := json.Unmarshal([]byte(payload), &ev); err != nil {
|
||||
xlog.Debug("localai-proxy: skip malformed SSE frame", "path", transcriptionsPath, "error", err)
|
||||
continue
|
||||
}
|
||||
switch ev.Type {
|
||||
case "transcript.text.delta":
|
||||
out <- &pb.TranscriptStreamResponse{Delta: ev.Delta}
|
||||
case "transcript.text.done":
|
||||
out <- &pb.TranscriptStreamResponse{FinalResult: ev.toProto()}
|
||||
return nil
|
||||
case "error":
|
||||
msg := "unknown error"
|
||||
if ev.Error != nil {
|
||||
msg = ev.Error.Message
|
||||
}
|
||||
xlog.Warn("localai-proxy: upstream stream error", "path", transcriptionsPath, "error", msg)
|
||||
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", transcriptionsPath, msg)
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return transportError(transcriptionsPath, err)
|
||||
}
|
||||
// The upstream always ends with a done or error frame, so a stream that
|
||||
// stops without one was cut off; returning nil would pass the partial
|
||||
// deltas off as the whole transcript.
|
||||
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream ended before the final transcript", transcriptionsPath)
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) {
|
||||
f := url.Values{}
|
||||
f.Set("model", p.model(""))
|
||||
// verbose_json keeps per-segment text; core strips it again for callers
|
||||
// that did not ask for it.
|
||||
f.Set("response_format", "verbose_json")
|
||||
if v := req.GetLanguage(); v != "" {
|
||||
f.Set("language", v)
|
||||
}
|
||||
for name, v := range map[string]int32{
|
||||
"num_speakers": req.GetNumSpeakers(), "min_speakers": req.GetMinSpeakers(), "max_speakers": req.GetMaxSpeakers(),
|
||||
} {
|
||||
if v != 0 {
|
||||
f.Set(name, strconv.Itoa(int(v)))
|
||||
}
|
||||
}
|
||||
for name, v := range map[string]float32{
|
||||
"clustering_threshold": req.GetClusteringThreshold(),
|
||||
"min_duration_on": req.GetMinDurationOn(),
|
||||
"min_duration_off": req.GetMinDurationOff(),
|
||||
} {
|
||||
if v != 0 {
|
||||
f.Set(name, formatFloat(v))
|
||||
}
|
||||
}
|
||||
if req.GetIncludeText() {
|
||||
f.Set("include_text", "true")
|
||||
}
|
||||
|
||||
var resp struct {
|
||||
Duration float64 `json:"duration"`
|
||||
Language string `json:"language"`
|
||||
NumSpeakers int32 `json:"num_speakers"`
|
||||
Segments []struct {
|
||||
ID int32 `json:"id"`
|
||||
Speaker string `json:"speaker"`
|
||||
Label string `json:"label"`
|
||||
Start float32 `json:"start"`
|
||||
End float32 `json:"end"`
|
||||
Text string `json:"text"`
|
||||
} `json:"segments"`
|
||||
}
|
||||
if err := p.postForm(context.Background(), "/v1/audio/diarization",
|
||||
multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}, &resp); err != nil {
|
||||
return pb.DiarizeResponse{}, err
|
||||
}
|
||||
var segments []*pb.DiarizeSegment
|
||||
for _, s := range resp.Segments {
|
||||
// The upstream renames speakers to SPEAKER_NN and keeps the backend's
|
||||
// own id in label. Core renames again, so hand it the raw id: the
|
||||
// caller then sees the same speakers and labels a direct call gives.
|
||||
speaker := s.Label
|
||||
if speaker == "" {
|
||||
speaker = s.Speaker
|
||||
}
|
||||
segments = append(segments, &pb.DiarizeSegment{
|
||||
Id: s.ID, Start: s.Start, End: s.End, Speaker: speaker, Text: s.Text,
|
||||
})
|
||||
}
|
||||
return pb.DiarizeResponse{
|
||||
Segments: segments, NumSpeakers: resp.NumSpeakers, Duration: float32(resp.Duration), Language: resp.Language,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) VAD(req *pb.VADRequest) (pb.VADResponse, error) {
|
||||
var resp struct {
|
||||
Segments []struct {
|
||||
Start float32 `json:"start"`
|
||||
End float32 `json:"end"`
|
||||
} `json:"segments"`
|
||||
}
|
||||
body := map[string]any{"model": p.model(""), "audio": req.GetAudio()}
|
||||
if err := p.postJSON(context.Background(), "/v1/vad", body, &resp); err != nil {
|
||||
return pb.VADResponse{}, err
|
||||
}
|
||||
var segments []*pb.VADSegment
|
||||
for _, s := range resp.Segments {
|
||||
segments = append(segments, &pb.VADSegment{Start: s.Start, End: s.End})
|
||||
}
|
||||
return pb.VADResponse{Segments: segments}, nil
|
||||
}
|
||||
|
||||
func (p *LocalAIProxy) SoundDetection(ctx context.Context, req *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) {
|
||||
f := url.Values{}
|
||||
f.Set("model", p.model(""))
|
||||
if v := req.GetTopK(); v != 0 {
|
||||
f.Set("top_k", strconv.Itoa(int(v)))
|
||||
}
|
||||
if v := req.GetThreshold(); v != 0 {
|
||||
f.Set("threshold", formatFloat(v))
|
||||
}
|
||||
var resp struct {
|
||||
Detections []struct {
|
||||
Index int32 `json:"index"`
|
||||
Label string `json:"label"`
|
||||
Score float32 `json:"score"`
|
||||
} `json:"detections"`
|
||||
}
|
||||
if err := p.postForm(ctx, "/v1/audio/classification",
|
||||
multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetSrc()}}}, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := &pb.SoundDetectionResponse{}
|
||||
for _, d := range resp.Detections {
|
||||
out.Detections = append(out.Detections, &pb.SoundClass{Index: d.Index, Label: d.Label, Score: d.Score})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// stemsHeader names the other outputs of a separation run: the body carries
|
||||
// one file, and the rest are served under /generated-audio/.
|
||||
const stemsHeader = "X-Audio-Stems"
|
||||
|
||||
// AudioTransform uploads the input (and reference) to /audio/transformations
|
||||
// and writes the returned audio to Dst. SampleRate and Samples stay 0: the
|
||||
// REST reply does not report them, and core only uses them for tracing.
|
||||
func (p *LocalAIProxy) AudioTransform(req *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) {
|
||||
f := url.Values{}
|
||||
f.Set("model", p.model(""))
|
||||
for k, v := range req.GetParams() {
|
||||
f.Set("params["+k+"]", v)
|
||||
}
|
||||
form := multipartForm{fields: f, files: []formFile{{field: "audio", path: req.GetAudioPath()}}}
|
||||
if ref := req.GetReferencePath(); ref != "" {
|
||||
form.files = append(form.files, formFile{field: "reference", path: ref})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
header, err := p.postMultipartToFile(ctx, "/audio/transformations", form, req.GetDst())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &pb.AudioTransformResult{
|
||||
Dst: req.GetDst(),
|
||||
ReferenceProvided: req.GetReferencePath() != "",
|
||||
Stems: p.fetchStems(ctx, header.Get(stemsHeader), req.GetDst()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// fetchStems downloads each stem the upstream names and writes it beside
|
||||
// dst, where core looks for them. A stem that cannot be fetched is dropped
|
||||
// with a warning rather than failing the call: the caller still gets the
|
||||
// audio it asked for.
|
||||
func (p *LocalAIProxy) fetchStems(ctx context.Context, header, dst string) []*pb.AudioTransformStem {
|
||||
if header == "" {
|
||||
return nil
|
||||
}
|
||||
var entries []struct {
|
||||
Name string `json:"name"`
|
||||
URL string `json:"url"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(header), &entries); err != nil {
|
||||
xlog.Warn("localai-proxy: ignore malformed stems header", "error", err)
|
||||
return nil
|
||||
}
|
||||
dir := filepath.Dir(dst)
|
||||
prefix := strings.TrimSuffix(filepath.Base(dst), filepath.Ext(dst))
|
||||
var stems []*pb.AudioTransformStem
|
||||
for _, e := range entries {
|
||||
// Only files under /generated-audio/ are stems. Anything else is not
|
||||
// a path this API hands out, so it is not fetched.
|
||||
escaped, ok := strings.CutPrefix(e.URL, "/generated-audio/")
|
||||
if !ok || e.Name == "" {
|
||||
xlog.Warn("localai-proxy: skip stem with unexpected url", "name", e.Name, "url", e.URL)
|
||||
continue
|
||||
}
|
||||
name, err := url.PathUnescape(escaped)
|
||||
if err != nil || name == "" || name != filepath.Base(name) || name == "." || name == ".." {
|
||||
xlog.Warn("localai-proxy: skip stem with unsafe file name", "name", e.Name, "url", e.URL)
|
||||
continue
|
||||
}
|
||||
local := filepath.Join(dir, prefix+"-"+name)
|
||||
if err := p.getToFile(ctx, e.URL, local); err != nil {
|
||||
xlog.Warn("localai-proxy: stem download failed", "name", e.Name, "error", err)
|
||||
continue
|
||||
}
|
||||
stems = append(stems, &pb.AudioTransformStem{Name: e.Name, Dst: local})
|
||||
}
|
||||
return stems
|
||||
}
|
||||
|
||||
// formatFloat prints a float32 form value without float64 noise
|
||||
// (0.2, not 0.20000000298023224).
|
||||
func formatFloat(v float32) string {
|
||||
return strconv.FormatFloat(float64(v), 'g', -1, 32)
|
||||
}
|
||||
@@ -0,0 +1,489 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
// writeInput writes content to a fresh temp file and returns its path, so a
|
||||
// spec can check the upstream received exactly those bytes.
|
||||
func writeInput(name, content string) string {
|
||||
path := filepath.Join(GinkgoT().TempDir(), name)
|
||||
Expect(os.WriteFile(path, []byte(content), 0o600)).To(Succeed())
|
||||
return path
|
||||
}
|
||||
|
||||
// drainBytes collects every chunk sent on ch until it is closed.
|
||||
func drainBytes(ch chan []byte) <-chan [][]byte {
|
||||
done := make(chan [][]byte, 1)
|
||||
go func() {
|
||||
var got [][]byte
|
||||
for c := range ch {
|
||||
got = append(got, c)
|
||||
}
|
||||
done <- got
|
||||
}()
|
||||
return done
|
||||
}
|
||||
|
||||
func drainTranscript(ch chan *pb.TranscriptStreamResponse) <-chan []*pb.TranscriptStreamResponse {
|
||||
done := make(chan []*pb.TranscriptStreamResponse, 1)
|
||||
go func() {
|
||||
var got []*pb.TranscriptStreamResponse
|
||||
for c := range ch {
|
||||
got = append(got, c)
|
||||
}
|
||||
done <- got
|
||||
}()
|
||||
return done
|
||||
}
|
||||
|
||||
// cutMidStream answers 200, flushes body, then drops the connection without
|
||||
// finishing the chunked encoding, as a crashed upstream would.
|
||||
func cutMidStream(contentType, body string) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(body))
|
||||
w.(http.Flusher).Flush()
|
||||
conn, _, err := w.(http.Hijacker).Hijack()
|
||||
if err == nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var _ = Describe("audio methods", func() {
|
||||
var up *fakeUpstream
|
||||
|
||||
BeforeEach(func() {
|
||||
up = newFakeUpstream()
|
||||
DeferCleanup(up.Close)
|
||||
})
|
||||
|
||||
Describe("TTS", func() {
|
||||
It("posts to /tts and writes the upstream audio to Dst", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-audio-bytes"})
|
||||
dst := filepath.Join(GinkgoT().TempDir(), "out.wav")
|
||||
lang := "it"
|
||||
instr := "cheerful"
|
||||
|
||||
Expect(p.TTS(&pb.TTSRequest{
|
||||
Text: "ciao", Model: "local-path", Dst: dst, Voice: "v1", Language: &lang,
|
||||
Instructions: &instr, Params: map[string]string{"speed": "1.2"},
|
||||
})).To(Succeed())
|
||||
|
||||
got, err := os.ReadFile(dst)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(got)).To(Equal("RIFF-audio-bytes"))
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/tts"))
|
||||
Expect(req.JSON).To(Equal(map[string]any{
|
||||
"model": "remote-model", "input": "ciao", "voice": "v1", "language": "it",
|
||||
"instructions": "cheerful", "params": map[string]any{"speed": "1.2"},
|
||||
}))
|
||||
})
|
||||
|
||||
It("maps an upstream failure and leaves no partial file", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/tts", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"})
|
||||
dst := filepath.Join(GinkgoT().TempDir(), "out.wav")
|
||||
err := p.TTS(&pb.TTSRequest{Text: "x", Dst: dst})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Expect(dst).NotTo(BeAnExistingFile())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("TTSStream", func() {
|
||||
It("forwards the chunked WAV body in order and closes the channel", func() {
|
||||
header := "RIFF" + string(make([]byte, 40))
|
||||
up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "audio/wav")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
for _, c := range []string{header, "pcm-1", "pcm-2"} {
|
||||
_, _ = w.Write([]byte(c))
|
||||
w.(http.Flusher).Flush()
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
})
|
||||
DeferCleanup(up.Close)
|
||||
p := loadProxy(up, nil)
|
||||
|
||||
out := make(chan []byte)
|
||||
done := drainBytes(out)
|
||||
Expect(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)).To(Succeed())
|
||||
|
||||
var chunks [][]byte
|
||||
Eventually(done).Should(Receive(&chunks))
|
||||
Expect(chunks).NotTo(BeEmpty())
|
||||
Expect(string(chunks[0][:4])).To(Equal("RIFF"))
|
||||
var all []byte
|
||||
for _, c := range chunks {
|
||||
all = append(all, c...)
|
||||
}
|
||||
Expect(string(all)).To(Equal(header + "pcm-1pcm-2"))
|
||||
})
|
||||
|
||||
It("asks the upstream to stream", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF"})
|
||||
out := make(chan []byte)
|
||||
done := drainBytes(out)
|
||||
Expect(p.TTSStream(&pb.TTSRequest{Text: "hi", Voice: "v"}, out)).To(Succeed())
|
||||
Eventually(done).Should(Receive())
|
||||
Expect(up.last().JSON).To(HaveKeyWithValue("stream", true))
|
||||
Expect(up.last().JSON).To(HaveKeyWithValue("input", "hi"))
|
||||
})
|
||||
|
||||
It("reports a mid-stream disconnect as Unavailable and still closes the channel", func() {
|
||||
cut := newFakeUpstreamWithHandler(cutMidStream("audio/wav", "RIFFpartial"))
|
||||
DeferCleanup(cut.Close)
|
||||
p := loadProxy(cut, nil)
|
||||
out := make(chan []byte)
|
||||
done := drainBytes(out)
|
||||
err := p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Eventually(done).Should(Receive())
|
||||
})
|
||||
|
||||
It("closes the channel when the upstream refuses the request", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/tts", scriptedResponse{Status: http.StatusBadRequest, Body: "bad"})
|
||||
out := make(chan []byte)
|
||||
done := drainBytes(out)
|
||||
Expect(codeOf(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out))).To(Equal(codes.InvalidArgument))
|
||||
Eventually(done).Should(Receive())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("SoundGeneration", func() {
|
||||
It("posts the ElevenLabs body to /v1/sound-generation and writes Dst", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/sound-generation", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-sfx"})
|
||||
dst := filepath.Join(GinkgoT().TempDir(), "sfx.wav")
|
||||
dur, temp, bpm := float32(4.5), float32(0.3), int32(120)
|
||||
sample, instrumental := true, false
|
||||
|
||||
Expect(p.SoundGeneration(&pb.SoundGenerationRequest{
|
||||
Text: "rain", Dst: dst, Duration: &dur, Temperature: &temp, Sample: &sample,
|
||||
Bpm: &bpm, Caption: ptr("cap"), Lyrics: ptr("la"), Keyscale: ptr("C major"),
|
||||
Language: ptr("en"), Timesignature: ptr("4/4"), Instrumental: &instrumental,
|
||||
})).To(Succeed())
|
||||
|
||||
got, err := os.ReadFile(dst)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(got)).To(Equal("RIFF-sfx"))
|
||||
Expect(up.last().JSON).To(Equal(map[string]any{
|
||||
"model_id": "remote-model", "text": "rain", "duration_seconds": 4.5,
|
||||
"prompt_influence": 0.3, "do_sample": true, "bpm": float64(120),
|
||||
"caption": "cap", "lyrics": "la", "keyscale": "C major", "language": "en",
|
||||
"timesignature": "4/4", "instrumental": false,
|
||||
}))
|
||||
})
|
||||
|
||||
It("refuses audio-conditioned generation it cannot upload", func() {
|
||||
p := loadProxy(up, nil)
|
||||
err := p.SoundGeneration(&pb.SoundGenerationRequest{Text: "x", Src: ptr("/tmp/in.wav")})
|
||||
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
|
||||
Expect(up.recorded()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("AudioTranscription", func() {
|
||||
It("uploads Dst with the request fields and maps the seconds-based result", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/audio/transcriptions", map[string]any{
|
||||
"text": "hello world", "language": "en", "duration": 2.5,
|
||||
"segments": []any{map[string]any{
|
||||
"id": 0, "start": 0.5, "end": 1.25, "text": "hello world", "tokens": []int{1, 2},
|
||||
"speaker": "SPEAKER_00",
|
||||
"words": []any{map[string]any{"start": 0.5, "end": 0.75, "text": "hello"}},
|
||||
}},
|
||||
})
|
||||
in := writeInput("speech.wav", "RIFF-speech")
|
||||
|
||||
res, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{
|
||||
Dst: in, Language: "en", Translate: true, Diarize: true, Prompt: "names: Ada",
|
||||
Temperature: 0.2, TimestampGranularities: []string{"word"},
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/audio/transcriptions"))
|
||||
Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-speech"}))
|
||||
Expect(req.Fields).To(Equal(map[string]string{
|
||||
"model": "remote-model", "language": "en", "translate": "true", "diarize": "true",
|
||||
"prompt": "names: Ada", "temperature": "0.2", "timestamp_granularities[]": "word",
|
||||
"response_format": "verbose_json",
|
||||
}))
|
||||
|
||||
Expect(res.Text).To(Equal("hello world"))
|
||||
Expect(res.Language).To(Equal("en"))
|
||||
Expect(res.Duration).To(BeNumerically("==", 2.5))
|
||||
Expect(res.Segments).To(HaveLen(1))
|
||||
seg := res.Segments[0]
|
||||
// pb carries nanoseconds (core reads them as time.Duration).
|
||||
Expect(seg.Start).To(Equal(int64(500 * time.Millisecond)))
|
||||
Expect(seg.End).To(Equal(int64(1250 * time.Millisecond)))
|
||||
Expect(seg.Text).To(Equal("hello world"))
|
||||
Expect(seg.Tokens).To(Equal([]int32{1, 2}))
|
||||
Expect(seg.Speaker).To(Equal("SPEAKER_00"))
|
||||
Expect(seg.Words).To(HaveLen(1))
|
||||
Expect(seg.Words[0].Start).To(Equal(int64(500 * time.Millisecond)))
|
||||
Expect(seg.Words[0].Text).To(Equal("hello"))
|
||||
})
|
||||
|
||||
It("sends diarize=false explicitly because the upstream defaults it on", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "x"})
|
||||
_, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
f := up.last().Fields
|
||||
Expect(f).To(HaveKeyWithValue("diarize", "false"))
|
||||
// Unset language and translate are left to the upstream model config.
|
||||
Expect(f).NotTo(HaveKey("language"))
|
||||
Expect(f).NotTo(HaveKey("translate"))
|
||||
})
|
||||
|
||||
It("maps a 4xx to InvalidArgument", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/audio/transcriptions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad audio"})
|
||||
_, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")})
|
||||
Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("AudioTranscriptionStream", func() {
|
||||
It("emits deltas, then the final result, and closes the channel", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}),
|
||||
sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "lo"}),
|
||||
sseJSON(map[string]any{
|
||||
"type": "transcript.text.done", "text": "hello", "language": "en", "duration": 1.5,
|
||||
"segments": []any{map[string]any{"id": 0, "start": 0.0, "end": 1.5, "text": "hello"}},
|
||||
}),
|
||||
"[DONE]",
|
||||
}})
|
||||
out := make(chan *pb.TranscriptStreamResponse)
|
||||
done := drainTranscript(out)
|
||||
|
||||
Expect(p.AudioTranscriptionStream(context.Background(),
|
||||
&pb.TranscriptRequest{Dst: writeInput("a.wav", "RIFF-a"), Stream: true}, out)).To(Succeed())
|
||||
|
||||
var got []*pb.TranscriptStreamResponse
|
||||
Eventually(done).Should(Receive(&got))
|
||||
Expect(got).To(HaveLen(3))
|
||||
Expect(got[0].Delta).To(Equal("hel"))
|
||||
Expect(got[1].Delta).To(Equal("lo"))
|
||||
final := got[2].FinalResult
|
||||
Expect(final).NotTo(BeNil())
|
||||
Expect(final.Text).To(Equal("hello"))
|
||||
Expect(final.Language).To(Equal("en"))
|
||||
Expect(final.Duration).To(BeNumerically("==", 1.5))
|
||||
Expect(final.Segments).To(HaveLen(1))
|
||||
Expect(final.Segments[0].End).To(Equal(int64(1500 * time.Millisecond)))
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Fields).To(HaveKeyWithValue("stream", "true"))
|
||||
Expect(req.Files).To(HaveKeyWithValue("file", "RIFF-a"))
|
||||
})
|
||||
|
||||
It("returns an upstream error event as Unavailable", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{
|
||||
sseJSON(map[string]any{"type": "error", "error": map[string]any{"message": "decoder died"}}),
|
||||
"[DONE]",
|
||||
}})
|
||||
out := make(chan *pb.TranscriptStreamResponse)
|
||||
done := drainTranscript(out)
|
||||
err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out)
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
Expect(err).To(MatchError(ContainSubstring("decoder died")))
|
||||
Eventually(done).Should(Receive())
|
||||
})
|
||||
|
||||
It("reports an upstream disconnect before the final result as Unavailable", func() {
|
||||
frame := "data: " + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}) + "\n\n"
|
||||
cut := newFakeUpstreamWithHandler(cutMidStream("text/event-stream", frame))
|
||||
DeferCleanup(cut.Close)
|
||||
p := loadProxy(cut, nil)
|
||||
out := make(chan *pb.TranscriptStreamResponse)
|
||||
done := drainTranscript(out)
|
||||
err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out)
|
||||
Expect(codeOf(err)).To(Equal(codes.Unavailable))
|
||||
var got []*pb.TranscriptStreamResponse
|
||||
Eventually(done).Should(Receive(&got))
|
||||
Expect(got).To(HaveLen(1))
|
||||
})
|
||||
|
||||
It("closes the channel when the local file is missing", func() {
|
||||
p := loadProxy(up, nil)
|
||||
out := make(chan *pb.TranscriptStreamResponse)
|
||||
done := drainTranscript(out)
|
||||
err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: "/nonexistent.wav"}, out)
|
||||
Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
|
||||
Eventually(done).Should(Receive())
|
||||
Expect(up.recorded()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("Diarize", func() {
|
||||
It("uploads Dst with the tuning fields and maps the segments", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/audio/diarization", map[string]any{
|
||||
"task": "diarize", "duration": 3.0, "language": "en", "num_speakers": 2,
|
||||
"segments": []any{
|
||||
map[string]any{"id": 0, "speaker": "SPEAKER_00", "label": "spk_a", "start": 0.0, "end": 1.5, "text": "hi"},
|
||||
map[string]any{"id": 1, "speaker": "SPEAKER_01", "start": 1.5, "end": 3.0},
|
||||
},
|
||||
})
|
||||
res, err := p.Diarize(&pb.DiarizeRequest{
|
||||
Dst: writeInput("talk.wav", "RIFF-talk"), Language: "en", NumSpeakers: 2, MinSpeakers: 1,
|
||||
MaxSpeakers: 3, ClusteringThreshold: 0.5, MinDurationOn: 0.1, MinDurationOff: 0.2, IncludeText: true,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/v1/audio/diarization"))
|
||||
Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-talk"}))
|
||||
Expect(req.Fields).To(Equal(map[string]string{
|
||||
"model": "remote-model", "language": "en", "num_speakers": "2", "min_speakers": "1",
|
||||
"max_speakers": "3", "clustering_threshold": "0.5", "min_duration_on": "0.1",
|
||||
"min_duration_off": "0.2", "include_text": "true", "response_format": "verbose_json",
|
||||
}))
|
||||
|
||||
Expect(res.NumSpeakers).To(Equal(int32(2)))
|
||||
Expect(res.Duration).To(BeNumerically("==", 3.0))
|
||||
Expect(res.Language).To(Equal("en"))
|
||||
Expect(res.Segments).To(HaveLen(2))
|
||||
// The raw upstream label survives, so core's own normalisation
|
||||
// yields the same speakers a direct call would.
|
||||
Expect(res.Segments[0].Speaker).To(Equal("spk_a"))
|
||||
Expect(res.Segments[0].Text).To(Equal("hi"))
|
||||
Expect(res.Segments[1].Speaker).To(Equal("SPEAKER_01"))
|
||||
Expect(res.Segments[1].Start).To(BeNumerically("==", 1.5))
|
||||
Expect(res.Segments[1].Id).To(Equal(int32(1)))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("VAD", func() {
|
||||
It("posts the samples to /v1/vad and maps the segments", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/vad", map[string]any{"segments": []any{map[string]any{"start": 0.25, "end": 1.5}}})
|
||||
res, err := p.VAD(&pb.VADRequest{Audio: []float32{0.5, -0.25}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(up.last().JSON).To(Equal(map[string]any{"model": "remote-model", "audio": []any{0.5, -0.25}}))
|
||||
Expect(res.Segments).To(HaveLen(1))
|
||||
Expect(res.Segments[0].Start).To(BeNumerically("==", 0.25))
|
||||
Expect(res.Segments[0].End).To(BeNumerically("==", 1.5))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("SoundDetection", func() {
|
||||
It("uploads Src to /v1/audio/classification and maps the detections", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.replyJSON("/v1/audio/classification", map[string]any{"model": "remote-model", "detections": []any{
|
||||
map[string]any{"index": 74, "label": "Dog", "score": 0.9},
|
||||
map[string]any{"index": 0, "label": "Speech", "score": 0.4},
|
||||
}})
|
||||
res, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{
|
||||
Src: writeInput("bark.wav", "RIFF-bark"), TopK: 5, Threshold: 0.25,
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
req := up.last()
|
||||
Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-bark"}))
|
||||
Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "top_k": "5", "threshold": "0.25"}))
|
||||
Expect(res.Detections).To(HaveLen(2))
|
||||
Expect(res.Detections[0].Label).To(Equal("Dog"))
|
||||
Expect(res.Detections[0].Index).To(Equal(int32(74)))
|
||||
Expect(res.Detections[0].Score).To(BeNumerically("~", 0.9, 1e-6))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("AudioTransform", func() {
|
||||
It("uploads audio and reference with params and writes the result to Dst", func() {
|
||||
p := loadProxy(up, nil)
|
||||
up.script("/audio/transformations", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-clean"})
|
||||
dst := filepath.Join(GinkgoT().TempDir(), "transform.wav")
|
||||
|
||||
res, err := p.AudioTransform(&pb.AudioTransformRequest{
|
||||
AudioPath: writeInput("mic.wav", "RIFF-mic"), ReferencePath: writeInput("ref.wav", "RIFF-ref"),
|
||||
Dst: dst, Params: map[string]string{"noise_gate": "true"},
|
||||
})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
got, err := os.ReadFile(dst)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(got)).To(Equal("RIFF-clean"))
|
||||
Expect(res.Dst).To(Equal(dst))
|
||||
Expect(res.ReferenceProvided).To(BeTrue())
|
||||
Expect(res.Stems).To(BeEmpty())
|
||||
|
||||
req := up.last()
|
||||
Expect(req.Path).To(Equal("/audio/transformations"))
|
||||
Expect(req.Files).To(Equal(map[string]string{"audio": "RIFF-mic", "reference": "RIFF-ref"}))
|
||||
Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "params[noise_gate]": "true"}))
|
||||
})
|
||||
|
||||
It("fetches the stems the upstream names and writes them beside Dst", func() {
|
||||
p := loadProxy(up, nil)
|
||||
stems := `[{"name":"vocals","url":"/generated-audio/sep%20vocals.wav"},{"name":"evil","url":"/etc/passwd"}]`
|
||||
up.script("/audio/transformations", scriptedResponse{
|
||||
Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-drums",
|
||||
Header: map[string]string{"X-Audio-Stems": stems},
|
||||
})
|
||||
up.script("/generated-audio/sep vocals.wav", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-vocals"})
|
||||
dir := GinkgoT().TempDir()
|
||||
dst := filepath.Join(dir, "transform.wav")
|
||||
|
||||
res, err := p.AudioTransform(&pb.AudioTransformRequest{AudioPath: writeInput("song.wav", "RIFF-song"), Dst: dst})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(res.ReferenceProvided).To(BeFalse())
|
||||
Expect(res.Stems).To(HaveLen(1))
|
||||
Expect(res.Stems[0].Name).To(Equal("vocals"))
|
||||
// Core keeps only stems that are direct children of Dst's directory.
|
||||
Expect(filepath.Dir(res.Stems[0].Dst)).To(Equal(dir))
|
||||
got, err := os.ReadFile(res.Stems[0].Dst)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(string(got)).To(Equal("RIFF-vocals"))
|
||||
|
||||
var paths []string
|
||||
for _, r := range up.recorded() {
|
||||
paths = append(paths, r.Method+" "+r.Path)
|
||||
}
|
||||
Expect(paths).To(Equal([]string{"POST /audio/transformations", "GET /generated-audio/sep vocals.wav"}))
|
||||
})
|
||||
})
|
||||
|
||||
It("names the upload file after the local file", func() {
|
||||
// The upstream saves the upload under its base name and some handlers
|
||||
// pick the decoder by extension, so the name must reach it intact.
|
||||
var name string
|
||||
named := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, fh, err := r.FormFile("file")
|
||||
if err == nil {
|
||||
name = fh.Filename
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = fmt.Fprint(w, `{"detections":[]}`)
|
||||
})
|
||||
DeferCleanup(named.Close)
|
||||
p := loadProxy(named, nil)
|
||||
_, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{Src: writeInput("clip.mp3", "ID3")})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(name).To(Equal("clip.mp3"))
|
||||
})
|
||||
})
|
||||
|
||||
func ptr[T any](v T) *T { return &v }
|
||||
@@ -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
|
||||
}
|
||||
@@ -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: <frame>" events and wins over Body. Header adds response headers.
|
||||
type scriptedResponse struct {
|
||||
Status int
|
||||
ContentType string
|
||||
Header map[string]string
|
||||
Body string
|
||||
SSE []string
|
||||
}
|
||||
|
||||
// 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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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:<mime>;base64,<payload>` 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
|
||||
}
|
||||
@@ -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))))
|
||||
})
|
||||
})
|
||||
})
|
||||
Executable
+13
@@ -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/
|
||||
@@ -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:<name>"]).
|
||||
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")
|
||||
}
|
||||
Executable
+6
@@ -0,0 +1,6 @@
|
||||
#!/bin/bash
|
||||
set -ex
|
||||
|
||||
CURDIR=$(dirname "$(realpath "$0")")
|
||||
|
||||
exec "$CURDIR"/localai-proxy "$@"
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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{"</s>"},
|
||||
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("</s>")))
|
||||
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)))
|
||||
})
|
||||
})
|
||||
@@ -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"
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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")))
|
||||
})
|
||||
})
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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 != "" {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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{
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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(""))
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
))
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 }
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
})
|
||||
@@ -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)")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 },
|
||||
])
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
@@ -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 })
|
||||
})
|
||||
})
|
||||
@@ -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 }) => {
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -60,7 +60,8 @@
|
||||
"middleware": "Middleware",
|
||||
"activity": "Aktivität",
|
||||
"overview": "Übersicht",
|
||||
"thisMachine": "Dieser Rechner"
|
||||
"thisMachine": "Dieser Rechner",
|
||||
"failover": "Failover"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -61,7 +61,8 @@
|
||||
"api": "API",
|
||||
"activity": "Activity",
|
||||
"overview": "Overview",
|
||||
"thisMachine": "This machine"
|
||||
"thisMachine": "This machine",
|
||||
"failover": "Failover"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -60,7 +60,8 @@
|
||||
"middleware": "Middleware",
|
||||
"activity": "Actividad",
|
||||
"overview": "Resumen",
|
||||
"thisMachine": "Esta máquina"
|
||||
"thisMachine": "Esta máquina",
|
||||
"failover": "Failover"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -61,7 +61,8 @@
|
||||
"api": "API",
|
||||
"activity": "Aktivitas",
|
||||
"overview": "Ikhtisar",
|
||||
"thisMachine": "Mesin ini"
|
||||
"thisMachine": "Mesin ini",
|
||||
"failover": "Failover"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -60,7 +60,8 @@
|
||||
"middleware": "Middleware",
|
||||
"activity": "Attività",
|
||||
"overview": "Panoramica",
|
||||
"thisMachine": "Questa macchina"
|
||||
"thisMachine": "Questa macchina",
|
||||
"failover": "Failover"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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": "페일오버 체인 만들기"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -60,7 +60,8 @@
|
||||
"api": "API",
|
||||
"activity": "활동",
|
||||
"overview": "개요",
|
||||
"thisMachine": "이 머신"
|
||||
"thisMachine": "이 머신",
|
||||
"failover": "페일오버"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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",
|
||||
|
||||
@@ -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": "创建故障转移链"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -60,7 +60,8 @@
|
||||
"middleware": "Middleware",
|
||||
"activity": "活动",
|
||||
"overview": "概览",
|
||||
"thisMachine": "本机"
|
||||
"thisMachine": "本机",
|
||||
"failover": "故障转移"
|
||||
},
|
||||
"footer": {
|
||||
"github": "GitHub",
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 (
|
||||
<div className="list-row">
|
||||
<div className="hstack hstack--between mb-xs">
|
||||
<div>
|
||||
<div className="text-base fw-medium"><FieldLabel field={field} /></div>
|
||||
<div className="text-meta mt-xs">{description}</div>
|
||||
</div>
|
||||
</div>
|
||||
<FailoverTargetsEditor value={value} onChange={handleChange} />
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
// PII detectors — a capability-filtered multi-select of token_classify
|
||||
// models (the consuming model's pii.detectors list).
|
||||
if (component === 'model-multi-select') {
|
||||
|
||||
@@ -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 (
|
||||
<section className="failover-status" aria-label={t('failover.title')}>
|
||||
<div className="failover-status__summary">
|
||||
<span className="failover-status__label">{t('failover.title')}</span>
|
||||
<span className="failover-status__state">
|
||||
<StatusPill status={chain.state} label={t(`failover.states.${chain.state}`, chain.state)} />
|
||||
</span>
|
||||
<span className="failover-status__meta">
|
||||
{t('failover.servedBy')} <code className="failover-status__active">{chain.active}</code>
|
||||
</span>
|
||||
{since && (
|
||||
<span className="text-muted" title={new Date(chain.active_since).toLocaleString(i18n.language)}>
|
||||
{t('failover.changed', { time: since })}
|
||||
</span>
|
||||
)}
|
||||
{chain.pinned && (
|
||||
<span className="failover-status__pinned">
|
||||
<i className="fas fa-thumbtack" aria-hidden="true" /> {t('failover.pinnedTo', { target: chain.pinned })}
|
||||
</span>
|
||||
)}
|
||||
{canPin && chain.pinned && (
|
||||
<button type="button" className="btn btn-ghost btn-sm failover-status__unpin" onClick={() => setConfirm({ kind: 'unpin' })}>
|
||||
{t('failover.actions.unpin')}
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="table-container">
|
||||
<table className="table failover-status__table">
|
||||
<thead>
|
||||
<tr>
|
||||
<th>{t('failover.columns.target')}</th>
|
||||
<th>{t('failover.columns.kind')}</th>
|
||||
<th>{t('failover.columns.warm')}</th>
|
||||
<th>{t('failover.columns.health')}</th>
|
||||
<th>{t('failover.columns.lastProbe')}</th>
|
||||
<th>{t('failover.columns.lastError')}</th>
|
||||
{canPin && <th><span className="sr-only">{t('failover.columns.actions')}</span></th>}
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{(chain.targets || []).map(target => (
|
||||
<tr key={target.model} className={target.model === chain.active ? 'failover-status__row--active' : undefined}>
|
||||
<td><code>{target.model}</code></td>
|
||||
<td>{t(`failover.kinds.${target.kind}`, target.kind)}</td>
|
||||
<td>{target.warm ? t('failover.warm') : <span className="text-muted">-</span>}</td>
|
||||
<td><StatusPill status={target.state} label={t(`failover.states.${target.state}`, target.state)} /></td>
|
||||
<td className="text-muted">{relative(target.last_probe, i18n.language) || t('failover.never')}</td>
|
||||
<td>
|
||||
{target.last_error
|
||||
? <span className="failover-status__error" title={target.last_error}>{target.last_error}</span>
|
||||
: <span className="text-muted">-</span>}
|
||||
</td>
|
||||
{canPin && (
|
||||
<td className="failover-status__action">
|
||||
{chain.pinned !== target.model && (
|
||||
<button type="button" className="btn btn-ghost btn-sm" onClick={() => setConfirm({ kind: 'pin', target: target.model })}>
|
||||
{t('failover.actions.pin')}
|
||||
</button>
|
||||
)}
|
||||
</td>
|
||||
)}
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
<ConfirmDialog
|
||||
open={!!confirm}
|
||||
title={confirm?.kind === 'pin'
|
||||
? t('failover.confirm.pinTitle', { target: confirm.target })
|
||||
: t('failover.confirm.unpinTitle', { chain: chain.name })}
|
||||
message={confirm?.kind === 'pin'
|
||||
? t('failover.confirm.pinMessage', { chain: chain.name, target: confirm.target })
|
||||
: t('failover.confirm.unpinMessage', { chain: chain.name })}
|
||||
confirmLabel={confirm?.kind === 'pin' ? t('failover.actions.pin') : t('failover.actions.unpin')}
|
||||
pendingLabel={confirm?.kind === 'pin' ? t('failover.actions.pinning') : t('failover.actions.unpinning')}
|
||||
pending={pending}
|
||||
onConfirm={runConfirmed}
|
||||
onCancel={() => setConfirm(null)}
|
||||
/>
|
||||
</section>
|
||||
)
|
||||
}
|
||||
|
||||
// ModelFailoverStatus is the Model Editor mount point: it renders the strip
|
||||
// only when the edited model is a failover chain, and owns the SSE
|
||||
// subscription so create mode never opens one.
|
||||
export function ModelFailoverStatus({ name, addToast }) {
|
||||
const { t } = useTranslation('models')
|
||||
const { isAdmin } = useAuth()
|
||||
const { byName, refresh } = useFailoverChains()
|
||||
const chain = byName[name]
|
||||
if (!chain) return null
|
||||
|
||||
const act = async (call, errorKey) => {
|
||||
try {
|
||||
await call()
|
||||
} catch (err) {
|
||||
addToast?.(t(errorKey, { error: err.message }), 'error')
|
||||
}
|
||||
refresh()
|
||||
}
|
||||
|
||||
return (
|
||||
<FailoverChainStatus
|
||||
chain={chain}
|
||||
canPin={isAdmin}
|
||||
onPin={(target) => act(() => failoverApi.pin(name, target), 'failover.errors.pin')}
|
||||
onUnpin={() => act(() => failoverApi.unpin(name), 'failover.errors.unpin')}
|
||||
/>
|
||||
)
|
||||
}
|
||||
Loaded 100 of 186 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user