From ee8d796f934ebe09865da7f1fa0ae549edcab906 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 14:51:26 +0000 Subject: [PATCH 01/79] docs: design model failover chains A model served from a remote upstream has no local fallback when that upstream is unhealthy, so clients reimplement failover and lose realtime context when they switch endpoints. Define a failover block on model configs: an ordered list of targets with active probes, in-request retry before the response is committed, fail-back with hysteresis, and switch events over REST, SSE and the realtime socket. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- ...2026-09-26-model-failover-chains-design.md | 398 ++++++++++++++++++ 1 file changed, 398 insertions(+) create mode 100644 docs/superpowers/specs/2026-09-26-model-failover-chains-design.md diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md new file mode 100644 index 000000000..574bc1883 --- /dev/null +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -0,0 +1,398 @@ +# Model failover chains + +Date: 2026-09-26 +Status: design approved in brainstorming, pending spec review + +## Problem + +A LocalAI instance that serves a model from a remote upstream (for example a +`cloud-proxy` model that points at a larger LocalAI cluster) has no way to fall +back to a local model when that upstream is unhealthy. Clients that want this +today build it themselves. The wingman voice assistant, for example, keeps an +ordered list of realtime WebSocket endpoints, quarantines an endpoint after +repeated backend errors, and reconnects to the next one. This has three costs: + +- Every client reimplements failover, health tracking and recovery. +- The switch happens at the session level. A realtime conversation loses its + context when the client reconnects to another endpoint. +- The client sees at least one failed request before it reacts. + +## Goal + +A model config can declare an ordered **failover chain** of target models. +LocalAI serves each request for the chain from the highest-priority healthy +target, retries on the next target when a target fails before the response is +committed, probes targets actively, fails back with hysteresis, and tells +clients when the active target changes. + +Because a chain is a model name, it works in every place that takes a model +name, including the `llm`, `transcription`, `tts` and `vad` fields of a +realtime pipeline. The realtime session stays on the local instance, so a +switch changes the stage that serves the next turn and the conversation +history survives. + +## Non-goals + +- A dedicated WebUI chain editor and live health view. That is a follow-up + (sub-project 3). This spec only registers the config fields in the field + metadata registry, so the generic model editor can show them. +- The `localai-proxy` backend (a remote target kind that covers the full + LocalAI API through gRPC). That is a follow-up (sub-project 1). This spec + defines the "remote target" probe path it will plug into. +- Shared failover state across several LocalAI frontends. State is in memory + and per instance. +- Transparent failover after a response has started to stream. + +## Config schema + +A chain is a model config with a `failover` block. Like an alias, it has no +backend of its own. + +```yaml +name: assistant-llm +failover: + targets: + - model: argus-llm # for example a cloud-proxy config + - model: gemma-local + warm: true # load at startup, never evict + probe: + interval: 15s # liveness probe interval + timeout: 5s + trip: + errors: 1 # retryable failures in `window` that mark a target down + window: 30s + recovery: + probes: 3 # consecutive inference probes to confirm recovery + min_dwell: 60s # minimum time on a lower target before fail-back +``` + +Only `targets` is required. The other values in the example are the defaults. + +### Validation + +The config loader rejects a chain when: + +- `failover` is set together with `alias` or `backend`. +- `targets` has fewer than 2 entries. +- A target does not exist, or the same target is listed twice. +- A target is itself a chain. Chains do not nest. + +A target can be an alias. The alias is resolved one hop, as it is today. + +The loader logs a warning, and does not reject, when: + +- The known usecases of the targets do not overlap (for example an LLM and a + TTS model in one chain). Usecases are often inferred, so this cannot be an + error. +- `warm: true` is set on a remote target. The flag has no effect there. + +### Target kinds + +The kind decides how a target is probed. It is inferred from the backend of the +target config: + +- **remote**: `cloud-proxy`, and `localai-proxy` when it exists. +- **local**: every other backend. + +### Naming in responses and accounting + +The behaviour matches aliases. Responses echo the chain name. Usage and traces +record `requested=` and `served=` through the existing +`ContextKeyRequestedModel` and `ContextKeyServedModel` keys. + +## Failover manager + +New package: `core/services/failover`. The application creates one `Manager` at +start and keeps it in sync with the model config loader. When a chain is added, +edited or removed (from YAML, the model editor API or the MCP tools), the +manager updates without a restart. + +### State + +- Health is tracked **per target**. A target that is in two chains is probed + once, and a failure marks it down for both. +- The active target is tracked **per chain**. + +Target states: + +``` + trip (errors within window, from requests or probes) + healthy ───────────────────────────────▶ down + ▲ │ liveness probe passes + │ recovery.probes consecutive ▼ + └──── inference probes pass ◀──── recovering ──(any failure)──▶ down +``` + +- At startup, targets are `healthy`. The manager runs one liveness pass + immediately, and in-request retry covers the gap until it completes. +- A target whose config is removed becomes `missing`. It is treated as `down`, + and the chain reports it. + +Chain states: + +- `primary`: the active target is target 0. +- `fallback`: the active target is a lower target. +- `degraded`: all targets are down. + +### Selecting the active target + +- The active target is the highest-priority `healthy` target. +- Failover to a lower target is immediate. +- Fail-back to a higher target happens only when that target is `healthy` + (which already needs `recovery.probes` passing inference probes) **and** the + current target has been active for at least `min_dwell`. +- When the chain is `degraded`, requests still try every target in priority + order. A probe can lag behind a recovery, so the manager does not fail fast. +- A manual pin (see API) forces the active target. While a pin is set, probes + continue and report state, but they do not change the active target. + +### Probes + +| Target | Liveness (steady state) | Recovery confirmation | +|---|---|---| +| remote | `GET /readyz`, and the upstream model is listed in `GET /v1/models` | one minimal real request, chosen by usecase | +| local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | the same minimal request, run in-process | +| local, cold | the config and model files exist and the backend is installed. The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | + +Minimal requests by usecase: + +- chat and completion: `max_tokens: 1` +- embeddings: the input `"ping"` +- transcription: 200 ms of silence +- TTS: the text `"ok"` +- expensive usecases (image, video, 3D): no inference probe. Liveness is the + confirmation. + +Probe load rules: + +- Each chain has one ticker with jitter. A target shared by chains is probed + once. +- A successful real request counts as a liveness pass, so a busy target is + almost never probed. +- Inference probes run only while a target is `recovering`. +- Remote probes use the URL and API key from the target's proxy config. +- Probe results go into the same trip counter as request failures. + +### Warm targets + +The manager loads `warm: true` targets at startup and marks them pinned in the +watchdog, so LRU and idle eviction skip them. They still count toward the +active backend limit. When pinned warm targets leave no room for another load, +that load fails with an error that names them. The docs state this. + +### Events + +Each change of a target state or of an active target produces an event on an +internal bus: + +``` +{chain, target, from, to, state, reason, error, at} +``` + +`reason` is one of `trip`, `recovery`, `manual`, `degraded`, `missing`. + +## Request path (HTTP) + +### Resolution + +In `core/http/middleware/request.go`, next to the alias block, a chain config is +resolved with `mgr.Plan(chain)`. The plan is the ordered list of attempts: + +1. the active target, +2. the other `healthy` targets in priority order, +3. the `down` targets, only when the chain is `degraded`. + +The middleware stores the plan in the request context and sets +`MODEL_CONFIG` to the config of the first target. Handlers do not change. The +plan is fixed when the request starts, so a config reload does not affect +requests that are in progress. + +### In-request retry + +A new middleware wraps the handler. It replaces the response writer with one +that records whether the response is committed: + +- A response with status 500 or higher is buffered until the handler returns, + as long as no body has been flushed. Error bodies are small, so the wrapper + can discard them. +- A streaming response (SSE, or any flushed body) is committed at the first + flush. + +When the handler returns a retryable error and the response is not committed, +the wrapper calls `mgr.ReportFailure(target, err)`, sets the config of the next +target in the plan, and runs the handler again. When the handler succeeds, the +wrapper calls `mgr.ReportSuccess(target)`. When the response is committed and +then fails, the wrapper reports the failure and does not retry. + +Request bodies: + +- JSON bodies are already parsed into the request context. +- Multipart bodies are cached by `ParseMultipartForm`. +- Other bodies are buffered up to a limit. A larger body gets no in-request + retry. The failure still counts toward the trip. + +### Retryable errors + +`failover.IsRetryable(err, status)` returns true for: + +- connection and dial errors +- timeouts, when the client did not cancel the request +- gRPC `Unavailable`, `Internal`, `DeadlineExceeded` and `Unknown` +- upstream HTTP 5xx +- model load failures + +It returns false for client cancellation, 4xx responses and validation errors. +These do not trip a target, because the next target would reject the same +request. + +### Handler audit + +Each endpoint family must be safe to run again before its response is +committed: chat, completions, embeddings, transcription, TTS, image generation, +rerank, VAD and sound detection. An endpoint that is not safe gets no +in-request retry (its failures still trip the target). The PR lists these +endpoints. + +### Response headers + +Every response for a chain carries: + +- `X-LocalAI-Served-Model: ` +- `X-LocalAI-Failover: fallback` or `degraded`, when target 0 did not serve the + request + +Plain HTTP clients can see failover without subscribing to events. + +## Request path (realtime) + +- In `core/http/endpoints/openai/realtime_model.go`, a pipeline stage that names + a chain is resolved **for each call** of `wrappedModel`, not once at session + start. A helper, `mgr.Do(ctx, chain, func(cfg *config.ModelConfig) error)`, + goes through the plan with the same classification as HTTP. +- Streaming stages (`Predict` with a token callback, `TTSStream`, + `TranscribeStream`) wrap the callback. A retry is allowed only until the first + token or audio chunk goes to the client. After that, the turn fails as it does + today, the target is tripped, and the next turn uses the next target. +- `TranscribeLive` is resolved when it opens. A failure in the middle of the + stream ends it in the existing way, and the next utterance opens it again on + the new target. +- The conversation history is in the realtime session on this instance. A + switch of the LLM stage keeps it. +- `Warmup` warms the active target of each chain stage. + +## API and events + +All endpoints use the global auth middleware. `GET` endpoints and the event +stream need standard auth. The pin endpoints are admin only. + +### REST + +`GET /api/failover` returns all chains: + +```json +{"chains":[{"name":"assistant-llm","state":"fallback","active":"gemma-local", + "active_since":"2026-09-26T10:00:00Z","pinned":null, + "targets":[ + {"model":"argus-llm","kind":"remote","warm":false,"state":"recovering", + "consecutive_ok":1,"last_probe":"2026-09-26T10:04:10Z", + "last_error":"503 no healthy nodes"}, + {"model":"gemma-local","kind":"local","warm":true,"state":"healthy"}]}]} +``` + +`GET /api/failover/{chain}` returns one chain. + +`POST /api/failover/{chain}/pin` with `{"target": ""}` forces a target. +`DELETE /api/failover/{chain}/pin` removes the pin. The pin is in memory and a +restart clears it. + +### Server-sent events + +`GET /api/failover/events`: + +- The first event is `snapshot`, with the same payload as `GET /api/failover`. + A new client knows the current state without a race against a separate GET. +- Then `chain.switched` with `{chain, from, to, state, reason, at}`, and + `target.state` with `{target, from, to, reason, error, at}`. +- A keepalive comment every 15 s. + +### Realtime server event + +`localai.model.failover`: + +```json +{"type":"localai.model.failover","chain":"assistant-llm","stage":"llm", + "from":"argus-llm","to":"gemma-local","state":"fallback","reason":"trip"} +``` + +- The server sends it to every session whose pipeline uses the chain when the + chain switches. +- The server also sends it once for each chain stage when the session starts, + with `reason: "initial"` and `from` empty. A client knows at the start whether + it runs on the primary or on a fallback. +- `stage` is one of `llm`, `transcription`, `tts`, `vad`, `sound_detection`. + +### Observability + +- Metrics: `localai_failover_switches_total{chain,from,to,reason}` and + `localai_failover_target_up{target}`. +- Each failed attempt in a request is recorded in the Traces UI, so a request + served by target 2 shows why target 1 was skipped. + +## Capability surfaces + +As required by `.agents/api-endpoints-and-auth.md`: + +- Handlers in `core/http/endpoints/localai/failover.go` with swagger blocks, tag + `failover`. Routes in `core/http/routes/localai.go`. `make swagger`. +- An `instructionDefs` entry for the new tag in + `core/http/endpoints/localai/api_instructions.go`, and the count in + `api_instructions_test.go`. +- `failover.*` fields in the config field metadata registry + (`core/config/meta/registry.go`), in a new `failover` section next to + `alias`, so the generic model editor can show and edit them. +- MCP tools in `pkg/mcp/localaitools/`: `list_failover_chains`, + `pin_failover_target` and `unpin_failover_target`, in the `inproc` and + `httpapi` clients, the skill prompts, and `toolToHTTPRoute` in + `coverage_test.go`. Chains are created and edited through the existing model + config tools. +- A docs page, `docs/content/features/model-failover.md`, linked from + `model-aliases.md`, `openai-realtime.md` and the cloud-proxy docs. + +## Testing + +Ginkgo and Gomega, like the rest of LocalAI. The coverage baseline must not go +down. + +- **State machine**, with a fake clock: trip; recovery after N inference probes; + `min_dwell` hysteresis; `degraded`; pin and unpin; a target shared by two + chains; a `missing` target. +- **`IsRetryable`**: a table of error and status cases. +- **Config validation**: nested chain, `alias` with `failover`, fewer than 2 + targets, missing target, duplicate target, the usecase warning. +- **HTTP integration**: two fake OpenAI-compatible upstreams (`httptest`) behind + `cloud-proxy` target configs. + - Upstream 1 fails. The request is served by upstream 2 with no client error, + `X-LocalAI-Served-Model` is set, and the SSE stream sends `chain.switched`. + - Upstream 1 recovers. Fail-back happens only after N inference probes and + `min_dwell`. + - Upstream 1 fails after the first SSE chunk. There is no retry, the client + gets the error, the target trips, and the next request goes to upstream 2. +- **Handler audit**: one retry test for each endpoint family that proves it is + safe to run again before commit. +- **Realtime**: a pipeline whose LLM stage is a chain of fake backends. The test + checks the `initial` event, the `localai.model.failover` event when the + primary fails, and that the conversation history is kept after the switch. +- **API**: authenticated and unauthenticated access to every endpoint; pin + requires admin. + +## Follow-ups + +1. `localai-proxy` backend: a fork of `cloud-proxy` that forwards every gRPC + method (Predict, Embedding, AudioTranscription, TTS, GenerateImage, Rerank, + VAD, sound detection) to the REST API of an upstream LocalAI. This lets a + remote model serve any realtime pipeline stage. +2. WebUI: a chain editor and a live health view built on + `GET /api/failover/events`. +3. Wingman: use one local LocalAI endpoint with chains for its pipeline stages, + and react to `localai.model.failover` events instead of its own endpoint + supervisor. From d6c6818dc784fe10f6e7609e8f1e1e03c26f1e43 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 14:55:39 +0000 Subject: [PATCH 02/79] docs: probe remote failover targets with /v1/models /readyz exists only on LocalAI upstreams, and a cloud-proxy target can point at any OpenAI-compatible provider. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .../specs/2026-09-26-model-failover-chains-design.md | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 574bc1883..c332082bf 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -150,8 +150,8 @@ Chain states: | Target | Liveness (steady state) | Recovery confirmation | |---|---|---| -| remote | `GET /readyz`, and the upstream model is listed in `GET /v1/models` | one minimal real request, chosen by usecase | -| local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | the same minimal request, run in-process | +| remote | `GET /v1/models` returns 2xx and lists the upstream model. `` is the scheme and host of `proxy.upstream_url` plus any path prefix before `/v1`. The upstream model is `proxy.upstream_model`, or the target name when it is empty. `/v1/models` works on any OpenAI-compatible upstream, and `/readyz` exists only on LocalAI. | one minimal real request, chosen by usecase | +| local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. | | local, cold | the config and model files exist and the backend is installed. The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | Minimal requests by usecase: @@ -173,6 +173,10 @@ Probe load rules: - Remote probes use the URL and API key from the target's proxy config. - Probe results go into the same trip counter as request failures. +When one target is in several chains, its `probe`, `trip` and `recovery` +settings come from the first of those chains in name order. The docs state +this. + ### Warm targets The manager loads `warm: true` targets at startup and marks them pinned in the From f4ce839bdc0aebf9184fcc9dbce76540dc93ef33 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:13:40 +0000 Subject: [PATCH 03/79] docs: plan model failover chains Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .../plans/2026-09-26-model-failover-chains.md | 5054 +++++++++++++++++ 1 file changed, 5054 insertions(+) create mode 100644 docs/superpowers/plans/2026-09-26-model-failover-chains.md diff --git a/docs/superpowers/plans/2026-09-26-model-failover-chains.md b/docs/superpowers/plans/2026-09-26-model-failover-chains.md new file mode 100644 index 000000000..a5a498d53 --- /dev/null +++ b/docs/superpowers/plans/2026-09-26-model-failover-chains.md @@ -0,0 +1,5054 @@ +# Model Failover Chains Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** A model config can declare an ordered `failover` chain of target models; LocalAI serves each request from the highest-priority healthy target, retries uncommitted failures on the next target, probes targets, fails back with hysteresis, and publishes switch events. + +**Architecture:** A new `core/services/failover` package holds a `Manager` (per-target health state machine, per-chain active target, probes, event bus). The HTTP request middleware resolves a chain to a target the way it resolves an alias, and a retry wrapper inside `SetModelAndConfig` re-runs the request on the next target while the response is uncommitted. Realtime pipeline stages resolve chains per call through `Manager.Do`. REST, SSE, a realtime server event, metrics and MCP tools expose the state. + +**Tech Stack:** Go, echo v4, Ginkgo v2 + Gomega, OpenTelemetry metrics, gRPC backend interface (`pkg/grpc`). + +**Spec:** `docs/superpowers/specs/2026-09-26-model-failover-chains-design.md` + +## Global Constraints + +- Worktree: `/home/mudler/_git/LocalAI/.wt/failover-chains`, branch `feat/failover-chains`. All paths below are relative to it. +- Commit trailer: `Assisted-by: Claude:claude-opus-5-5`. Never add `Co-Authored-By` or `Signed-off-by` (the human adds the DCO sign-off). Subjects use the repo's conventional style, for example `feat(failover): ...`. +- `docs/superpowers/` is excluded by the local `.git/info/exclude`; add files there with `git add -f`. +- Logging: `github.com/mudler/xlog`. Use `any`, never `interface{}`. Comments explain why, not what. +- Defaults, exactly: probe interval `15s`, probe timeout `5s`, trip errors `1`, trip window `30s`, recovery probes `3`, min dwell `60s`. +- Response headers, exactly: `X-LocalAI-Served-Model`, `X-LocalAI-Failover` (values `fallback`, `degraded`). +- Event names, exactly: SSE `snapshot`, `chain.switched`, `target.state`; realtime `localai.model.failover`. Reasons: `trip`, `recovery`, `manual`, `degraded`, `missing`, `initial`. +- Metrics, exactly: `localai_failover_switches_total{chain,from,to,reason}`, `localai_failover_target_up{target}`. +- Coverage baseline (`coverage-baseline.txt`, 54.2) must not go down. Never edit the baseline. +- Docs change ships in the same PR (`docs/content/features/model-failover.md`). +- Tests need generated protos: run `make protogen-go` once in the worktree before the first `go test`. E2E tests need `make build-mock-backend`. +- Run a package's tests with: `go run github.com/onsi/ginkgo/v2/ginkgo -v ./` (or `go test .//...`). + +## Review Focus + +1. A handler that mutates the parsed request on attempt 1 must not leak that mutation into attempt 2: each attempt re-binds from the replayed body (Task 8, spec "gives each attempt a fresh request"). +2. A client that disconnects mid-request must not trip the target or trigger a retry (Task 8, spec "does not retry when the client cancelled"). +3. A 4xx (bad request, context overflow) must not retry and must not trip the target (Task 8, spec "does not retry or trip on 4xx"). +4. A degraded chain (all targets down) must try every target in priority order and return the last target's error (Task 8, spec "degraded"). +5. A multipart upload (transcription) must be readable again by the second attempt (Task 8, spec "replays a multipart body"). + +--- + +## File Structure + +| File | Responsibility | +|---|---| +| `core/config/model_config_failover.go` (new) | `FailoverConfig` types, defaults, `IsFailover`, `validateFailover`, `WarmFailoverTargets` | +| `core/config/model_config.go` (modify) | `Failover` field, call `validateFailover` from `Validate`, `ProxyConfig.ResolveAPIKey` | +| `core/config/model_config_loader.go` (modify) | `ValidateFailoverTargets`, load-time pruning and usecase warning | +| `core/config/meta/registry.go`, `types.go` (modify) | `failover` section and field metadata | +| `core/services/failover/types.go` (new) | states, reasons, events, status DTOs, `KindOf`, `MergePinned` | +| `core/services/failover/classify.go` (new) | `IsRetryable` | +| `core/services/failover/manager.go` (new) | `Manager`: sync, state machine, plan/attempt, pin, events, status, `Do` | +| `core/services/failover/schedule.go` (new) | `Prober` interface, `Run`, `Tick`, probe scheduling | +| `core/services/failover/prober.go` (new) | `DefaultProber`: remote HTTP and local gRPC probes | +| `core/services/failover/metrics.go` (new) | OTel counter and gauge | +| `core/services/failover/trace.go` (new) | `RecordAttemptTrace` | +| `core/trace/backend_trace.go` (modify) | `BackendTraceFailover` type | +| `core/application/{application.go,startup.go,watchdog.go,failover.go}` | wiring, warm pinning/preload | +| `core/http/middleware/failover.go` (new) | chain resolution, retry wrapper, writer, headers | +| `core/http/middleware/request.go` (modify) | call resolution, wrap `SetModelAndConfig` | +| `core/http/app.go` (modify) | `SetFailoverManager` | +| `core/http/endpoints/localai/failover.go` (new) | REST + SSE handlers | +| `core/http/routes/localai.go` (modify) | routes | +| `core/http/endpoints/localai/api_instructions.go` (modify) | instruction entry | +| `core/http/endpoints/openai/realtime_model.go`, `realtime.go`, `realtime_failover.go` (new), `types/failover.go` (new), `types/server_events.go` | realtime per-call resolution and events | +| `pkg/mcp/localaitools/...` | MCP tools | +| `tests/e2e/...` | e2e specs, mock-backend load-failure trigger, fake upstream `/v1/models` | +| `docs/content/features/model-failover.md` (new) + cross-links | docs | + +--- + +### Task 1: Config schema, per-config validation, field metadata + +**Files:** +- Create: `core/config/model_config_failover.go` +- Modify: `core/config/model_config.go` (field next to `Alias` at line ~76; `Validate()` at line ~1593, before the alias block at ~1650) +- Modify: `core/config/meta/types.go` (`DefaultSections()`), `core/config/meta/registry.go` (next to the `alias` entry at ~414) +- Test: `core/config/model_config_failover_test.go`, `core/config/meta/registry_test.go` + +**Interfaces:** +- Produces: `config.FailoverConfig`, `config.FailoverTarget`, `(ModelConfig).IsFailover() bool`, `(FailoverConfig).ProbeInterval()/ProbeTimeout()/TripWindow()/MinDwell() time.Duration`, `(FailoverConfig).TripErrors()/RecoveryProbes() int`, `(ModelConfig).WarmFailoverTargets() []string`, field `ModelConfig.Failover *FailoverConfig`. + +- [ ] **Step 1: Write the failing tests** + +`core/config/model_config_failover_test.go` (package `config`, like `model_config_test.go`): + +```go +package config + +import ( + "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"), + ) +}) +``` + +In `core/config/meta/registry_test.go`, next to the alias assertions (lines 13-31), add: + +```go + It("registers the failover section", func() { + reg := meta.DefaultRegistry() + Expect(reg).To(HaveKey("failover.targets")) + Expect(reg["failover.targets"].Section).To(Equal("failover")) + var ids []string + for _, s := range meta.DefaultSections() { + ids = append(ids, s.ID) + } + Expect(ids).To(ContainElement("failover")) + }) +``` + +(Match the surrounding style: if the file uses `It` inside an existing `Describe`, put this `It` there.) + +- [ ] **Step 2: Run the tests to verify they fail** + +Run: `go test ./core/config/... 2>&1 | tail -20` +Expected: compile failure, `undefined: FailoverConfig`. + +- [ ] **Step 3: Implement the config types** + +`core/config/model_config_failover.go`: + +```go +package config + +import ( + "fmt" + "time" +) + +// 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 } + +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 +} +``` + +In `core/config/model_config.go`, add the field directly after `Alias` (line ~76): + +```go + // 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"` +``` + +In `Validate()`, directly before the `if c.IsAlias() {` block (line ~1650), add: + +```go + if c.IsFailover() { + if err := c.validateFailover(); err != nil { + return false, err + } + return true, nil + } +``` + +If `Validate()` rejects configs without a backend earlier than line 1650 (the artifact check at ~1619 applies to aliases), move this block above that point so a chain returns before any backend-specific check. + +- [ ] **Step 4: Register field metadata** + +In `core/config/meta/types.go` `DefaultSections()`, add after the `alias` line: + +```go + {ID: "failover", Label: "Failover", Icon: "shuffle", Order: 6}, +``` + +(If `shuffle` is not an icon used elsewhere in `DefaultSections()`, reuse `git-merge`.) + +In `core/config/meta/registry.go`, after the `alias` entry, add: + +```go + // --- 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: "json-editor", + 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", + }, +``` + +Run `go test ./core/config/meta/...`. `TestAllFieldsHaveRegistryEntries` lists every reflected path that has no entry. If it names paths different from the seven above (for example `failover.targets[].model`), register those exact paths with matching labels in the same section, and remove entries it reports as unknown. + +- [ ] **Step 5: Run the tests to verify they pass** + +Run: `go test ./core/config/... 2>&1 | tail -20` +Expected: PASS. + +- [ ] **Step 6: Commit** + +```bash +git add core/config +git commit -m "feat(config): add failover chain block to model configs + +A chain is a model config with an ordered list of target models, probe, +trip and recovery settings. Like an alias it has no backend. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 2: Cross-config validation in the loader and admin paths + +**Files:** +- Modify: `core/config/model_config_loader.go` (new methods next to `ValidateAliasTarget` at ~487; load-time check after the alias check at ~824-840) +- Modify: `core/services/modeladmin/config.go` (next to the `ValidateAliasTarget` calls at ~169 and ~286), `core/http/endpoints/localai/import_model.go` (~186) +- Test: `core/config/model_config_loader_test.go` + +**Interfaces:** +- Consumes: Task 1 types. +- Produces: `(*ModelConfigLoader).ValidateFailoverTargets(cfg *ModelConfig) error`, `(*ModelConfigLoader).FailoverTargetsShareUsecase(cfg *ModelConfig) bool`. + +- [ ] **Step 1: Write the failing tests** + +Append to `core/config/model_config_loader_test.go` (the file seeds `loader.configs` directly, see line ~304): + +```go +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()) + }) +}) +``` + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/config/... 2>&1 | tail -5` +Expected: `undefined: ... ValidateFailoverTargets`. + +- [ ] **Step 3: Implement** + +In `core/config/model_config_loader.go`: + +```go +// 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 +} + +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 +} +``` + +In `loadModelConfigsFromPath`, directly after the alias warning loop (~824-840), add a pass over `bcl.configs`. Check whether that code runs with `bcl.Mutex` held: if it does, pass a lock-free lookup (`func(n string) (ModelConfig, bool) { c, ok := bcl.configs[n]; return c, ok }`) instead of `bcl.GetModelConfig`, otherwise it deadlocks. + +```go + 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) + } + } +``` + +Use the strict-mode variable already in scope in that function (it is the one used at line ~813 for `invalid model config`). If the function has no strict flag at that point, drop the `if strict` branch. + +In `core/services/modeladmin/config.go` (both sites) and `core/http/endpoints/localai/import_model.go`, directly after each `ValidateAliasTarget(...)` call, add the equivalent call and return the error the same way the alias error is returned: + +```go + if err := s.Loader.ValidateFailoverTargets(&cfg); err != nil { + return /* same error shape as the ValidateAliasTarget branch above */ err + } +``` + +Copy the exact receiver and variable names from the adjacent `ValidateAliasTarget` call. Do not change what the alias branch returns. + +- [ ] **Step 4: Run tests** + +Run: `go test ./core/config/... ./core/services/modeladmin/... 2>&1 | tail -10` +Expected: PASS. + +- [ ] **Step 5: Commit** + +```bash +git add core/config core/services/modeladmin core/http/endpoints/localai/import_model.go +git commit -m "feat(config): validate failover chain targets across configs + +Reject chains whose targets are missing or are chains, at load and on +create or edit, and warn when the targets share no usecase. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 3: Failover package types and error classification + +**Files:** +- Create: `core/services/failover/failover_suite_test.go`, `types.go`, `classify.go`, `classify_test.go`, `types_test.go` + +**Interfaces:** +- Produces: types `TargetState`, `ChainState`, `Kind`, `Reason`, `EventType`, `Event`, `TargetStatus`, `ChainStatus`; constants listed below; `KindOf(cfg config.ModelConfig) Kind`; `MergePinned(pinned, warm []string) []string`; `IsRetryable(err error, status int) bool`. + +- [ ] **Step 1: Write the failing tests** + +`core/services/failover/failover_suite_test.go`: + +```go +package failover + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestFailover(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Failover test suite") +} +``` + +`core/services/failover/classify_test.go`: + +```go +package failover + +import ( + "context" + "errors" + "fmt" + "net/http" + + "github.com/labstack/echo/v4" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +var _ = DescribeTable("IsRetryable", + func(err error, status int, want bool) { + Expect(IsRetryable(err, status)).To(Equal(want)) + }, + Entry("nil error, no status", nil, 0, false), + Entry("held 503", nil, http.StatusServiceUnavailable, true), + Entry("held 500", nil, http.StatusInternalServerError, true), + Entry("held 501", nil, http.StatusNotImplemented, false), + Entry("client cancel", context.Canceled, 0, false), + Entry("wrapped client cancel", fmt.Errorf("predict: %w", context.Canceled), 0, false), + Entry("deadline", context.DeadlineExceeded, 0, true), + Entry("echo 502", echo.NewHTTPError(http.StatusBadGateway, "x"), 0, true), + Entry("echo 400", echo.NewHTTPError(http.StatusBadRequest, "x"), 0, false), + Entry("echo 404", echo.NewHTTPError(http.StatusNotFound, "x"), 0, false), + Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), 0, true), + Entry("grpc internal", grpcstatus.Error(codes.Internal, "x"), 0, true), + Entry("grpc deadline", grpcstatus.Error(codes.DeadlineExceeded, "x"), 0, true), + Entry("grpc unknown", grpcstatus.Error(codes.Unknown, "x"), 0, true), + Entry("grpc invalid argument", grpcstatus.Error(codes.InvalidArgument, "x"), 0, false), + Entry("cloud-proxy upstream 503", errors.New("cloud-proxy: upstream 503: no healthy nodes"), 0, true), + Entry("cloud-proxy upstream 429 stays 4xx", errors.New("cloud-proxy: upstream 429: slow down"), 0, false), + Entry("context overflow", errors.New("the request exceeds the available context size"), 0, false), + Entry("dial error", errors.New("dial tcp 10.0.0.1:8080: connect: connection refused"), 0, true), +) +``` + +`core/services/failover/types_test.go`: + +```go +package failover + +import ( + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("KindOf", func() { + It("treats proxy backends as remote", func() { + Expect(KindOf(config.ModelConfig{Backend: "cloud-proxy"})).To(Equal(KindRemote)) + Expect(KindOf(config.ModelConfig{Backend: "localai-proxy"})).To(Equal(KindRemote)) + Expect(KindOf(config.ModelConfig{Backend: "llama-cpp"})).To(Equal(KindLocal)) + }) +}) + +var _ = Describe("MergePinned", func() { + It("adds warm targets without duplicates", func() { + Expect(MergePinned([]string{"a", "b"}, []string{"b", "c"})).To(Equal([]string{"a", "b", "c"})) + Expect(MergePinned(nil, nil)).To(BeEmpty()) + }) +}) +``` + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: compile failure (`undefined: IsRetryable`). + +- [ ] **Step 3: Implement** + +`core/services/failover/types.go`: + +```go +// Package failover serves a model name from an ordered chain of target +// models, moving to the next target when one fails and back when it +// recovers. +package failover + +import ( + "slices" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +type TargetState string + +const ( + StateHealthy TargetState = "healthy" + StateDown TargetState = "down" + StateRecovering TargetState = "recovering" + StateMissing TargetState = "missing" +) + +type ChainState string + +const ( + ChainPrimary ChainState = "primary" + ChainFallback ChainState = "fallback" + ChainDegraded ChainState = "degraded" +) + +type Kind string + +const ( + KindLocal Kind = "local" + KindRemote Kind = "remote" +) + +type Reason string + +const ( + ReasonTrip Reason = "trip" + ReasonRecovery Reason = "recovery" + ReasonManual Reason = "manual" + ReasonDegraded Reason = "degraded" + ReasonMissing Reason = "missing" + ReasonInitial Reason = "initial" +) + +type EventType string + +const ( + EventChainSwitched EventType = "chain.switched" + EventTargetState EventType = "target.state" +) + +// Event is one change of a target state or of a chain's active target. +type Event struct { + Type EventType `json:"type"` + Chain string `json:"chain,omitempty"` + Target string `json:"target,omitempty"` + From string `json:"from"` + To string `json:"to"` + State string `json:"state,omitempty"` + Reason Reason `json:"reason"` + Error string `json:"error,omitempty"` + At time.Time `json:"at"` +} + +type TargetStatus struct { + Model string `json:"model"` + Kind Kind `json:"kind"` + Warm bool `json:"warm"` + State TargetState `json:"state"` + ConsecutiveOK int `json:"consecutive_ok"` + LastProbe *time.Time `json:"last_probe,omitempty"` + LastError string `json:"last_error,omitempty"` +} + +type ChainStatus struct { + Name string `json:"name"` + State ChainState `json:"state"` + Active string `json:"active"` + ActiveSince time.Time `json:"active_since"` + Pinned *string `json:"pinned"` + Targets []TargetStatus `json:"targets"` +} + +// KindOf decides how a target is probed: proxy backends forward to another +// server and are checked over HTTP, everything else runs in this instance. +func KindOf(cfg config.ModelConfig) Kind { + switch cfg.Backend { + case "cloud-proxy", "localai-proxy": + return KindRemote + } + return KindLocal +} + +// MergePinned adds warm failover targets to the config-pinned model list, so +// the watchdog never evicts them. +func MergePinned(pinned, warm []string) []string { + out := slices.Clone(pinned) + for _, w := range warm { + if !slices.Contains(out, w) { + out = append(out, w) + } + } + return out +} +``` + +`core/services/failover/classify.go`: + +```go +package failover + +import ( + "context" + "errors" + "net/http" + "regexp" + "strconv" + "strings" + + "github.com/labstack/echo/v4" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +// cloud-proxy translate mode reports upstream failures as plain text, so the +// status is only visible in the message. +var upstreamStatusRe = regexp.MustCompile(`upstream (\d{3})`) + +// Errors that the next target would reject in the same way. +var requestErrorMarkers = []string{ + "exceeds the available context size", + "is larger than the max context size", + "maximum context length", +} + +// IsRetryable reports whether a failed attempt should move to the next +// target. status is the HTTP status a handler wrote, or 0 when it returned err +// without writing. +func IsRetryable(err error, status int) bool { + if errors.Is(err, context.Canceled) { + return false + } + if status != 0 { + return retryableStatus(status) + } + if err == nil { + return false + } + var he *echo.HTTPError + if errors.As(err, &he) { + return retryableStatus(he.Code) + } + if errors.Is(err, context.DeadlineExceeded) { + return true + } + if st, ok := grpcstatus.FromError(err); ok { + switch st.Code() { + case codes.Unavailable, codes.Internal, codes.DeadlineExceeded, codes.Unknown: + return !isRequestError(st.Message()) + default: + return false + } + } + msg := err.Error() + if m := upstreamStatusRe.FindStringSubmatch(msg); m != nil { + code, _ := strconv.Atoi(m[1]) + return retryableStatus(code) + } + // Anything else is usually a dial or load failure of this target. + return !isRequestError(msg) +} + +func retryableStatus(code int) bool { + return code >= 500 && code != http.StatusNotImplemented +} + +func isRequestError(msg string) bool { + for _, m := range requestErrorMarkers { + if strings.Contains(msg, m) { + return true + } + } + return false +} +``` + +- [ ] **Step 4: Run tests** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: PASS. + +- [ ] **Step 5: Commit** + +```bash +git add core/services/failover +git commit -m "feat(failover): add types and retryable error classification + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 4: Manager state machine, plans, pins, events + +**Files:** +- Create: `core/services/failover/manager.go`, `core/services/failover/manager_test.go`, `core/services/failover/fakes_test.go` + +**Interfaces:** +- Consumes: Task 1 config types, Task 3 types and `IsRetryable`. +- Produces (used by Tasks 5-12): + - `type ConfigSource interface { GetModelConfig(string) (config.ModelConfig, bool); GetAllModelsConfigs() []config.ModelConfig }` + - `type Clock interface { Now() time.Time }` + - `func New(src ConfigSource, opts ...Option) *Manager`; options `WithClock(Clock)`, `WithProber(Prober)`, `WithOnWarmChanged(func([]string))`; the `Prober` interface (declared at the end of `manager.go`, implemented in Task 6) + - `(*Manager).Sync()`, `Reevaluate()`, `Plan(chain string) (*Attempt, error)`, `ReportFailure(target string, err error)`, `ReportSuccess(target string)`, `Pin(chain, target string) error`, `Unpin(chain string) error`, `Status() []ChainStatus`, `ChainStatus(name string) (ChainStatus, bool)`, `Subscribe(buffer int) (<-chan Event, func())`, `WarmTargets() []string`, `Do(ctx, chain string, fn func(ctx context.Context, target string, commit func()) error) error` + - `(*Attempt).Chain() string`, `Target() string`, `Primary() string`, `Degraded() bool`, `Fail(err error) bool`, `Report(err error)`, `Succeed()` + - errors `ErrChainNotFound`, `ErrTargetNotInChain`, `ErrNoTarget` + +- [ ] **Step 1: Write test fakes** + +`core/services/failover/fakes_test.go`: + +```go +package failover + +import ( + "sort" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +type fakeClock struct { + mu sync.Mutex + now time.Time +} + +func newFakeClock() *fakeClock { return &fakeClock{now: time.Date(2026, 9, 26, 10, 0, 0, 0, time.UTC)} } +func (c *fakeClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.now +} +func (c *fakeClock) Advance(d time.Duration) { + c.mu.Lock() + c.now = c.now.Add(d) + c.mu.Unlock() +} + +type fakeSource struct { + mu sync.Mutex + cfgs map[string]config.ModelConfig +} + +func newFakeSource(cfgs ...config.ModelConfig) *fakeSource { + s := &fakeSource{cfgs: map[string]config.ModelConfig{}} + for _, c := range cfgs { + s.cfgs[c.Name] = c + } + return s +} +func (s *fakeSource) Put(c config.ModelConfig) { s.mu.Lock(); s.cfgs[c.Name] = c; s.mu.Unlock() } +func (s *fakeSource) Delete(name string) { s.mu.Lock(); delete(s.cfgs, name); s.mu.Unlock() } +func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { + s.mu.Lock() + defer s.mu.Unlock() + c, ok := s.cfgs[n] + return c, ok +} +func (s *fakeSource) GetAllModelsConfigs() []config.ModelConfig { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]config.ModelConfig, 0, len(s.cfgs)) + for _, c := range s.cfgs { + out = append(out, c) + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out +} + +func local(name string) config.ModelConfig { return config.ModelConfig{Name: name, Backend: "llama-cpp"} } +func remote(name string) config.ModelConfig { return config.ModelConfig{Name: name, Backend: "cloud-proxy"} } + +// chainCfg builds a chain; fc may be nil for defaults. +func chainCfg(name string, fc *config.FailoverConfig, targets ...config.FailoverTarget) config.ModelConfig { + f := config.FailoverConfig{} + if fc != nil { + f = *fc + } + f.Targets = targets + return config.ModelConfig{Name: name, Failover: &f} +} + +func t(model string) config.FailoverTarget { return config.FailoverTarget{Model: model} } +func warmT(model string) config.FailoverTarget { return config.FailoverTarget{Model: model, Warm: true} } + +// drain returns the events buffered so far without blocking. +func drain(ch <-chan Event) []Event { + var out []Event + for { + select { + case ev, ok := <-ch: + if !ok { + return out + } + out = append(out, ev) + default: + return out + } + } +} +``` + +- [ ] **Step 2: Write the failing manager tests** + +`core/services/failover/manager_test.go`: + +```go +package failover + +import ( + "context" + "errors" + "time" + + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var errBoom = errors.New("dial tcp: connection refused") + +var _ = Describe("Manager", func() { + var ( + clock *fakeClock + src *fakeSource + m *Manager + ) + + BeforeEach(func() { + clock = newFakeClock() + src = newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m = New(src, WithClock(clock)) + }) + + switched := func(evs []Event) []Event { + var out []Event + for _, e := range evs { + if e.Type == EventChainSwitched { + out = append(out, e) + } + } + return out + } + + It("plans the primary first on a fresh chain", func() { + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Primary()).To(Equal("a")) + Expect(att.Degraded()).To(BeFalse()) + st, ok := m.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(st.State).To(Equal(ChainPrimary)) + Expect(st.Targets[0].Kind).To(Equal(KindRemote)) + Expect(st.Targets[1].Kind).To(Equal(KindLocal)) + }) + + It("returns ErrChainNotFound for an unknown chain", func() { + _, err := m.Plan("nope") + Expect(errors.Is(err, ErrChainNotFound)).To(BeTrue()) + }) + + It("trips on the first failure by default and switches with an event", func() { + events, cancel := m.Subscribe(16) + defer cancel() + att, _ := m.Plan("chain") + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(st.State).To(Equal(ChainFallback)) + Expect(st.Targets[0].State).To(Equal(StateDown)) + Expect(st.Targets[0].LastError).To(ContainSubstring("connection refused")) + sw := switched(drain(events)) + Expect(sw).To(HaveLen(1)) + Expect(sw[0]).To(MatchFields(IgnoreExtras, Fields{ + "Chain": Equal("chain"), "From": Equal("a"), "To": Equal("b"), + "State": Equal("fallback"), "Reason": Equal(ReasonTrip), + })) + }) + + It("counts failures inside the trip window only", func() { + src.Put(chainCfg("chain", &config.FailoverConfig{Trip: config.FailoverTrip{Errors: 2, Window: "30s"}}, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + clock.Advance(31 * time.Second) + m.ReportFailure("a", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + clock.Advance(time.Second) + m.ReportFailure("a", errBoom) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("fails back only after recovery probes and min_dwell", func() { + m.Plan("chain") + m.ReportFailure("a", errBoom) + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a real success counts like a passed inference probe + } + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Active).To(Equal("b"), "min_dwell has not passed") + clock.Advance(61 * time.Second) + events, cancel := m.Subscribe(16) + defer cancel() + m.Reevaluate() + st, _ = m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + Expect(switched(drain(events))[0].Reason).To(Equal(ReasonRecovery)) + }) + + It("moves up at once when the active target itself goes down", func() { + src.Put(chainCfg("chain", nil, t("a"), t("b"), t("c"))) + src.Put(local("c")) + m.Sync() + m.ReportFailure("a", errBoom) // active: b + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a healthy again, but dwell not passed + } + m.ReportFailure("b", errBoom) // b down: go to a now, not c + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + }) + + It("goes degraded when all targets are down and plans all of them in priority order", func() { + m.ReportFailure("a", errBoom) + m.ReportFailure("b", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.State).To(Equal(ChainDegraded)) + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Degraded()).To(BeTrue()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("pins a target regardless of health", func() { + Expect(m.Pin("chain", "b")).To(Succeed()) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse(), "a pin allows only the pinned target") + st, _ := m.ChainStatus("chain") + Expect(*st.Pinned).To(Equal("b")) + Expect(st.Active).To(Equal("b")) + Expect(m.Unpin("chain")).To(Succeed()) + st, _ = m.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + Expect(errors.Is(m.Pin("chain", "zzz"), ErrTargetNotInChain)).To(BeTrue()) + Expect(errors.Is(m.Pin("nope", "a"), ErrChainNotFound)).To(BeTrue()) + }) + + It("shares target health across chains", func() { + src.Put(chainCfg("chain2", nil, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + s1, _ := m.ChainStatus("chain") + s2, _ := m.ChainStatus("chain2") + Expect(s1.Active).To(Equal("b")) + Expect(s2.Active).To(Equal("b")) + }) + + It("marks a removed target missing and leaves it out of plans", func() { + m.Plan("chain") + src.Delete("a") + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateMissing)) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("resets a chain whose target list changed", func() { + m.ReportFailure("a", errBoom) + src.Put(local("c")) + src.Put(chainCfg("chain", nil, t("c"), t("b"))) + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("c")) + }) + + It("reports warm local targets and ignores warm on remote ones", func() { + var got []string + m = New(src, WithClock(clock), WithOnWarmChanged(func(w []string) { got = w })) + src.Put(chainCfg("chain", nil, warmT("a"), warmT("b"))) + m.Sync() + Expect(got).To(Equal([]string{"b"})) + Expect(m.WarmTargets()).To(Equal([]string{"b"})) + }) + + It("closes a subscription on cancel", func() { + events, cancel := m.Subscribe(1) + cancel() + _, ok := <-events + Expect(ok).To(BeFalse()) + cancel() // idempotent + }) + + Describe("Do", func() { + It("retries on the next target until commit", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + if target == "a" { + return errBoom + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + }) + + It("does not retry after commit but still trips the target", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + commit() + return errBoom + }) + Expect(err).To(MatchError(errBoom)) + Expect(tried).To(Equal([]string{"a"})) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("does not retry or trip on a non-retryable error", func() { + bad := errors.New("the request exceeds the available context size") + err := m.Do(context.Background(), "chain", func(_ context.Context, _ string, _ func()) error { return bad }) + Expect(err).To(MatchError(bad)) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + }) +}) +``` + +The tests use `MatchFields`, so add `. "github.com/onsi/gomega/gstruct"` to the imports. + +- [ ] **Step 3: Run to verify failure** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: compile failure (`undefined: New`). + +- [ ] **Step 4: Implement the manager** + +`core/services/failover/manager.go`: + +```go +package failover + +import ( + "context" + "errors" + "fmt" + "slices" + "sort" + "sync" + "sync/atomic" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/xlog" +) + +var ( + ErrChainNotFound = errors.New("failover chain not found") + ErrTargetNotInChain = errors.New("target is not in this failover chain") + ErrNoTarget = errors.New("failover chain has no usable target") +) + +// ConfigSource is the part of ModelConfigLoader the manager reads. +type ConfigSource interface { + GetModelConfig(name string) (config.ModelConfig, bool) + GetAllModelsConfigs() []config.ModelConfig +} + +type Clock interface{ Now() time.Time } + +type realClock struct{} + +func (realClock) Now() time.Time { return time.Now() } + +type Option func(*Manager) + +func WithClock(c Clock) Option { return func(m *Manager) { m.clock = c } } +func WithProber(p Prober) Option { return func(m *Manager) { m.prober = p } } + +// WithOnWarmChanged is called outside the manager lock when the set of warm +// local targets changes. The application pins and preloads them. +func WithOnWarmChanged(fn func(warm []string)) Option { return func(m *Manager) { m.onWarm = fn } } + +// Manager tracks health per target and the active target per chain. +type Manager struct { + mu sync.Mutex + src ConfigSource + clock Clock + prober Prober + onWarm func([]string) + targets map[string]*targetState + chains map[string]*chainState + subs map[int]chan Event + nextSub int + warm []string + warmPending bool + closed bool +} + +type targetState struct { + name string + kind Kind + warm bool + state TargetState + failures []time.Time + consecutiveOK int + downSince time.Time + lastProbe time.Time + lastActivity time.Time + lastError string + // params come from the first chain, in name order, that lists the target. + params config.FailoverConfig +} + +func (ts *targetState) cold() bool { return ts.kind == KindLocal && !ts.warm } + +type chainState struct { + name string + cfg config.FailoverConfig + targets []string + active int + activeSince time.Time + pinned string + state ChainState +} + +func New(src ConfigSource, opts ...Option) *Manager { + m := &Manager{ + src: src, + clock: realClock{}, + targets: map[string]*targetState{}, + chains: map[string]*chainState{}, + subs: map[int]chan Event{}, + } + for _, o := range opts { + o(m) + } + return m +} + +// Sync reconciles chains with the config source. There is no config-change +// hook in the loader, so this runs on every tick and on a lookup miss. +func (m *Manager) Sync() { + m.mu.Lock() + m.syncLocked() + warm, deliver := m.takeWarmLocked() + m.mu.Unlock() + if deliver && m.onWarm != nil { + m.onWarm(warm) + } +} + +func (m *Manager) syncLocked() { + now := m.clock.Now() + seenChains := map[string]bool{} + claimed := map[string]bool{} + for _, c := range m.src.GetAllModelsConfigs() { + if !c.IsFailover() { + continue + } + seenChains[c.Name] = true + names := make([]string, 0, len(c.Failover.Targets)) + for _, t := range c.Failover.Targets { + names = append(names, t.Model) + } + ch := m.chains[c.Name] + if ch == nil || !slices.Equal(ch.targets, names) { + pinned := "" + if ch != nil && slices.Contains(names, ch.pinned) { + pinned = ch.pinned + } + ch = &chainState{name: c.Name, targets: names, activeSince: now, state: ChainPrimary, pinned: pinned} + m.chains[c.Name] = ch + } + ch.cfg = *c.Failover + for _, t := range c.Failover.Targets { + ts := m.targets[t.Model] + if ts == nil { + ts = &targetState{name: t.Model, state: StateHealthy} + m.targets[t.Model] = ts + } + if !claimed[t.Model] { + claimed[t.Model] = true + ts.params = *c.Failover + ts.warm = false + } + tc, ok := m.lookupTarget(t.Model) + if !ok { + m.setTargetLocked(ts, StateMissing, ReasonMissing, "target config not found") + continue + } + ts.kind = KindOf(tc) + if t.Warm && ts.kind == KindLocal { + ts.warm = true + } + if ts.state == StateMissing { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } + } + } + for name := range m.chains { + if !seenChains[name] { + delete(m.chains, name) + } + } + for name := range m.targets { + if !claimed[name] { + delete(m.targets, name) + } + } + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } + var warm []string + for name, ts := range m.targets { + if ts.warm { + warm = append(warm, name) + } + } + sort.Strings(warm) + if !slices.Equal(warm, m.warm) { + m.warm = warm + m.warmPending = true + } +} + +func (m *Manager) takeWarmLocked() ([]string, bool) { + if !m.warmPending { + return nil, false + } + m.warmPending = false + return slices.Clone(m.warm), true +} + +// lookupTarget returns the config that serves a target, one alias hop deep. +func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { + c, ok := m.src.GetModelConfig(name) + if ok && c.IsAlias() { + return m.src.GetModelConfig(c.Alias) + } + return c, ok +} + +func (m *Manager) chainLocked(name string) *chainState { + if ch := m.chains[name]; ch != nil { + return ch + } + m.syncLocked() + return m.chains[name] +} + +// WarmTargets returns the warm local targets, sorted. +func (m *Manager) WarmTargets() []string { + m.mu.Lock() + defer m.mu.Unlock() + return slices.Clone(m.warm) +} + +// Reevaluate recomputes every chain. Dwell-based fail-back needs no event, so +// the scheduler calls this on every tick. +func (m *Manager) Reevaluate() { + m.mu.Lock() + defer m.mu.Unlock() + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } +} + +func (m *Manager) setTargetLocked(ts *targetState, to TargetState, reason Reason, errMsg string) { + if ts.state == to { + return + } + from := ts.state + now := m.clock.Now() + ts.state = to + switch to { + case StateDown: + ts.downSince = now + ts.consecutiveOK = 0 + ts.failures = nil + case StateRecovering, StateHealthy: + ts.consecutiveOK = 0 + ts.failures = nil + } + m.emitLocked(Event{Type: EventTargetState, Target: ts.name, From: string(from), To: string(to), Reason: reason, Error: errMsg, At: now}) +} + +// recomputeLocked picks the active target. override replaces the reason of a +// resulting switch (pin and unpin are always "manual"). +func (m *Manager) recomputeLocked(ch *chainState, override Reason) { + now := m.clock.Now() + prev := ch.active + next := prev + reason := ReasonTrip + best := -1 + for i, name := range ch.targets { + if ts := m.targets[name]; ts != nil && ts.state == StateHealthy { + best = i + break + } + } + switch { + case ch.pinned != "": + next = slices.Index(ch.targets, ch.pinned) + reason = ReasonManual + case best == -1: + // Nothing is healthy: keep the active target, Plan tries all of them. + case best > prev: + next = best // the active target is not healthy + case best < prev: + cur := m.targets[ch.targets[prev]] + curHealthy := cur != nil && cur.state == StateHealthy + if !curHealthy { + next = best + } else if now.Sub(ch.activeSince) >= ch.cfg.MinDwell() { + next = best + reason = ReasonRecovery + } + } + if override != "" { + reason = override + } + var state ChainState + switch { + case ch.pinned == "" && best == -1: + state = ChainDegraded + case next == 0: + state = ChainPrimary + default: + state = ChainFallback + } + switch { + case next != prev: + ch.active = next + ch.activeSince = now + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: reason, At: now}) + case state == ChainDegraded && ch.state != ChainDegraded: + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonDegraded, At: now}) + } + ch.state = state +} + +func (m *Manager) recomputeForLocked(target string) { + for _, ch := range m.chains { + if slices.Contains(ch.targets, target) { + m.recomputeLocked(ch, "") + } + } +} + +// Attempt walks the targets of one request in order. +type Attempt struct { + m *Manager + chain string + primary string + degraded bool + targets []string + i int +} + +// Plan returns the attempt order for one request: the active target, then the +// other healthy targets. A degraded chain tries every target in priority +// order; a pinned chain only the pinned target. +func (m *Manager) Plan(chain string) (*Attempt, error) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return nil, fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + att := &Attempt{m: m, chain: ch.name, primary: ch.targets[0], degraded: ch.state == ChainDegraded} + usable := func(name string) bool { + ts := m.targets[name] + if ts == nil || ts.state == StateMissing { + return false + } + return att.degraded || ts.state == StateHealthy + } + switch { + case ch.pinned != "": + att.targets = []string{ch.pinned} + case att.degraded: + for _, name := range ch.targets { + if usable(name) { + att.targets = append(att.targets, name) + } + } + default: + active := ch.targets[ch.active] + if usable(active) { + att.targets = append(att.targets, active) + } + for _, name := range ch.targets { + if name != active && usable(name) { + att.targets = append(att.targets, name) + } + } + } + if len(att.targets) == 0 { + return nil, fmt.Errorf("%w: %q", ErrNoTarget, chain) + } + return att, nil +} + +func (a *Attempt) Chain() string { return a.chain } +func (a *Attempt) Target() string { return a.targets[a.i] } +func (a *Attempt) Primary() string { return a.primary } +func (a *Attempt) Degraded() bool { return a.degraded } + +// Fail records err against the current target and moves to the next one. It +// returns false when no target is left. +func (a *Attempt) Fail(err error) bool { + a.m.ReportFailure(a.Target(), err) + if a.i+1 >= len(a.targets) { + return false + } + a.i++ + return true +} + +// Report records err against the current target without moving on: the +// response was already committed, so nothing is left to retry. +func (a *Attempt) Report(err error) { a.m.ReportFailure(a.Target(), err) } + +func (a *Attempt) Succeed() { a.m.ReportSuccess(a.Target()) } + +func (m *Manager) ReportFailure(target string, err error) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targets[target] + if ts == nil { + return + } + msg := "" + if err != nil { + msg = err.Error() + } + m.recordFailureLocked(ts, msg) + m.recomputeForLocked(target) +} + +func (m *Manager) recordFailureLocked(ts *targetState, msg string) { + now := m.clock.Now() + ts.lastError = msg + switch ts.state { + case StateRecovering: + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + case StateHealthy: + cut := now.Add(-ts.params.TripWindow()) + kept := ts.failures[:0] + for _, f := range ts.failures { + if f.After(cut) { + kept = append(kept, f) + } + } + ts.failures = append(kept, now) + if len(ts.failures) >= ts.params.TripErrors() { + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + } + } +} + +func (m *Manager) ReportSuccess(target string) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targets[target] + if ts == nil { + return + } + ts.lastActivity = m.clock.Now() + m.recordPassLocked(ts) + m.recomputeForLocked(target) +} + +// recordPassLocked counts a served request or a passed inference probe. +func (m *Manager) recordPassLocked(ts *targetState) { + switch ts.state { + case StateHealthy: + ts.failures = nil + return + case StateMissing: + return + case StateDown: + if ts.cold() { + // Cold targets are never probed; a served request is proof enough. + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + return + } + m.setTargetLocked(ts, StateRecovering, ReasonRecovery, "") + } + ts.consecutiveOK++ + if ts.consecutiveOK >= ts.params.RecoveryProbes() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } +} + +func (m *Manager) Pin(chain, target string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + if !slices.Contains(ch.targets, target) { + return fmt.Errorf("%w: %q", ErrTargetNotInChain, target) + } + ch.pinned = target + m.recomputeLocked(ch, ReasonManual) + return nil +} + +func (m *Manager) Unpin(chain string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + ch.pinned = "" + m.recomputeLocked(ch, ReasonManual) + return nil +} + +// Status returns every chain, sorted by name. +func (m *Manager) Status() []ChainStatus { + m.mu.Lock() + defer m.mu.Unlock() + m.syncLocked() + names := make([]string, 0, len(m.chains)) + for name := range m.chains { + names = append(names, name) + } + sort.Strings(names) + out := make([]ChainStatus, 0, len(names)) + for _, name := range names { + out = append(out, m.statusLocked(m.chains[name])) + } + return out +} + +func (m *Manager) ChainStatus(name string) (ChainStatus, bool) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(name) + if ch == nil { + return ChainStatus{}, false + } + return m.statusLocked(ch), true +} + +func (m *Manager) statusLocked(ch *chainState) ChainStatus { + cs := ChainStatus{Name: ch.name, State: ch.state, Active: ch.targets[ch.active], ActiveSince: ch.activeSince} + if ch.pinned != "" { + p := ch.pinned + cs.Pinned = &p + } + for _, name := range ch.targets { + st := TargetStatus{Model: name} + if ts := m.targets[name]; ts != nil { + st.Kind, st.Warm, st.State = ts.kind, ts.warm, ts.state + st.ConsecutiveOK, st.LastError = ts.consecutiveOK, ts.lastError + if !ts.lastProbe.IsZero() { + lp := ts.lastProbe + st.LastProbe = &lp + } + } + cs.Targets = append(cs.Targets, st) + } + return cs +} + +// Subscribe returns a buffered event channel and a cancel func. A subscriber +// that does not keep up loses events rather than blocking the manager. +func (m *Manager) Subscribe(buffer int) (<-chan Event, func()) { + m.mu.Lock() + defer m.mu.Unlock() + ch := make(chan Event, buffer) + if m.closed { + close(ch) + return ch, func() {} + } + id := m.nextSub + m.nextSub++ + m.subs[id] = ch + var once sync.Once + return ch, func() { + once.Do(func() { + m.mu.Lock() + defer m.mu.Unlock() + if c, ok := m.subs[id]; ok { + delete(m.subs, id) + close(c) + } + }) + } +} + +func (m *Manager) emitLocked(ev Event) { + for _, c := range m.subs { + select { + case c <- ev: + default: + xlog.Warn("failover: dropping event for a slow subscriber", "type", ev.Type, "chain", ev.Chain, "target", ev.Target) + } + } +} + +func (m *Manager) close() { + m.mu.Lock() + defer m.mu.Unlock() + m.closed = true + for id, c := range m.subs { + close(c) + delete(m.subs, id) + } +} + +// Do runs fn against the chain's targets in plan order. fn calls commit once +// output has reached the client; after that a failure is not retried. +func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Context, target string, commit func()) error) error { + att, err := m.Plan(chain) + if err != nil { + return err + } + for { + var committed atomic.Bool + err := fn(ctx, att.Target(), func() { committed.Store(true) }) + switch { + case err == nil: + att.Succeed() + return nil + case ctx.Err() != nil || !IsRetryable(err, 0): + return err + case committed.Load(): + att.Report(err) + return err + case !att.Fail(err): + return err + } + } +} + +// Prober checks targets. Implemented by DefaultProber (prober.go). +type Prober interface { + // Liveness is the cheap steady-state check. + Liveness(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error + // Inference sends one minimal real request to confirm recovery. + Inference(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error +} +``` + +- [ ] **Step 5: Run tests** + +Run: `go test -race ./core/services/failover/... 2>&1 | tail -10` +Expected: PASS, no race reports. + +- [ ] **Step 6: Commit** + +```bash +git add core/services/failover +git commit -m "feat(failover): add chain manager with trip, fail-back and pins + +Health is tracked per target and the active target per chain. Fail-back +waits for recovery probes and a minimum time on the fallback. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 5: Probe scheduling + +**Files:** +- Create: `core/services/failover/schedule.go`, `core/services/failover/schedule_test.go` + +**Interfaces:** +- Consumes: Task 4 `Manager`, `Prober`, internal helpers `recordFailureLocked`, `recordPassLocked`, `setTargetLocked`, `recomputeForLocked`, `lookupTarget`, `close`. +- Produces: `(*Manager).Run(ctx)`, `(*Manager).Tick(ctx)`. + +- [ ] **Step 1: Write the failing tests** + +`core/services/failover/schedule_test.go`: + +```go +package failover + +import ( + "context" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type probeCall struct { + target string + inference bool +} + +type fakeProber struct { + mu sync.Mutex + calls []probeCall + fail map[string]error // target -> error returned by every probe +} + +func (p *fakeProber) record(target string, inference bool) error { + p.mu.Lock() + defer p.mu.Unlock() + p.calls = append(p.calls, probeCall{target, inference}) + return p.fail[target] +} +func (p *fakeProber) Liveness(_ context.Context, c config.ModelConfig, _ Kind, _ bool) error { + return p.record(c.Name, false) +} +func (p *fakeProber) Inference(_ context.Context, c config.ModelConfig, _ Kind, _ bool) error { + return p.record(c.Name, true) +} +func (p *fakeProber) take() []probeCall { + p.mu.Lock() + defer p.mu.Unlock() + out := p.calls + p.calls = nil + return out +} + +var _ = Describe("Manager probes", func() { + var ( + clock *fakeClock + src *fakeSource + prober *fakeProber + m *Manager + ctx = context.Background() + ) + + BeforeEach(func() { + clock = newFakeClock() + prober = &fakeProber{fail: map[string]error{}} + src = newFakeSource(remote("a"), local("b"), local("cold"), + chainCfg("chain", nil, t("a"), warmT("b"))) + m = New(src, WithClock(clock), WithProber(prober)) + }) + + It("probes idle targets on the first tick and not again before the interval", func() { + m.Tick(ctx) + Expect(prober.take()).To(ConsistOf(probeCall{"a", false}, probeCall{"b", false})) + clock.Advance(5 * time.Second) + m.Tick(ctx) + Expect(prober.take()).To(BeEmpty()) + }) + + It("skips the liveness probe for a target with recent traffic", func() { + m.Tick(ctx) + prober.take() + clock.Advance(14 * time.Second) + m.ReportSuccess("a") + clock.Advance(2 * time.Second) + m.Tick(ctx) + Expect(prober.take()).To(ConsistOf(probeCall{"b", false})) + }) + + It("trips a target whose liveness probe fails", func() { + prober.fail["a"] = errBoom + m.Tick(ctx) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + Expect(st.Active).To(Equal("b")) + }) + + It("recovers through liveness, then inference probes, then fails back after dwell", func() { + prober.fail["a"] = errBoom + m.Tick(ctx) + delete(prober.fail, "a") + prober.take() + + clock.Advance(15 * time.Second) + m.Tick(ctx) // liveness passes: recovering + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateRecovering)) + + for i := 0; i < 3; i++ { + clock.Advance(15 * time.Second) + m.Tick(ctx) + } + calls := prober.take() + Expect(calls).To(ContainElement(probeCall{"a", true})) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Active).To(Equal("a"), "60s min_dwell passed during the 4 ticks") + }) + + It("sends a recovering target back down when an inference probe fails", func() { + m.ReportFailure("a", errBoom) + clock.Advance(15 * time.Second) + m.Tick(ctx) // liveness passes: recovering + prober.fail["a"] = errBoom + clock.Advance(15 * time.Second) + m.Tick(ctx) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("never probes a down cold target and restores it after min_dwell", func() { + src.Put(chainCfg("chain", nil, t("cold"), warmT("b"))) + m.Sync() + m.ReportFailure("cold", errBoom) + prober.take() + clock.Advance(30 * time.Second) + m.Tick(ctx) + for _, c := range prober.take() { + Expect(c.target).ToNot(Equal("cold")) + } + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + clock.Advance(31 * time.Second) + m.Tick(ctx) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + + It("probes a target shared by two chains once per tick", func() { + src.Put(chainCfg("chain2", nil, t("a"), warmT("b"))) + m.Tick(ctx) + calls := prober.take() + n := 0 + for _, c := range calls { + if c.target == "a" { + n++ + } + } + Expect(n).To(Equal(1)) + }) + + It("closes subscriptions when Run stops", func() { + events, _ := m.Subscribe(1) + rctx, cancel := context.WithCancel(ctx) + done := make(chan struct{}) + go func() { m.Run(rctx); close(done) }() + cancel() + Eventually(done).Should(BeClosed()) + Eventually(events).Should(BeClosed()) + }) +}) +``` + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: compile failure (`m.Tick undefined`). + +- [ ] **Step 3: Implement** + +`core/services/failover/schedule.go`: + +```go +package failover + +import ( + "context" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +// Run drives probes and dwell-based fail-back until ctx ends. +func (m *Manager) Run(ctx context.Context) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + m.Tick(ctx) + for { + select { + case <-ctx.Done(): + m.close() + return + case <-ticker.C: + m.Tick(ctx) + } + } +} + +// Tick runs one pass: sync configs, run due probes, recompute chains. It is +// exported so tests can drive the manager without a real ticker. +func (m *Manager) Tick(ctx context.Context) { + m.Sync() + var wg sync.WaitGroup + for _, j := range m.dueProbes() { + wg.Add(1) + go func(j probeJob) { + defer wg.Done() + m.runProbe(ctx, j) + }(j) + } + wg.Wait() + m.Reevaluate() +} + +type probeJob struct { + target string + cfg config.ModelConfig + kind Kind + warm bool + inference bool + timeout time.Duration +} + +func (m *Manager) dueProbes() []probeJob { + m.mu.Lock() + defer m.mu.Unlock() + now := m.clock.Now() + var jobs []probeJob + for _, ts := range m.targets { + interval := ts.params.ProbeInterval() + inference := false + switch ts.state { + case StateMissing: + continue + case StateHealthy: + // A served request is as good as a liveness probe. + if now.Sub(ts.lastActivity) < interval || now.Sub(ts.lastProbe) < interval { + continue + } + case StateDown: + if ts.cold() { + // Loading a cold model only to probe it could evict others. + if now.Sub(ts.downSince) >= ts.params.MinDwell() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + m.recomputeForLocked(ts.name) + } + continue + } + if now.Sub(ts.lastProbe) < interval { + continue + } + case StateRecovering: + if ts.cold() || now.Sub(ts.lastProbe) < interval { + continue + } + inference = true + } + if m.prober == nil { + continue + } + cfg, ok := m.lookupTarget(ts.name) + if !ok { + continue + } + ts.lastProbe = now + jobs = append(jobs, probeJob{ + target: ts.name, cfg: cfg, kind: ts.kind, warm: ts.warm, + inference: inference, timeout: ts.params.ProbeTimeout(), + }) + } + return jobs +} + +func (m *Manager) runProbe(ctx context.Context, j probeJob) { + pctx, cancel := context.WithTimeout(ctx, j.timeout) + defer cancel() + var err error + if j.inference { + err = m.prober.Inference(pctx, j.cfg, j.kind, j.warm) + } else { + err = m.prober.Liveness(pctx, j.cfg, j.kind, j.warm) + } + if ctx.Err() != nil { + return // shutting down: a cancelled probe says nothing about the target + } + m.applyProbe(j, err) +} + +func (m *Manager) applyProbe(j probeJob, err error) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targets[j.target] + if ts == nil || ts.state == StateMissing { + return + } + if err != nil { + if ts.state == StateDown { + ts.lastError = err.Error() + } else { + m.recordFailureLocked(ts, err.Error()) + } + m.recomputeForLocked(ts.name) + return + } + switch ts.state { + case StateHealthy: + ts.lastActivity = m.clock.Now() + ts.failures = nil + case StateDown: + m.setTargetLocked(ts, StateRecovering, ReasonRecovery, "") + case StateRecovering: + if j.inference { + m.recordPassLocked(ts) + } + } + m.recomputeForLocked(ts.name) +} +``` + +- [ ] **Step 4: Run tests** + +Run: `go test -race ./core/services/failover/... 2>&1 | tail -10` +Expected: PASS. + +- [ ] **Step 5: Commit** + +```bash +git add core/services/failover +git commit -m "feat(failover): schedule liveness and recovery probes + +Idle targets get a liveness probe each interval, recovering targets an +inference probe. Cold local targets are never loaded to be probed. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 6: Default prober (remote HTTP and local gRPC) + +**Files:** +- Create: `core/services/failover/prober.go`, `core/services/failover/prober_test.go` +- Modify: `core/config/model_config.go` (add `ResolveAPIKey` next to `ProxyConfig`, line ~238-310) +- Test: `core/config/model_config_failover_test.go` (add ResolveAPIKey specs) +- Modify: `docs/superpowers/specs/2026-09-26-model-failover-chains-design.md` (cold-local row: "the model file exists"; see Step 5) + +**Interfaces:** +- Consumes: `Prober` (Task 4), `KindRemote/KindLocal`. +- Produces: `type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error)`, `func NewProber(load LoadFunc, modelPath string) *DefaultProber`, `func UpstreamBase(raw string) (string, error)`, `func UpstreamModel(cfg config.ModelConfig) string`, `(config.ProxyConfig).ResolveAPIKey() (string, error)`. + +- [ ] **Step 1: Write the failing tests** + +Add to `core/config/model_config_failover_test.go`: + +```go +var _ = Describe("ProxyConfig.ResolveAPIKey", func() { + It("reads the env var", func() { + GinkgoT().Setenv("FAILOVER_TEST_KEY", "k1") + Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey()).To(Equal("k1")) + }) + It("fails on an unset env var", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey() + 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()).To(Equal("k2")) + }) + It("returns empty when nothing is set", func() { + Expect(ProxyConfig{}.ResolveAPIKey()).To(Equal("")) + }) +}) +``` + +(add `os` and `path/filepath` imports) + +`core/services/failover/prober_test.go`: + +```go +package failover + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" +) + +type fakeUpstream struct { + mu sync.Mutex + srv *httptest.Server + models []string + status int + paths []string + auth string +} + +func newFakeUpstream() *fakeUpstream { + u := &fakeUpstream{status: http.StatusOK} + u.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + u.mu.Lock() + u.paths = append(u.paths, r.Method+" "+r.URL.Path) + u.auth = r.Header.Get("Authorization") + status, models := u.status, u.models + u.mu.Unlock() + _, _ = io.Copy(io.Discard, r.Body) + if status != http.StatusOK { + w.WriteHeader(status) + return + } + if r.URL.Path == "/v1/models" { + var data []map[string]string + for _, m := range models { + data = append(data, map[string]string{"id": m}) + } + _ = json.NewEncoder(w).Encode(map[string]any{"data": data}) + return + } + _, _ = w.Write([]byte(`{}`)) + })) + return u +} + +type fakeBackend struct { + grpc.Backend + healthy bool + predictErr error + predicted bool +} + +func (b *fakeBackend) HealthCheck(context.Context) (bool, error) { return b.healthy, nil } +func (b *fakeBackend) Predict(context.Context, *pb.PredictOptions, ...ggrpc.CallOption) (*pb.Reply, error) { + b.predicted = true + return &pb.Reply{}, b.predictErr +} + +var _ = Describe("DefaultProber", func() { + var ( + up *fakeUpstream + p *DefaultProber + ctx = context.Background() + ) + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.srv.Close) + p = NewProber(nil, "") + }) + + proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { + c := config.ModelConfig{Name: name, Backend: "cloud-proxy", KnownUsecaseStrings: usecases} + c.KnownUsecases = config.GetUsecasesFromYAML(usecases) + c.Proxy.UpstreamURL = up.srv.URL + "/v1/chat/completions" + c.Proxy.UpstreamModel = upstreamModel + return c + } + + DescribeTable("UpstreamBase", + func(in, want string) { + got, err := UpstreamBase(in) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(want)) + }, + Entry("full endpoint", "https://h:8080/v1/chat/completions", "https://h:8080"), + Entry("path prefix", "https://h/api/v1/chat/completions", "https://h/api"), + Entry("bare host", "https://h", "https://h"), + Entry("bare host slash", "https://h/", "https://h"), + ) + + It("passes liveness when the upstream lists the model", func() { + up.models = []string{"big-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", "big-llm"), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("GET /v1/models")) + }) + + It("uses the target name when upstream_model is empty", func() { + up.models = []string{"argus-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(Succeed()) + }) + + It("fails liveness when the model is not listed or the upstream errors", func() { + up.models = []string{"other"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("does not list"))) + up.status = http.StatusServiceUnavailable + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("503"))) + }) + + It("sends the API key as a bearer token", func() { + GinkgoT().Setenv("FAILOVER_PROBE_KEY", "sekret") + up.models = []string{"argus-llm"} + c := proxied("argus-llm", "") + c.Proxy.APIKeyEnv = "FAILOVER_PROBE_KEY" + Expect(p.Liveness(ctx, c, KindRemote, false)).To(Succeed()) + Expect(up.auth).To(Equal("Bearer sekret")) + }) + + DescribeTable("remote inference hits the usecase endpoint", + func(usecase, path string) { + Expect(p.Inference(ctx, proxied("m", "", usecase), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("POST " + path)) + }, + Entry("chat", "chat", "/v1/chat/completions"), + Entry("embeddings", "embeddings", "/v1/embeddings"), + Entry("transcription", "transcript", "/v1/audio/transcriptions"), + Entry("tts", "tts", "/v1/audio/speech"), + ) + + It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { + b := &fakeBackend{healthy: true} + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }, "") + c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) + Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) + b.healthy = false + Expect(p.Liveness(ctx, c, KindLocal, true)).To(HaveOccurred()) + Expect(p.Inference(ctx, c, KindLocal, true)).To(Succeed()) + Expect(b.predicted).To(BeTrue()) + b.predictErr = errors.New("boom") + Expect(p.Inference(ctx, c, KindLocal, true)).To(HaveOccurred()) + }) + + It("checks the model file for cold local liveness without loading", func() { + dir := GinkgoT().TempDir() + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { + Fail("cold liveness must not load the model") + return nil, nil + }, dir) + c := config.ModelConfig{Name: "cold", Backend: "llama-cpp"} + c.Model = "weights.gguf" + Expect(p.Liveness(ctx, c, KindLocal, false)).To(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(dir, "weights.gguf"), []byte("x"), 0o600)).To(Succeed()) + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + c.Model = "org/some-hf-repo" // no extension: downloaded on demand + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + }) +}) +``` + +If `c.Model` is not directly assignable (it lives in an embedded struct), set it through that struct, e.g. `c.PredictionOptions.Model = "weights.gguf"`; use whatever `validateFailover` reads as `c.Model`. + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/config/... ./core/services/failover/... 2>&1 | tail -5` +Expected: compile failure. + +- [ ] **Step 3: Implement `ResolveAPIKey`** + +In `core/config/model_config.go` after the `ProxyConfig` constants: + +```go +// ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or +// "" when neither is set. The cloud-proxy backend applies the same rules. +func (p ProxyConfig) ResolveAPIKey() (string, error) { + switch { + case p.APIKeyEnv != "": + v, ok := os.LookupEnv(p.APIKeyEnv) + if !ok { + return "", fmt.Errorf("proxy api_key_env %q is not set", p.APIKeyEnv) + } + return v, nil + case p.APIKeyFile != "": + b, err := os.ReadFile(p.APIKeyFile) + if err != nil { + return "", fmt.Errorf("proxy api_key_file: %w", err) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} +``` + +- [ ] **Step 4: Implement the prober** + +`core/services/failover/prober.go`: + +```go +package failover + +import ( + "bytes" + "context" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "mime/multipart" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// LoadFunc returns the backend for a local target, loading it if needed. +type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) + +// DefaultProber probes remote targets over the upstream's OpenAI-compatible +// API and local targets through their gRPC backend. +type DefaultProber struct { + HTTP *http.Client + Load LoadFunc + ModelPath string +} + +func NewProber(load LoadFunc, modelPath string) *DefaultProber { + return &DefaultProber{HTTP: &http.Client{}, Load: load, ModelPath: modelPath} +} + +func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { + switch { + case kind == KindRemote: + return p.remoteLiveness(ctx, cfg) + case warm: + return p.localHealth(ctx, cfg) + } + return p.coldLiveness(cfg) +} + +func (p *DefaultProber) Inference(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { + if kind == KindRemote { + return p.remoteInference(ctx, cfg) + } + return p.localInference(ctx, cfg) +} + +// UpstreamBase strips the endpoint path from a cloud-proxy upstream_url: +// everything from "/v1" on, so a path prefix before it survives. +func UpstreamBase(raw string) (string, error) { + u, err := url.Parse(raw) + if err != nil || u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf("invalid upstream_url %q", raw) + } + path := u.Path + if i := strings.Index(path, "/v1"); i >= 0 { + path = path[:i] + } + return u.Scheme + "://" + u.Host + strings.TrimSuffix(path, "/"), nil +} + +// UpstreamModel is the model name the upstream knows the target by. +func UpstreamModel(cfg config.ModelConfig) string { + if cfg.Proxy.UpstreamModel != "" { + return cfg.Proxy.UpstreamModel + } + return cfg.Name +} + +func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error { + key, err := cfg.Proxy.ResolveAPIKey() + if err != nil || key == "" { + return err + } + if cfg.Proxy.Provider == config.ProxyProviderAnthropic { + req.Header.Set("x-api-key", key) + req.Header.Set("anthropic-version", "2023-06-01") + return nil + } + req.Header.Set("Authorization", "Bearer "+key) + return nil +} + +func (p *DefaultProber) do(req *http.Request, cfg config.ModelConfig) (*http.Response, error) { + if err := p.authorize(req, cfg); err != nil { + return nil, err + } + resp, err := p.HTTP.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode/100 != 2 { + resp.Body.Close() + return nil, fmt.Errorf("upstream %s: HTTP %d", req.URL.Path, resp.StatusCode) + } + return resp, nil +} + +func (p *DefaultProber) remoteLiveness(ctx context.Context, cfg config.ModelConfig) error { + base, err := UpstreamBase(cfg.Proxy.UpstreamURL) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, base+"/v1/models", nil) + if err != nil { + return err + } + resp, err := p.do(req, cfg) + if err != nil { + return err + } + defer resp.Body.Close() + var list struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&list); err != nil { + return fmt.Errorf("upstream /v1/models: %w", err) + } + want := UpstreamModel(cfg) + for _, d := range list.Data { + if d.ID == want { + return nil + } + } + return fmt.Errorf("upstream does not list model %q", want) +} + +func (p *DefaultProber) remoteInference(ctx context.Context, cfg config.ModelConfig) error { + base, err := UpstreamBase(cfg.Proxy.UpstreamURL) + if err != nil { + return err + } + model := UpstreamModel(cfg) + ping := []map[string]string{{"role": "user", "content": "ping"}} + switch { + case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION): + if cfg.Proxy.Provider == config.ProxyProviderAnthropic { + return p.postJSON(ctx, cfg, base+"/v1/messages", map[string]any{"model": model, "max_tokens": 1, "messages": ping}) + } + return p.postJSON(ctx, cfg, base+"/v1/chat/completions", map[string]any{"model": model, "max_tokens": 1, "messages": ping}) + case cfg.HasUsecases(config.FLAG_EMBEDDINGS): + return p.postJSON(ctx, cfg, base+"/v1/embeddings", map[string]any{"model": model, "input": "ping"}) + case cfg.HasUsecases(config.FLAG_TRANSCRIPT): + return p.postTranscription(ctx, cfg, base, model) + case cfg.HasUsecases(config.FLAG_TTS): + return p.postJSON(ctx, cfg, base+"/v1/audio/speech", map[string]any{"model": model, "input": "ok"}) + } + // Image, video and other costly usecases: liveness is the confirmation. + return p.remoteLiveness(ctx, cfg) +} + +func (p *DefaultProber) postJSON(ctx context.Context, cfg config.ModelConfig, endpoint string, body any) error { + b, err := json.Marshal(body) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(b)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + resp, err := p.do(req, cfg) + if err != nil { + return err + } + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + return resp.Body.Close() +} + +func (p *DefaultProber) postTranscription(ctx context.Context, cfg config.ModelConfig, base, model string) error { + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + _ = mw.WriteField("model", model) + fw, err := mw.CreateFormFile("file", "probe.wav") + if err != nil { + return err + } + _, _ = fw.Write(silenceWAV()) + if err := mw.Close(); err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, base+"/v1/audio/transcriptions", &buf) + if err != nil { + return err + } + req.Header.Set("Content-Type", mw.FormDataContentType()) + resp, err := p.do(req, cfg) + if err != nil { + return err + } + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + return resp.Body.Close() +} + +// silenceWAV is 200 ms of 16 kHz mono 16-bit silence. +func silenceWAV() []byte { + const rate, samples = 16000, 3200 + data := samples * 2 + b := make([]byte, 44+data) + copy(b[0:], "RIFF") + binary.LittleEndian.PutUint32(b[4:], uint32(36+data)) + copy(b[8:], "WAVE") + copy(b[12:], "fmt ") + binary.LittleEndian.PutUint32(b[16:], 16) + binary.LittleEndian.PutUint16(b[20:], 1) // PCM + binary.LittleEndian.PutUint16(b[22:], 1) // mono + binary.LittleEndian.PutUint32(b[24:], rate) + binary.LittleEndian.PutUint32(b[28:], rate*2) + binary.LittleEndian.PutUint16(b[32:], 2) + binary.LittleEndian.PutUint16(b[34:], 16) + copy(b[36:], "data") + binary.LittleEndian.PutUint32(b[40:], uint32(data)) + return b +} + +func (p *DefaultProber) localHealth(ctx context.Context, cfg config.ModelConfig) error { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + // Load returns the running backend, or starts it again after a crash. + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + ok, err := b.HealthCheck(ctx) + if err != nil { + return err + } + if !ok { + return errors.New("backend health check failed") + } + return nil +} + +func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConfig) error { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + switch { + case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION): + _, err = b.Predict(ctx, &pb.PredictOptions{Prompt: "ping", Tokens: 1}) + return err + case cfg.HasUsecases(config.FLAG_EMBEDDINGS): + _, err = b.Embeddings(ctx, &pb.PredictOptions{Embeddings: "ping"}) + return err + } + // A backend process that answers HealthCheck rarely fails only for TTS or + // transcription, so a real request adds little here. + return p.localHealth(ctx, cfg) +} + +// coldLiveness checks the model file without loading the model. +func (p *DefaultProber) coldLiveness(cfg config.ModelConfig) error { + f := cfg.Model + if f == "" || p.ModelPath == "" || strings.Contains(f, "://") { + return nil + } + path := f + if !filepath.IsAbs(path) { + path = filepath.Join(p.ModelPath, f) + } + if _, err := os.Stat(path); err != nil { + if errors.Is(err, fs.ErrNotExist) && filepath.Ext(f) == "" { + return nil // a repository id, downloaded on demand + } + return fmt.Errorf("model file %s: %w", f, err) + } + return nil +} +``` + +Check the `pb.PredictOptions` field names (`Prompt`, `Tokens`, `Embeddings`) in `pkg/grpc/proto` and adjust only the field names if they differ. + +- [ ] **Step 5: Align the spec with cold liveness** + +In the spec's probe table, change the "local, cold" liveness cell to: "the model file exists (skipped for URLs and repository ids). The model is never loaded only to probe it." (The installed-backend check is left out: the loader installs backends on demand, so a missing backend is not a health signal.) `git add -f` the spec. + +- [ ] **Step 6: Run tests** + +Run: `go test -race ./core/config/... ./core/services/failover/... 2>&1 | tail -10` +Expected: PASS. + +- [ ] **Step 7: Commit** + +```bash +git add core/config core/services/failover +git add -f docs/superpowers/specs +git commit -m "feat(failover): probe remote targets over HTTP and local ones over gRPC + +Remote liveness uses /v1/models, which every OpenAI-compatible upstream +serves. Recovery sends one minimal request for the target's usecase. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 7: Application wiring, warm targets, metrics, traces + +**Files:** +- Create: `core/services/failover/metrics.go`, `core/services/failover/trace.go`, `core/application/failover.go` +- Modify: `core/services/failover/manager.go` (`emitLocked` records metrics), `core/application/application.go` (field + accessor), `core/application/startup.go` (construct + run), `core/application/watchdog.go` (`SyncPinnedModelsToWatchdog`), `core/trace/backend_trace.go` (new type) +- Test: `core/services/failover/metrics_test.go` + +**Interfaces:** +- Consumes: Tasks 4-6. +- Produces: `(*application.Application).FailoverManager() *failover.Manager`, `failover.RegisterMetrics(m *Manager)`, `failover.RecordAttemptTrace(enabled bool, chain, target string, err error)`, `trace.BackendTraceFailover`. + +- [ ] **Step 1: Write the failing test** + +`core/services/failover/metrics_test.go`: + +```go +package failover + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("metrics", func() { + It("registers and records without a meter provider", func() { + src := newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m := New(src, WithClock(newFakeClock())) + Expect(func() { RegisterMetrics(m) }).ToNot(Panic()) + Expect(func() { m.ReportFailure("a", errBoom) }).ToNot(Panic()) + }) + It("records an attempt trace only when enabled", func() { + Expect(func() { RecordAttemptTrace(false, "chain", "a", errBoom) }).ToNot(Panic()) + Expect(func() { RecordAttemptTrace(true, "chain", "a", errBoom) }).ToNot(Panic()) + }) +}) +``` + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: `undefined: RegisterMetrics`. + +- [ ] **Step 3: Implement metrics and traces** + +`core/services/failover/metrics.go`: + +```go +package failover + +import ( + "context" + "sync" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" +) + +var ( + metricsOnce sync.Once + switches metric.Int64Counter +) + +func initMetrics() { + metricsOnce.Do(func() { + meter := otel.Meter("github.com/mudler/LocalAI") + switches, _ = meter.Int64Counter("localai_failover_switches_total", + metric.WithDescription("Failover chain switches between targets")) + }) +} + +func recordSwitch(ev Event) { + initMetrics() + if switches == nil { + return + } + switches.Add(context.Background(), 1, metric.WithAttributes( + attribute.String("chain", ev.Chain), + attribute.String("from", ev.From), + attribute.String("to", ev.To), + attribute.String("reason", string(ev.Reason)), + )) +} + +// RegisterMetrics exports target health as a gauge. The application calls it +// once for its manager; tests create many managers and skip it. +func RegisterMetrics(m *Manager) { + meter := otel.Meter("github.com/mudler/LocalAI") + _, _ = meter.Int64ObservableGauge("localai_failover_target_up", + metric.WithDescription("1 when a failover target is healthy, 0 otherwise"), + metric.WithInt64Callback(func(_ context.Context, o metric.Int64Observer) error { + for name, state := range m.targetStates() { + v := int64(0) + if state == StateHealthy { + v = 1 + } + o.Observe(v, metric.WithAttributes(attribute.String("target", name))) + } + return nil + })) +} +``` + +Add to `manager.go`: + +```go +func (m *Manager) targetStates() map[string]TargetState { + m.mu.Lock() + defer m.mu.Unlock() + out := make(map[string]TargetState, len(m.targets)) + for name, ts := range m.targets { + out[name] = ts.state + } + return out +} +``` + +and at the top of `emitLocked`: + +```go + if ev.Type == EventChainSwitched { + recordSwitch(ev) + } +``` + +In `core/trace/backend_trace.go`, add to the `BackendTraceType` constants: + +```go + BackendTraceFailover BackendTraceType = "failover" +``` + +Check `core/http/react-ui/src` for a label map of backend trace types (`grep -rn "image_generation" core/http/react-ui/src`); if one exists, add `failover: 'Failover'` in the same style. + +`core/services/failover/trace.go`: + +```go +package failover + +import ( + "fmt" + "time" + + "github.com/mudler/LocalAI/core/trace" +) + +// RecordAttemptTrace shows in the Traces UI why a target was skipped. +func RecordAttemptTrace(enabled bool, chain, target string, err error) { + if !enabled || err == nil { + return + } + trace.RecordBackendTrace(trace.BackendTrace{ + Timestamp: time.Now(), + Type: trace.BackendTraceFailover, + ModelName: target, + Summary: fmt.Sprintf("failover chain %s: %s failed, trying the next target", chain, target), + Error: err.Error(), + Data: map[string]any{"chain": chain}, + }) +} +``` + +(If `BackendTrace` field names differ, match `core/trace/backend_trace.go`.) + +- [ ] **Step 4: Wire the application** + +`core/application/application.go`: add field `failoverManager *failover.Manager` to `Application` and: + +```go +// FailoverManager serves failover chains. Never nil after New. +func (a *Application) FailoverManager() *failover.Manager { return a.failoverManager } +``` + +`core/application/failover.go`: + +```go +package application + +import ( + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/xlog" +) + +// applyFailoverWarmTargets pins warm failover targets in the watchdog and +// loads them, so a switch does not wait for a cold load. +func (a *Application) applyFailoverWarmTargets(warm []string) { + a.SyncPinnedModelsToWatchdog() + for _, name := range warm { + if _, err := backend.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) + } + } +} +``` + +`core/application/watchdog.go` in `SyncPinnedModelsToWatchdog`, before `wd.SetPinnedModels(pinned)`: + +```go + if a.failoverManager != nil { + pinned = failover.MergePinned(pinned, a.failoverManager.WarmTargets()) + } +``` + +`core/application/startup.go`: next to `application.routerRegistry = router.NewRegistry()` (~251), construct the manager: + +```go + application.failoverManager = failover.New(application.ModelConfigLoader(), + failover.WithProber(failover.NewProber(func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) { + return application.ModelLoader().Load(backend.ModelOptions(cfg, options)...) + }, application.ModelLoader().ModelPath)), + failover.WithOnWarmChanged(application.applyFailoverWarmTargets), + ) +``` + +and directly after the `LoadToMemory` preload loop (~539-548), which runs after `initializeWatchdog`, start it: + +```go + failover.RegisterMetrics(application.failoverManager) + go application.failoverManager.Run(options.Context) +``` + +`grpc` here is `github.com/mudler/LocalAI/pkg/grpc`. If `startup.go` already imports a different package as `grpc`, alias this one `lagrpc`. + +- [ ] **Step 5: Build and test** + +Run: `go build ./... && go test ./core/services/failover/... ./core/application/... 2>&1 | tail -10` +Expected: build OK, tests PASS. + +- [ ] **Step 6: Commit** + +```bash +git add core/services/failover core/application core/trace core/http/react-ui/src +git commit -m "feat(failover): run the chain manager and keep warm targets loaded + +Warm local targets are pinned in the watchdog and preloaded. Switches +and target health are exported as metrics, skipped attempts as traces. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 8: HTTP chain resolution and in-request retry + +**Files:** +- Create: `core/http/middleware/failover.go`, `core/http/middleware/failover_test.go` +- Modify: `core/http/middleware/request.go` (`RequestExtractor` field; `SetModelAndConfig` wraps its body with `failoverRetry`; resolution after the alias block at ~181-195) +- Modify: `core/http/middleware/context_keys.go` (new key) +- Modify: `core/http/app.go:457` (call `SetFailoverManager`) + +**Interfaces:** +- Consumes: `failover.Manager.Plan`, `Attempt` methods, `failover.IsRetryable`, `failover.RecordAttemptTrace`. +- Produces: `(*RequestExtractor).SetFailoverManager(*failover.Manager)`, `ContextKeyFailoverAttempt = "failover.attempt"`, `HeaderServedModel = "X-LocalAI-Served-Model"`, `HeaderFailover = "X-LocalAI-Failover"`, `MaxFailoverReplayBody = 32 << 20`. + +Design note for the implementer: the retry loop wraps the whole `SetModelAndConfig` body plus `next`, so every attempt binds the request again from the replayed body. No route file changes. `failoverWriter` holds back a response with status >= 500 only while the request is resolved to a chain, so other requests are unaffected. + +- [ ] **Step 1: Write the failing tests** + +`core/http/middleware/failover_test.go` (use the same package and suite as `request_config_revision_test.go`): + +```go +package middleware + +import ( + "bytes" + "context" + "errors" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + + "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/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("failover chains in the request pipeline", func() { + var ( + app *echo.Echo + fm *failover.Manager + 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") + + 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) + 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) })) + }) + + 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("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")) + }) + + It("replays a multipart body for the next target", func() { + 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, "/v1/audio/transcriptions", &body) + req.Header.Set("Content-Type", mw.FormDataContentType()) + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(calls).To(Equal([]string{"a", "b"})) + }) + + 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()) + }) +}) +``` + +If `SetModelAndConfig` rejects `backend: fake-a` configs in this fixture (for example through the existence check), give the fixture configs the backend the revision test uses (`llama-cpp`); the handler never loads a backend. + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/http/middleware/... 2>&1 | tail -5` +Expected: compile failure (`SetFailoverManager undefined`). + +- [ ] **Step 3: Implement** + +In `core/http/middleware/context_keys.go` add: + +```go + // ContextKeyFailoverAttempt holds the *failoverState of a request whose + // model is a failover chain. + ContextKeyFailoverAttempt = "failover.attempt" +``` + +`core/http/middleware/failover.go`: + +```go +package middleware + +import ( + "bufio" + "bytes" + "errors" + "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()) + if err == nil { + 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 failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + req := c.Request() + src := req.Body + if src == nil { + src = http.NoBody + } + rec := &replayBody{src: src, limit: MaxFailoverReplayBody} + req.Body = rec + 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 := appConfig != nil && appConfig.EnableTracing + for { + w := &failoverWriter{ResponseWriter: orig, active: active} + resp.Writer = w + 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 { + att.Succeed() + return nil + } + 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 + } + req.Body = rec.replay() + req.MultipartForm, req.Form, req.PostForm = nil, nil, nil + resetResponse(resp, baseHeader) + } + } +} + +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 +} + +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.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) +} + +func (w *failoverWriter) Flush() { + w.release() + w.committed = true + if f, ok := w.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + +func (w *failoverWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + w.committed = true + h, ok := w.ResponseWriter.(http.Hijacker) + if !ok { + return nil, nil, errors.New("failover: response writer cannot hijack") + } + return h.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 +} + +func (r *replayBody) Read(p []byte) (int, error) { + n, err := r.src.Read(p) + if n > 0 && !r.overflow { + if r.buf.Len()+n > r.limit { + r.overflow = true + r.buf.Reset() + } 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 } + +// 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 +} +``` + +In `core/http/middleware/request.go`: + +1. Add field `failover *failover.Manager` to `RequestExtractor`. +2. In `SetModelAndConfig`, wrap the returned handler: + +```go +func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIRequest) echo.MiddlewareFunc { + return func(next echo.HandlerFunc) echo.HandlerFunc { + return failoverRetry(re.applicationConfig, func(c echo.Context) error { + // ... existing body, unchanged ... + }) + } +} +``` + +3. After the alias block (after `cfg = resolved` at ~195) and before the disabled check, add: + +```go + // 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 + } +``` + +In `core/http/app.go` after line 457: + +```go + requestExtractor.SetFailoverManager(application.FailoverManager()) +``` + +- [ ] **Step 4: Run tests** + +Run: `go test -race ./core/http/middleware/... 2>&1 | tail -15` +Expected: PASS, including the existing middleware specs. + +- [ ] **Step 5: Run the wider HTTP suite** + +Run: `make build-mock-backend && go test ./core/http/... 2>&1 | tail -15` +Expected: PASS (the retry wrapper is a no-op for plain models). + +- [ ] **Step 6: Commit** + +```bash +git add core/http +git commit -m "feat(failover): resolve chains per request and retry on the next target + +The retry wraps SetModelAndConfig, so each attempt binds the request +again from a replayed body. A 5xx of a chain request is held back until +the handler returns, and a streamed response is never retried. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 9: REST and SSE endpoints + +**Files:** +- Create: `core/http/endpoints/localai/failover.go`, `core/http/endpoints/localai/failover_test.go` +- Modify: `core/http/routes/localai.go` (register routes), `core/http/endpoints/localai/api_instructions.go` (+ test count 19 → 20 in `api_instructions_test.go:42`) +- Test: auth coverage next to the existing aliases auth tests (find them with `grep -rn '"/api/aliases"' core/http --include=*_test.go`) + +**Interfaces:** +- Consumes: `failover.Manager` `Status`, `ChainStatus`, `Pin`, `Unpin`, `Subscribe`, errors. +- Produces: routes `GET /api/failover`, `GET /api/failover/events`, `GET /api/failover/:chain`, `POST /api/failover/:chain/pin`, `DELETE /api/failover/:chain/pin`; type `localai.FailoverChainsResponse`. + +- [ ] **Step 1: Write the failing tests** + +`core/http/endpoints/localai/failover_test.go` (match the package name and suite used by the other `_test.go` files in that directory): + +```go +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 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")) + }) +}) +``` + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/http/endpoints/localai/... 2>&1 | tail -5` +Expected: compile failure. + +- [ ] **Step 3: Implement the handlers** + +`core/http/endpoints/localai/failover.go`: + +```go +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 + } + } + } + } +} +``` + +- [ ] **Step 4: Register routes and instructions** + +In `core/http/routes/localai.go`, next to `/api/aliases` (~92), where `app` / `application` is in scope (use the variable holding `*application.Application`): + +```go + fm := application.FailoverManager() + router.GET("/api/failover", localai.ListFailoverChainsEndpoint(fm)) + router.GET("/api/failover/events", localai.FailoverEventsEndpoint(fm)) + router.GET("/api/failover/:chain", localai.GetFailoverChainEndpoint(fm)) + router.POST("/api/failover/:chain/pin", localai.PinFailoverTargetEndpoint(fm), adminMiddleware) + router.DELETE("/api/failover/:chain/pin", localai.UnpinFailoverTargetEndpoint(fm), adminMiddleware) +``` + +In `api_instructions.go` `instructionDefs`, add: + +```go + { + 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.", + }, +``` + +and change `HaveLen(19)` to `HaveLen(20)` in `api_instructions_test.go`. + +- [ ] **Step 5: Auth tests** + +Open the test that covers `/api/aliases` auth (found with the grep above) and add the same cases for: +- `GET /api/failover` without credentials → 401 when auth is enabled; with a user key → 200. +- `POST /api/failover//pin` with a non-admin user → 403; with an admin → 200 (or 404 for an unknown chain, which still proves the admin gate passed). + +Mirror the existing assertions exactly; do not invent a new auth harness. + +- [ ] **Step 6: Swagger and tests** + +Run: `make swagger && go test ./core/http/... 2>&1 | tail -15` +Expected: swagger regenerates without errors; tests PASS. + +- [ ] **Step 7: Commit** + +```bash +git add core/http swagger +git commit -m "feat(failover): expose chain status, pins and events over REST and SSE + +Assisted-by: Claude:claude-opus-5-5" +``` + +(Use the actual swagger output directory if it is not `swagger/`; `git status` shows it.) + +--- + +### Task 10: End-to-end HTTP failover and endpoint audit + +**Files:** +- Modify: `tests/e2e/mock-backend/main.go` (`LoadModel` failure trigger, line ~89) +- Modify: `tests/e2e/cloud_proxy_helpers_test.go` (fake upstream serves `GET /v1/models`) +- Modify: `tests/e2e/e2e_suite_test.go` (model configs) +- Create: `tests/e2e/e2e_failover_test.go` + +**Interfaces:** +- Consumes: the running app from the e2e suite, Task 8 headers, Task 9 endpoints. + +- [ ] **Step 1: Add the mock load-failure trigger** + +In `tests/e2e/mock-backend/main.go` `LoadModel`, before the success return: + +```go + // Lets e2e specs build a failover target whose backend cannot load. + if strings.HasPrefix(in.Model, "fail-load") { + return &pb.Result{Message: "mock: load failure", Success: false}, nil + } +``` + +(add `strings` to the imports if missing) + +- [ ] **Step 2: Make the fake upstream answer `/v1/models`** + +In `newFakeOpenAIUpstream()` (`cloud_proxy_helpers_test.go:40-97`), add a models list and, at the top of the handler: + +```go + if r.Method == http.MethodGet && r.URL.Path == "/v1/models" { + u.mu.Lock() + ids := slices.Clone(u.models) + u.mu.Unlock() + data := make([]map[string]string, 0, len(ids)) + for _, id := range ids { + data = append(data, map[string]string{"id": id}) + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"data": data}) + return + } +``` + +with a `models []string` field and: + +```go +func (u *fakeOpenAIUpstream) SetModels(ids ...string) { + u.mu.Lock() + u.models = ids + u.mu.Unlock() +} +``` + +Use the helper's actual type, mutex and field names. The existing cloud-proxy specs must still pass. + +- [ ] **Step 3: Add configs to the suite** + +In `tests/e2e/e2e_suite_test.go`, where the existing mock-backend configs are written (lines ~105-125), write these additional configs with the same helper and style. For each family `f` in `chat, completion, embeddings, transcription, tts, image, rerank, vad`: + +```yaml +# fail-.yaml +name: fail- +backend: mock-backend +parameters: + model: fail-load- +``` + +```yaml +# chain-.yaml +name: chain- +failover: + targets: + - model: fail- + - model: +``` + +The remote chain needs the fake upstreams' URLs, which exist only at runtime. Write those configs in the `BeforeAll` of the remote spec (Step 4) and load them the way `e2e_cloud_proxy_test.go` registers its cloud-proxy models (read that file first and reuse its mechanism, for example writing YAML to the models dir and calling the loader, or `POST /models/import`). + +- [ ] **Step 4: Write the e2e specs** + +`tests/e2e/e2e_failover_test.go`. Use the suite's base URL variable and HTTP helpers (check `e2e_suite_test.go` for their names; below they are `apiURL` and `http.Post`): + +```go +package e2e_test + +import ( + "bytes" + "encoding/json" + "io" + "mime/multipart" + "net/http" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Failover chains", Label("failover"), func() { + postJSON := func(path string, body map[string]any) *http.Response { + b, _ := json.Marshal(body) + resp, err := http.Post(apiURL+path, "application/json", bytes.NewReader(b)) + Expect(err).ToNot(HaveOccurred()) + return resp + } + expectServedByMock := func(resp *http.Response) { + defer resp.Body.Close() + body, _ := io.ReadAll(resp.Body) + Expect(resp.StatusCode).To(BeNumerically("<", 300), string(body)) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).ToNot(HavePrefix("fail-")) + Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) + } + + DescribeTable("retries every endpoint family on the next target", + func(path string, body func(model string) map[string]any) { + expectServedByMock(postJSON(path, body("chain-"+CurrentSpecReport().LeafNodeText))) + }, + Entry("chat", "/v1/chat/completions", func(m string) map[string]any { + return map[string]any{"model": m, "messages": []map[string]string{{"role": "user", "content": "hi"}}} + }), + Entry("completion", "/v1/completions", func(m string) map[string]any { + return map[string]any{"model": m, "prompt": "hi"} + }), + Entry("embeddings", "/v1/embeddings", func(m string) map[string]any { + return map[string]any{"model": m, "input": "hi"} + }), + Entry("tts", "/v1/audio/speech", func(m string) map[string]any { + return map[string]any{"model": m, "input": "hi"} + }), + Entry("image", "/v1/images/generations", func(m string) map[string]any { + return map[string]any{"model": m, "prompt": "a cat", "size": "256x256"} + }), + Entry("rerank", "/v1/rerank", func(m string) map[string]any { + return map[string]any{"model": m, "query": "q", "documents": []string{"a", "b"}} + }), + Entry("vad", "/v1/vad", func(m string) map[string]any { + return map[string]any{"model": m, "audio": []float32{0, 0, 0, 0}} + }), + ) + + It("retries transcription with the multipart body", func() { + var body bytes.Buffer + mw := multipart.NewWriter(&body) + _ = mw.WriteField("model", "chain-transcription") + fw, _ := mw.CreateFormFile("file", "a.wav") + _, _ = fw.Write(testWAV()) + Expect(mw.Close()).To(Succeed()) + resp, err := http.Post(apiURL+"/v1/audio/transcriptions", mw.FormDataContentType(), &body) + Expect(err).ToNot(HaveOccurred()) + expectServedByMock(resp) + }) + + Describe("remote targets", Ordered, func() { + var up1, up2 *fakeOpenAIUpstream + + BeforeAll(func() { + up1, up2 = newFakeOpenAIUpstream(), newFakeOpenAIUpstream() + up1.SetModels("up-1") + up2.SetModels("up-2") + // Register up-1 and up-2 as cloud-proxy passthrough models pointing at + // up1.URL()+"/v1/chat/completions" and up2.URL()+"/v1/chat/completions", + // and chain-remote with probe interval 1s, recovery probes 2 and + // min_dwell 2s, using the same mechanism as e2e_cloud_proxy_test.go. + registerFailoverRemoteModels(up1.URL(), up2.URL()) + }) + + It("fails over when the primary upstream errors and fails back when it recovers", func() { + up1.SetScript(func([]byte) (int, string, string) { return 503, `{"error":"no healthy nodes"}`, "application/json" }) + up2.SetScript(func([]byte) (int, string, string) { + return 200, `{"id":"x","object":"chat.completion","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`, "application/json" + }) + resp := postJSON("/v1/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) + defer resp.Body.Close() + Expect(resp.StatusCode).To(Equal(200)) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-2")) + + up1.SetScript(func([]byte) (int, string, string) { + return 200, `{"id":"x","object":"chat.completion","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`, "application/json" + }) + Eventually(func() string { + r, err := http.Get(apiURL + "/api/failover/chain-remote") + if err != nil { + return "" + } + defer r.Body.Close() + var st struct { + Active string `json:"active"` + } + _ = json.NewDecoder(r.Body).Decode(&st) + return st.Active + }, 30*time.Second, 500*time.Millisecond).Should(Equal("up-1")) + }) + }) +}) +``` + +Add this helper to the same file (200 ms of 16 kHz mono 16-bit silence): + +```go +func testWAV() []byte { + const rate, samples = 16000, 3200 + data := samples * 2 + b := make([]byte, 44+data) + copy(b[0:], "RIFF") + binary.LittleEndian.PutUint32(b[4:], uint32(36+data)) + copy(b[8:], "WAVEfmt ") + binary.LittleEndian.PutUint32(b[16:], 16) + binary.LittleEndian.PutUint16(b[20:], 1) + binary.LittleEndian.PutUint16(b[22:], 1) + binary.LittleEndian.PutUint32(b[24:], rate) + binary.LittleEndian.PutUint32(b[28:], rate*2) + binary.LittleEndian.PutUint16(b[32:], 2) + binary.LittleEndian.PutUint16(b[34:], 16) + copy(b[36:], "data") + binary.LittleEndian.PutUint32(b[40:], uint32(data)) + return b +} +``` + +(add `encoding/binary` to the imports; if the package already defines `testWAV`, use the existing one.) + +Implement `registerFailoverRemoteModels(url1, url2 string)` in the same file using the cloud-proxy registration mechanism found in Step 3. `CurrentSpecReport().LeafNodeText` gives the entry name (`chat`, `completion`, ...), which matches the chain names from Step 3. For VAD and rerank, copy the request shape from existing e2e specs if they differ from the ones above. If the mock backend cannot serve a family at all (the request fails on the mock target too), remove that entry and list the family in the PR description as "covered by unit tests only". + +- [ ] **Step 5: Run the e2e specs** + +Run: `make build-mock-backend && go run github.com/onsi/ginkgo/v2/ginkgo --label-filter=failover -v ./tests/e2e 2>&1 | tail -30` +Expected: PASS. Then run the cloud-proxy specs to check the helper change: `go run github.com/onsi/ginkgo/v2/ginkgo --focus="cloud" -v ./tests/e2e 2>&1 | tail -10`. + +- [ ] **Step 6: Commit** + +```bash +git add tests/e2e +git commit -m "test(failover): cover retry per endpoint family and remote fail-back + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 11: Realtime pipelines + +**Files:** +- Modify: `core/http/endpoints/openai/realtime_model.go` (`RealtimeRoutingContext`, `buildRealtimeRoutingContext` ~860-880, `wrappedModel` ~39-90, `newModel` ~882, stage methods) +- Create: `core/http/endpoints/openai/realtime_failover.go`, `core/http/endpoints/openai/realtime_failover_test.go`, `core/http/endpoints/openai/types/failover.go` +- Modify: `core/http/endpoints/openai/types/server_events.go` (event type constant) +- Modify: `core/http/endpoints/openai/realtime.go` (start events after `newModel` succeeds, ~643-660) +- Test: `tests/e2e/realtime_ws_test.go` (new spec) + +**Interfaces:** +- Consumes: `failover.Manager.Do`, `ChainStatus`, `Subscribe`, `RecordAttemptTrace`. +- Produces: `types.ModelFailoverEvent`, `types.ServerEventTypeModelFailover = "localai.model.failover"`, `(*wrappedModel).stageCall`. + +- [ ] **Step 1: Add the event type** + +In `types/server_events.go`, next to `ServerEventTypeClassifierResult`: + +```go + ServerEventTypeModelFailover ServerEventType = "localai.model.failover" +``` + +`types/failover.go`, following the `ClassifierResultEvent` pattern in `types/classifier.go:426-469` (copy its `ServerEventBase` embedding and `MarshalJSON` shape exactly): + +```go +package types + +import "encoding/json" + +// ModelFailoverEvent tells a client which target serves a pipeline stage +// that names a failover chain. +type ModelFailoverEvent struct { + ServerEventBase + Chain string `json:"chain"` + Stage string `json:"stage"` + From string `json:"from"` + To string `json:"to"` + State string `json:"state"` + Reason string `json:"reason"` +} + +func (ModelFailoverEvent) ServerEventType() ServerEventType { return ServerEventTypeModelFailover } + +func (e ModelFailoverEvent) MarshalJSON() ([]byte, error) { + type alias ModelFailoverEvent + return json.Marshal(struct { + Type ServerEventType `json:"type"` + alias + }{Type: ServerEventTypeModelFailover, alias: alias(e)}) +} +``` + +- [ ] **Step 2: Write the failing unit tests** + +`core/http/endpoints/openai/realtime_failover_test.go` (package and suite of the other tests in that directory; `fakeTransport` is in `realtime_doubles_test.go:29`): + +```go +package openai + +import ( + "context" + "errors" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" + . "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"}}}}, + }) + }) + + It("routes a chain stage through the plan and retries before commit", func() { + m := &wrappedModel{failover: fm, stageChains: map[string]string{"tts": "chain"}, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }} + var tried []string + err := m.stageCall(context.Background(), "tts", 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("calls a plain stage once with its own config", func() { + m := &wrappedModel{} + base := &config.ModelConfig{Name: "plain"} + err := m.stageCall(context.Background(), "tts", base, func(cfg *config.ModelConfig, _ func()) error { + Expect(cfg).To(BeIdenticalTo(base)) + return nil + }) + Expect(err).ToNot(HaveOccurred()) + }) + + 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{"llm": "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("To", "b")))) + stop() + }) +}) +``` + +`t.events()` must return the events sent so far as `[]types.ServerEvent`. Read `realtime_doubles_test.go`: if `fakeTransport` has no such accessor, add a mutex-guarded `events()` method there. + +- [ ] **Step 3: Run to verify failure** + +Run: `go test ./core/http/endpoints/openai/... 2>&1 | tail -5` +Expected: compile failure (`unknown field failover`). + +- [ ] **Step 4: Implement stage resolution** + +In `realtime_model.go`: + +1. `RealtimeRoutingContext`: add `Failover *failover.Manager`. In `buildRealtimeRoutingContext`, set `Failover: a.FailoverManager()`. +2. `wrappedModel`: add + +```go + // failover and stageChains route pipeline stages that name a failover + // chain; stageChains maps a stage ("llm", "tts", ...) to its chain. + failover *failover.Manager + stageChains map[string]string + stageTargetConfig func(name string) (*config.ModelConfig, error) + appTracing bool +``` + +3. In `newModel`, after each stage config is loaded with `LoadResolvedModelConfig` and before `Validate()`: when the loaded config `IsFailover()`, record the chain and swap in its active target's config, so everything that inspects stage configs at session start (voice, reasoning, templates) sees a real model: + +```go + resolveStage := func(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { + if cfg == nil || !cfg.IsFailover() { + return cfg, nil + } + if routing == nil || routing.Failover == nil { + return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) + } + st, ok := routing.Failover.ChainStatus(cfg.Name) + if !ok { + return nil, fmt.Errorf("failover chain %q not found", cfg.Name) + } + stageChains[stage] = cfg.Name + return cl.LoadResolvedModelConfig(st.Active, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + } +``` + +Call it for `vad`, `transcription`, `llm`, `tts` and `sound_detection` right after each load. Declare `stageChains := map[string]string{}` before the loads. When building `&wrappedModel{...}`, set: + +```go + stageChains: stageChains, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { + return cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + }, + appTracing: appConfig.EnableTracing, +``` + +and, when `routing != nil`, `failover: routing.Failover`. + +4. `realtime_failover.go`: + +```go +package openai + +import ( + "context" + "sort" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" +) + +// stageCall runs fn with the config that should serve stage now. A plain +// stage uses base. A chain stage goes through the failover plan and retries +// on the next target until fn calls commit. +func (m *wrappedModel) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { + chain, ok := m.stageChains[stage] + if !ok || m.failover == nil { + return fn(base, func() {}) + } + return m.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { + cfg, err := m.stageTargetConfig(target) + if err != nil { + return err + } + err = fn(cfg, commit) + if err != nil { + failover.RecordAttemptTrace(m.appTracing, chain, target, err) + } + return err + }) +} + +// 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() { + 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 +} +``` + +5. Route each stage method through `stageCall`. Non-streaming example (`TTS`): + +```go +func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) { + var ( + out string + res *proto.Result + ) + err := m.stageCall(ctx, "tts", 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 +} +``` + +Apply the same shape to `VAD` (stage `vad`, `m.VADConfig`), `Transcribe` (`transcription`, `m.TranscriptionConfig`) and `SoundDetection` (`sound_detection`, `m.SoundDetectionConfig`). + +Streaming stages wrap the callback so the first delivered chunk commits: +- `TTSStream`: `onAudio` becomes `func(pcm []byte, sr int) error { commit(); return onAudio(pcm, sr) }`. +- `TranscribeStream`: `onDelta` becomes `func(s string) { commit(); onDelta(s) }`. +- `TranscribeLive`: call `commit()` right after the session opens successfully; the stage call only covers opening. + +`Predict` returns a closure that runs inference later. Keep the early part (message and tool preparation that does not depend on the config) where it is, and move everything that reads `turnCfg` (templating, `routeTurn` result use, `backend.ModelInference`) into the returned closure: + +```go + return func() (backend.LLMResponse, error) { + var resp backend.LLMResponse + err := m.stageCall(ctx, "llm", turnCfg, func(cfg *config.ModelConfig, commit func()) error { + cb := func(s string, u backend.TokenUsage) bool { + commit() + return tokenCallback(s, u) + } + // build predInput from cfg (not turnCfg) here, then: + predict, err := backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, *cfg, m.confLoader, m.appConfig, cb, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil) + if err != nil { + return err + } + resp, err = predict() + return err + }) + return resp, err + }, nil +``` + +Handle a nil `tokenCallback` (call `commit()` only). Apply `applyPipelineReasoning`/`applyPipelineThinking` to `cfg` inside the closure the same way `newModel` applies them to `cfgLLM`. A chain stage used together with a router (`routeTurn`) is out of scope: when `routeTurn` swapped the config, call the backend directly as today. + +6. In `realtime.go`, directly after the successful `newModel` for the main session (~643-660), start the events for the life of the session handler: + +```go + if wrapped, ok := m.(*wrappedModel); ok && wrapped.failover != nil && len(wrapped.stageChains) > 0 { + stopFailoverEvents := startFailoverEvents(t, wrapped.failover, wrapped.stageChains) + defer stopFailoverEvents() + } +``` + +Check that the enclosing function runs for the whole session (the defer must fire when the session ends, not when setup returns). If it returns earlier, store the stop func on the session and call it where the session is torn down. + +- [ ] **Step 5: Run unit tests** + +Run: `go test -race ./core/http/endpoints/openai/... 2>&1 | tail -15` +Expected: PASS, including the existing realtime specs. + +- [ ] **Step 6: Add the e2e realtime spec** + +In `tests/e2e/realtime_ws_test.go`, add a spec (Label `failover`) next to the existing WebSocket session spec, reusing its connection and turn helpers: +- Config `rt-failover` whose pipeline uses the same VAD/transcription/TTS models as the existing spec and `llm: chain-rt`, where `chain-rt` targets `[fail-rt, ]` and `fail-rt` has `parameters.model: fail-load-rt` (write these in the suite as in Task 10 Step 3). +- Assert: after connect, a `localai.model.failover` event with `stage: llm`, `reason: initial`, `to: fail-rt`. +- Send one user turn: a `localai.model.failover` event with `to` = the mock LLM and `reason: trip` arrives, and the turn completes with `response.done`. +- Send a second turn: it completes, and the conversation still holds the first turn's items (the item ids from the first turn's `conversation.item.created` events are still retrievable, or the second request's messages include them — use whatever the existing spec can observe). + +Run: `go run github.com/onsi/ginkgo/v2/ginkgo --label-filter=failover -v ./tests/e2e 2>&1 | tail -30` +Expected: PASS. + +- [ ] **Step 7: Commit** + +```bash +git add core/http/endpoints/openai tests/e2e +git commit -m "feat(failover): switch realtime pipeline stages per call + +A stage that names a chain is resolved on every call, so a switch keeps +the session and its conversation. Clients get localai.model.failover +events at session start and on every switch. + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 12: MCP admin tools + +**Files:** +- Modify: `pkg/mcp/localaitools/tools.go` (constants, `mutatingToolNames`), `client.go` (interface), `dto.go` (DTOs), `inproc/client.go`, `httpapi/client.go`, `httpapi/routes.go`, `server.go` (register), `prompts/20_tools.md`, `prompts/10_safety.md` +- Create: `pkg/mcp/localaitools/tools_failover.go` +- Modify tests: `coverage_test.go`, `server_test.go`, `fakes_test.go`, `inproc/client_test.go`, `httpapi/client_test.go`, `parity_test.go` +- Modify: where the in-process client is constructed (find with `grep -rn "inproc.Client{\|inproc.New" core pkg`) to pass the failover manager + +**Interfaces:** +- Consumes: `failover.Manager.Status/Pin/Unpin`. +- Produces: tools `list_failover_chains` (read-only), `pin_failover_target`, `unpin_failover_target` (mutating). + +- [ ] **Step 1: Write the failing tests** + +- `coverage_test.go` `toolToHTTPRoute`: + +```go + ToolListFailoverChains: "GET /api/failover", + ToolPinFailoverTarget: "POST /api/failover/:chain/pin", + ToolUnpinFailoverTarget: "DELETE /api/failover/:chain/pin", +``` + +- `server_test.go`: add `ToolListFailoverChains` to `expectedReadOnlyCatalog`; add dispatch rows in the table at ~154: + +```go + {ToolListFailoverChains, map[string]any{}, "ListFailoverChains"}, + {ToolPinFailoverTarget, map[string]any{"chain": "c", "target": "b"}, "PinFailoverTarget"}, + {ToolUnpinFailoverTarget, map[string]any{"chain": "c"}, "UnpinFailoverTarget"}, +``` + +- `fakes_test.go`: fake methods that record the call name like the existing alias fakes (lines ~157-170). +- `inproc/client_test.go` and `httpapi/client_test.go`: one spec each for list and pin, in the style of the alias specs. The httpapi spec asserts the request method and path; the inproc spec builds a `failover.Manager` over a small in-memory source with one chain. + +Run: `go test ./pkg/mcp/localaitools/... 2>&1 | tail -5` +Expected: compile failure. + +- [ ] **Step 2: Implement** + +`tools.go`: add `ToolListFailoverChains = "list_failover_chains"` to the read-only block, `ToolPinFailoverTarget = "pin_failover_target"` and `ToolUnpinFailoverTarget = "unpin_failover_target"` to the mutating block, and both mutating names to `mutatingToolNames`. + +`dto.go`: + +```go +type FailoverTargetInfo struct { + Model string `json:"model"` + Kind string `json:"kind"` + Warm bool `json:"warm"` + State string `json:"state"` + LastError string `json:"last_error,omitempty"` +} + +type FailoverChainInfo struct { + Name string `json:"name"` + State string `json:"state"` + Active string `json:"active"` + Pinned string `json:"pinned,omitempty"` + Targets []FailoverTargetInfo `json:"targets"` +} +``` + +`client.go` interface, next to the alias methods: + +```go + ListFailoverChains(ctx context.Context) ([]FailoverChainInfo, error) + PinFailoverTarget(ctx context.Context, chain, target string) error + UnpinFailoverTarget(ctx context.Context, chain string) error +``` + +`inproc/client.go`: add field `Failover *failover.Manager` to the client struct, set it where the client is constructed (from `application.FailoverManager()`), and: + +```go +func (c *Client) ListFailoverChains(_ context.Context) ([]localaitools.FailoverChainInfo, error) { + out := []localaitools.FailoverChainInfo{} + if c.Failover == nil { + return out, nil + } + for _, ch := range c.Failover.Status() { + info := localaitools.FailoverChainInfo{Name: ch.Name, State: string(ch.State), Active: ch.Active} + if ch.Pinned != nil { + info.Pinned = *ch.Pinned + } + for _, t := range ch.Targets { + info.Targets = append(info.Targets, localaitools.FailoverTargetInfo{ + Model: t.Model, Kind: string(t.Kind), Warm: t.Warm, State: string(t.State), LastError: t.LastError, + }) + } + out = append(out, info) + } + return out, nil +} + +func (c *Client) PinFailoverTarget(_ context.Context, chain, target string) error { + if c.Failover == nil { + return errors.New("failover is not running") + } + return c.Failover.Pin(chain, target) +} + +func (c *Client) UnpinFailoverTarget(_ context.Context, chain string) error { + if c.Failover == nil { + return errors.New("failover is not running") + } + return c.Failover.Unpin(chain) +} +``` + +`httpapi/routes.go`: `routeFailover = "/api/failover"`. `httpapi/client.go`: + +```go +func (c *Client) ListFailoverChains(ctx context.Context) ([]localaitools.FailoverChainInfo, error) { + var out struct { + Chains []localaitools.FailoverChainInfo `json:"chains"` + } + if err := c.do(ctx, http.MethodGet, routeFailover, nil, &out); err != nil { + return nil, err + } + return out.Chains, nil +} + +func (c *Client) PinFailoverTarget(ctx context.Context, chain, target string) error { + return c.do(ctx, http.MethodPost, routeFailover+"/"+url.PathEscape(chain)+"/pin", map[string]string{"target": target}, nil) +} + +func (c *Client) UnpinFailoverTarget(ctx context.Context, chain string) error { + return c.do(ctx, http.MethodDelete, routeFailover+"/"+url.PathEscape(chain)+"/pin", nil, nil) +} +``` + +The REST list returns `pinned` as a JSON string or null and `failover.TargetStatus` fields; `FailoverChainInfo` decodes the fields it shares. Check that `c.do` accepts a body map and a nil out; match its real signature. + +`tools_failover.go`, following `tools_aliases.go`: + +```go +package localaitools + +import ( + "context" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +func registerFailoverTools(s *mcp.Server, client LocalAIClient, opts Options) { + mcp.AddTool(s, &mcp.Tool{ + Name: ToolListFailoverChains, + Description: "List model failover chains, the target serving each one now, and the health of every target.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, any, error) { + chains, err := client.ListFailoverChains(ctx) + if err != nil { + return errorResult(err), nil, nil + } + return jsonResult(chains) + }) + + if opts.DisableMutating { + return + } + + mcp.AddTool(s, &mcp.Tool{ + Name: ToolPinFailoverTarget, + Description: "Force a failover chain to serve every request from one target, regardless of health, until it is unpinned.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, args struct { + Chain string `json:"chain" jsonschema:"failover chain name"` + Target string `json:"target" jsonschema:"target model to pin"` + }) (*mcp.CallToolResult, any, error) { + if err := client.PinFailoverTarget(ctx, args.Chain, args.Target); err != nil { + return errorResult(err), nil, nil + } + return jsonResult(map[string]string{"chain": args.Chain, "pinned": args.Target}) + }) + + mcp.AddTool(s, &mcp.Tool{ + Name: ToolUnpinFailoverTarget, + Description: "Remove the pin from a failover chain so health decides the target again.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, args struct { + Chain string `json:"chain" jsonschema:"failover chain name"` + }) (*mcp.CallToolResult, any, error) { + if err := client.UnpinFailoverTarget(ctx, args.Chain); err != nil { + return errorResult(err), nil, nil + } + return jsonResult(map[string]string{"chain": args.Chain, "pinned": ""}) + }) +} +``` + +Match the import path of the MCP SDK and the exact signatures of `errorResult`/`jsonResult` used in `tools_aliases.go`. + +`server.go`: add `registerFailoverTools(s, client, opts)` to the register list (~45-56). + +`prompts/20_tools.md`: under `## Read-only` add ``- `list_failover_chains` — List failover chains, their active target and target health.``; under `## Mutating` add ``- `pin_failover_target` — Force a failover chain to one target.`` and ``- `unpin_failover_target` — Remove a failover pin.``. `prompts/10_safety.md` line 5: add both mutating names to the backticked list. + +- [ ] **Step 3: Run tests** + +Run: `go test ./pkg/mcp/localaitools/... 2>&1 | tail -10` +Expected: PASS (including `TestToolHTTPRouteMappingComplete`, the prompts test and parity). + +- [ ] **Step 4: Commit** + +```bash +git add pkg/mcp core +git commit -m "feat(failover): add MCP tools to list chains and pin targets + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 13: Documentation and final verification + +**Files:** +- Create: `docs/content/features/model-failover.md` +- Modify: `docs/content/features/model-aliases.md`, `docs/content/features/openai-realtime.md`, and the cloud-proxy section (find it with `grep -rln "cloud-proxy" docs/content`; add the link on the page that documents the `proxy:` block) + +- [ ] **Step 1: Write the docs page** + +`docs/content/features/model-failover.md`: + +````markdown ++++ +disableToc = false +title = "Model Failover" +weight = 15 +url = "/features/model-failover/" ++++ + +A **failover chain** is a model name that is served by an ordered list of +other models. LocalAI sends each request to the first healthy target. When a +target fails, the request moves to the next target, and later requests stay +there until the first target has recovered. + +Use it to serve a model from a remote LocalAI or another OpenAI-compatible +provider, and to fall back to a local model when the remote one is down. + +## Declaring a chain + +```yaml +name: assistant-llm +failover: + targets: + - model: argus-llm # for example a cloud-proxy model + - model: gemma-local + warm: true # keep it loaded +``` + +Clients call `assistant-llm`. Each target is a normal model config. A chain +has no `backend` and no `parameters.model`. + +Optional settings, with their defaults: + +```yaml +failover: + probe: + interval: 15s # how often an idle target is checked + timeout: 5s + trip: + errors: 1 # failures within the window that mark a target down + window: 30s + recovery: + probes: 3 # test requests a target must pass before it is used again + min_dwell: 60s # minimum time on a lower target before moving back +``` + +Rules: + +- A chain needs at least 2 targets. A target can be an alias, but not another + chain. +- A chain cannot also set `alias` or `backend`. +- Responses name the chain as the model. The `X-LocalAI-Served-Model` header + names the target that served the request. + +## How the target is chosen + +- The active target is the first healthy target in the list. +- When a target fails, LocalAI marks it down and moves to the next target at + once. +- LocalAI moves back to a higher target only when that target has passed + `recovery.probes` test requests **and** the current target has been active + for at least `recovery.min_dwell`. This stops an unstable upstream from + moving traffic back and forth. +- When all targets are down, the chain is `degraded`. Each request still tries + every target in order. + +## Retry inside a request + +When a target fails before the response starts, LocalAI sends the same +request to the next target. The client does not see the failure. + +- LocalAI does not retry after the first byte of a response is sent (for + example after the first streamed token). The request fails, the target is + marked down, and the next request uses the next target. +- LocalAI does not retry client errors (4xx), such as a prompt that is too + long, because the next target would reject it too. +- Request bodies larger than 32 MiB are not retried. + +When the primary did not serve the request, the response has the header +`X-LocalAI-Failover: fallback`, or `X-LocalAI-Failover: degraded` when all +targets were down. + +## Health checks + +| Target | Regular check | Check before moving back | +|---|---|---| +| Remote (`cloud-proxy`) | `GET /v1/models` on the upstream lists the model | one small real request, for example a 1-token completion | +| Local, `warm: true` | the backend answers a health check | one small real request | +| Local, not warm | the model file exists | none: the target is used again after `min_dwell` | + +A request that succeeds counts as a check, so a busy target is almost never +probed. A model that is not warm is never loaded only to check it. + +When a target is in more than one chain, its check settings come from the +first of those chains in name order. + +## Warm targets + +`warm: true` loads a local target at startup and protects it from idle and +LRU eviction, so a switch does not wait for the model to load. Warm targets +count toward the active backend limit. If warm targets fill that limit, other +models cannot load, and the error names the warm targets. + +## Realtime pipelines + +A pipeline stage can name a chain: + +```yaml +name: assistant +pipeline: + vad: silero-vad + transcription: whisper-chain + llm: assistant-llm + tts: voice-chain +``` + +LocalAI resolves the chain for every call of the stage. When a chain switches, +the session stays open and keeps its conversation. The next turn uses the new +target. + +The session receives a `localai.model.failover` event for each chain stage when +it starts (`reason: initial`) and each time a chain switches: + +```json +{"type":"localai.model.failover","chain":"assistant-llm","stage":"llm", + "from":"argus-llm","to":"gemma-local","state":"fallback","reason":"trip"} +``` + +A chain used as a candidate of a router in a realtime pipeline is not +resolved per call. + +## Watching failover + +- `GET /api/failover` lists every chain, its active target and the state of + each target. +- `GET /api/failover/{chain}` returns one chain. +- `GET /api/failover/events` is a server-sent event stream. The first event is + `snapshot` with the full state. Then `chain.switched` and `target.state` + events follow. +- Metrics: `localai_failover_switches_total{chain,from,to,reason}` and + `localai_failover_target_up{target}`. +- With tracing on, each skipped target appears in the Traces view with the + error that made LocalAI skip it. + +## Pinning a target + +An admin can force a chain to one target, for example during maintenance: + +```bash +curl -X POST http://localhost:8080/api/failover/assistant-llm/pin \ + -H 'Content-Type: application/json' -d '{"target":"gemma-local"}' +curl -X DELETE http://localhost:8080/api/failover/assistant-llm/pin +``` + +While a chain is pinned, only the pinned target serves it. Health checks +continue. A restart removes the pin. + +## Assistant and MCP + +The LocalAI Assistant and `local-ai mcp-server` offer `list_failover_chains`, +`pin_failover_target` and `unpin_failover_target`. Create and edit chains with +the model config tools, like any other model. + +## Limits + +- Failover state is kept in memory by each LocalAI instance. Several frontends + in distributed mode each keep their own view. +- Chains do not nest. +- See also [model aliases]({{%relref "features/model-aliases" %}}) and the + [realtime API]({{%relref "features/openai-realtime" %}}). +```` + +- [ ] **Step 2: Cross-links** + +- `model-aliases.md`: add at the end of `## Rules and behavior`: `To serve a name from several models with automatic fallback, use a [failover chain]({{%relref "features/model-failover" %}}).` +- `openai-realtime.md`: in the pipeline section add: `A pipeline stage can name a [failover chain]({{%relref "features/model-failover" %}}); the stage then switches targets without closing the session.` +- The page that documents `proxy:` / `cloud-proxy`: add `To fall back to a local model when the upstream is down, list the proxy model in a [failover chain]({{%relref "features/model-failover" %}}).` + +- [ ] **Step 3: Final verification** + +Run, in order, and read each output: + +```bash +make protogen-go build-mock-backend +go vet ./core/... ./pkg/mcp/... +go test -race ./core/services/failover/... ./core/config/... ./core/http/... ./pkg/mcp/localaitools/... +go run github.com/onsi/ginkgo/v2/ginkgo --label-filter=failover -v ./tests/e2e +make swagger && git diff --stat -- swagger +make test-coverage-check +``` + +Expected: vet clean, all tests PASS, swagger has no uncommitted changes, coverage at or above the baseline. If coverage dropped, add tests; never edit `coverage-baseline.txt`. + +- [ ] **Step 4: Commit** + +```bash +git add docs/content +git commit -m "docs: document model failover chains + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +## PR notes (for the finishing step) + +The PR description must include: +- The MCP decision: tools added for list, pin and unpin; chain create/edit reuses the model config tools. +- Endpoint families without in-request retry, if Task 10 found any. +- The limits from the docs page (per-instance state, router candidates in realtime). +- A note that the human submitter adds `Signed-off-by` (DCO). From 63221ee779f799fc66ce214f6145fea7f5346594 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:25:26 +0000 Subject: [PATCH 04/79] feat(config): add failover chain block to model configs A chain is a model config with an ordered list of target models, probe, trip and recovery settings. Like an alias it has no backend. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/meta/registry.go | 33 +++++ core/config/meta/registry_test.go | 11 ++ core/config/meta/types.go | 1 + core/config/model_config.go | 11 ++ core/config/model_config_failover.go | 150 ++++++++++++++++++++++ core/config/model_config_failover_test.go | 68 ++++++++++ 6 files changed, 274 insertions(+) create mode 100644 core/config/model_config_failover.go create mode 100644 core/config/model_config_failover_test.go diff --git a/core/config/meta/registry.go b/core/config/meta/registry.go index 3e53a81fd..8cf4e4050 100644 --- a/core/config/meta/registry.go +++ b/core/config/meta/registry.go @@ -419,6 +419,39 @@ func DefaultRegistry() map[string]FieldMetaOverride { Order: 0, }, + // --- Failover --- + "failover.targets": { + Section: "failover", + Label: "Failover targets", + Description: "Ordered list of models that serve this chain. The first healthy target serves each request; later targets take over when it fails. Mark a local target warm to keep it loaded.", + Component: "json-editor", + Order: 0, + }, + "failover.probe.interval": { + Section: "failover", Label: "Probe interval", Component: "input", Order: 1, Advanced: true, + Description: "How often an idle target is checked, as a duration (default 15s).", Placeholder: "15s", + }, + "failover.probe.timeout": { + Section: "failover", Label: "Probe timeout", Component: "input", Order: 2, Advanced: true, + Description: "How long one probe may take (default 5s).", Placeholder: "5s", + }, + "failover.trip.errors": { + Section: "failover", Label: "Errors to trip", Component: "number", Order: 3, Advanced: true, + Description: "Failures within the trip window that mark a target down (default 1).", + }, + "failover.trip.window": { + Section: "failover", Label: "Trip window", Component: "input", Order: 4, Advanced: true, + Description: "Window in which failures are counted (default 30s).", Placeholder: "30s", + }, + "failover.recovery.probes": { + Section: "failover", Label: "Recovery probes", Component: "number", Order: 5, Advanced: true, + Description: "Consecutive real test requests a target must pass before it is used again (default 3).", + }, + "failover.recovery.min_dwell": { + Section: "failover", Label: "Minimum time on fallback", Component: "input", Order: 6, Advanced: true, + Description: "Minimum time on a lower target before traffic moves back to a recovered higher one (default 60s).", Placeholder: "60s", + }, + // --- Pipeline --- "pipeline.llm": { Section: "pipeline", diff --git a/core/config/meta/registry_test.go b/core/config/meta/registry_test.go index c3f1a494d..eff0cbf58 100644 --- a/core/config/meta/registry_test.go +++ b/core/config/meta/registry_test.go @@ -28,6 +28,17 @@ 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")) + var ids []string + for _, s := range meta.DefaultSections() { + ids = append(ids, s.ID) + } + Expect(ids).To(ContainElement("failover")) + }) }) var _ = Describe("MCP field metadata", func() { diff --git a/core/config/meta/types.go b/core/config/meta/types.go index a29e66967..b7774d157 100644 --- a/core/config/meta/types.go +++ b/core/config/meta/types.go @@ -70,6 +70,7 @@ func DefaultSections() []Section { return []Section{ {ID: "general", Label: "General", Icon: "settings", Order: 0}, {ID: "alias", Label: "Alias", Icon: "git-merge", Order: 5}, + {ID: "failover", Label: "Failover", Icon: "git-merge", Order: 6}, {ID: "llm", Label: "LLM", Icon: "cpu", Order: 10}, {ID: "parameters", Label: "Parameters", Icon: "sliders", Order: 20}, {ID: "templates", Label: "Templates", Icon: "file-text", Order: 30}, diff --git a/core/config/model_config.go b/core/config/model_config.go index 9754ef6ac..33fe12d21 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -75,6 +75,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"` @@ -1644,6 +1648,13 @@ func (c *ModelConfig) Validate() (bool, error) { return false, fmt.Errorf("a config with artifacts must declare exactly one %q target, found %d", modelartifacts.TargetModel, primaries) } + if c.IsFailover() { + if err := c.validateFailover(); err != nil { + return false, err + } + return true, nil + } + // An alias is a pure redirect: validate only its own shape here. Target // existence and the no-chain rule need the full config set, so the loader // (load-time) and the create/swap endpoints enforce those. diff --git a/core/config/model_config_failover.go b/core/config/model_config_failover.go new file mode 100644 index 000000000..e33a51ea8 --- /dev/null +++ b/core/config/model_config_failover.go @@ -0,0 +1,150 @@ +package config + +import ( + "fmt" + "time" +) + +// 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 } + +func (f FailoverConfig) ProbeInterval() time.Duration { + return durationOr(f.Probe.Interval, DefaultFailoverProbeInterval) +} +func (f FailoverConfig) ProbeTimeout() time.Duration { + return durationOr(f.Probe.Timeout, DefaultFailoverProbeTimeout) +} +func (f FailoverConfig) TripWindow() time.Duration { + return durationOr(f.Trip.Window, DefaultFailoverTripWindow) +} +func (f FailoverConfig) MinDwell() time.Duration { + return durationOr(f.Recovery.MinDwell, DefaultFailoverMinDwell) +} +func (f FailoverConfig) TripErrors() int { + if f.Trip.Errors <= 0 { + return DefaultFailoverTripErrors + } + return f.Trip.Errors +} +func (f FailoverConfig) RecoveryProbes() int { + if f.Recovery.Probes <= 0 { + return DefaultFailoverRecoveryProbes + } + return f.Recovery.Probes +} + +// WarmFailoverTargets returns the targets marked warm, in chain order. +func (c ModelConfig) WarmFailoverTargets() []string { + if c.Failover == nil { + return nil + } + var out []string + for _, t := range c.Failover.Targets { + if t.Warm { + out = append(out, t.Model) + } + } + return out +} + +func durationOr(s string, def time.Duration) time.Duration { + if s == "" { + return def + } + d, err := time.ParseDuration(s) + if err != nil || d <= 0 { + return def + } + return d +} + +// validateFailover checks what a chain can check without other configs. +// Target existence is checked by ModelConfigLoader.ValidateFailoverTargets. +func (c *ModelConfig) validateFailover() error { + if c.Name == "" { + return fmt.Errorf("failover config requires a name") + } + if c.IsAlias() { + return fmt.Errorf("model %q cannot set both alias and failover", c.Name) + } + if c.Backend != "" || c.Model != "" { + return fmt.Errorf("failover config %q must not set backend or parameters.model: a chain is a pure redirect", c.Name) + } + f := c.Failover + if len(f.Targets) < 2 { + return fmt.Errorf("failover chain %q needs at least 2 targets", c.Name) + } + seen := map[string]bool{} + for _, t := range f.Targets { + switch { + case t.Model == "": + return fmt.Errorf("failover chain %q has a target with no model", c.Name) + case t.Model == c.Name: + return fmt.Errorf("failover chain %q cannot list itself", c.Name) + case seen[t.Model]: + return fmt.Errorf("failover chain %q lists %q twice", c.Name, t.Model) + } + seen[t.Model] = true + } + for key, v := range map[string]string{ + "probe.interval": f.Probe.Interval, + "probe.timeout": f.Probe.Timeout, + "trip.window": f.Trip.Window, + "recovery.min_dwell": f.Recovery.MinDwell, + } { + if v == "" { + continue + } + if d, err := time.ParseDuration(v); err != nil || d <= 0 { + return fmt.Errorf("failover chain %q: invalid %s %q", c.Name, key, v) + } + } + if f.Trip.Errors < 0 { + return fmt.Errorf("failover chain %q: trip.errors must not be negative", c.Name) + } + if f.Recovery.Probes < 0 { + return fmt.Errorf("failover chain %q: recovery.probes must not be negative", c.Name) + } + return nil +} diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go new file mode 100644 index 000000000..e79ed5f45 --- /dev/null +++ b/core/config/model_config_failover_test.go @@ -0,0 +1,68 @@ +package config + +import ( + "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"), + ) +}) From f7c0bb5f4bb546a75a5ddf1f7015140ffdd5b890 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:32:34 +0000 Subject: [PATCH 05/79] feat(config): validate failover chain targets across configs Reject chains whose targets are missing or are chains, at load and on create or edit, and warn when the targets share no usecase. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/model_config_loader.go | 96 +++++++++++++++++++++ core/config/model_config_loader_test.go | 40 +++++++++ core/http/endpoints/localai/import_model.go | 6 ++ core/services/modeladmin/config.go | 6 ++ 4 files changed, 148 insertions(+) diff --git a/core/config/model_config_loader.go b/core/config/model_config_loader.go index 7733a7392..363fcc652 100644 --- a/core/config/model_config_loader.go +++ b/core/config/model_config_loader.go @@ -501,6 +501,80 @@ 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 +} + +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 +911,28 @@ 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) + } + } + return nil } diff --git a/core/config/model_config_loader_test.go b/core/config/model_config_loader_test.go index 1a3e9b03a..9912a365b 100644 --- a/core/config/model_config_loader_test.go +++ b/core/config/model_config_loader_test.go @@ -402,3 +402,43 @@ 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()) + }) +}) diff --git a/core/http/endpoints/localai/import_model.go b/core/http/endpoints/localai/import_model.go index b3cf491eb..73448430a 100644 --- a/core/http/endpoints/localai/import_model.go +++ b/core/http/endpoints/localai/import_model.go @@ -187,6 +187,12 @@ func ImportModelEndpoint(cl *config.ModelConfigLoader, gs *galleryop.GalleryServ return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) } + // Reject failover chains whose targets are missing or are themselves + // chains, for the same reason. + if err := cl.ValidateFailoverTargets(&modelConfig); err != nil { + return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) + } + // Create the configuration file configPath := filepath.Join(appConfig.SystemState.Model.ModelsPath, modelConfig.Name+".yaml") if err := utils.VerifyPath(modelConfig.Name+".yaml", appConfig.SystemState.Model.ModelsPath); err != nil { diff --git a/core/services/modeladmin/config.go b/core/services/modeladmin/config.go index 2cadfcaed..463ffe3f3 100644 --- a/core/services/modeladmin/config.go +++ b/core/services/modeladmin/config.go @@ -169,6 +169,9 @@ func (s *ConfigService) patchConfig(ctx context.Context, name string, patch map[ if err := s.Loader.ValidateAliasTarget(&updated); err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) } + if err := s.Loader.ValidateFailoverTargets(&updated); err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) + } var result *PatchResult err = s.withMutationRollback([]string{configPath}, func() error { if err := writeFileAtomic(configPath, yamlData, 0644); err != nil { @@ -286,6 +289,9 @@ func (s *ConfigService) editYAML(ctx context.Context, name string, body []byte) if err := s.Loader.ValidateAliasTarget(&req); err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) } + if err := s.Loader.ValidateFailoverTargets(&req); err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) + } configPath := existing.GetModelConfigFile() modelsPath := s.modelsPath() From bc610465b552222d6e010fd973f4b0c049721832 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:35:47 +0000 Subject: [PATCH 06/79] feat(failover): add types and retryable error classification Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/classify.go | 75 ++++++++++++ core/services/failover/classify_test.go | 39 +++++++ core/services/failover/failover_suite_test.go | 13 +++ core/services/failover/types.go | 107 ++++++++++++++++++ core/services/failover/types_test.go | 22 ++++ 5 files changed, 256 insertions(+) create mode 100644 core/services/failover/classify.go create mode 100644 core/services/failover/classify_test.go create mode 100644 core/services/failover/failover_suite_test.go create mode 100644 core/services/failover/types.go create mode 100644 core/services/failover/types_test.go diff --git a/core/services/failover/classify.go b/core/services/failover/classify.go new file mode 100644 index 000000000..58868c089 --- /dev/null +++ b/core/services/failover/classify.go @@ -0,0 +1,75 @@ +package failover + +import ( + "context" + "errors" + "net/http" + "regexp" + "strconv" + "strings" + + "github.com/labstack/echo/v4" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +// cloud-proxy translate mode reports upstream failures as plain text, so the +// status is only visible in the message. +var upstreamStatusRe = regexp.MustCompile(`upstream (\d{3})`) + +// Errors that the next target would reject in the same way. +var requestErrorMarkers = []string{ + "exceeds the available context size", + "is larger than the max context size", + "maximum context length", +} + +// IsRetryable reports whether a failed attempt should move to the next +// target. status is the HTTP status a handler wrote, or 0 when it returned err +// without writing. +func IsRetryable(err error, status int) bool { + if errors.Is(err, context.Canceled) { + return false + } + if status != 0 { + return retryableStatus(status) + } + if err == nil { + return false + } + var he *echo.HTTPError + if errors.As(err, &he) { + return retryableStatus(he.Code) + } + if errors.Is(err, context.DeadlineExceeded) { + return true + } + if st, ok := grpcstatus.FromError(err); ok { + switch st.Code() { + case codes.Unavailable, codes.Internal, codes.DeadlineExceeded, codes.Unknown: + return !isRequestError(st.Message()) + default: + return false + } + } + msg := err.Error() + if m := upstreamStatusRe.FindStringSubmatch(msg); m != nil { + code, _ := strconv.Atoi(m[1]) + return retryableStatus(code) + } + // Anything else is usually a dial or load failure of this target. + return !isRequestError(msg) +} + +func retryableStatus(code int) bool { + return code >= 500 && code != http.StatusNotImplemented +} + +func isRequestError(msg string) bool { + for _, m := range requestErrorMarkers { + if strings.Contains(msg, m) { + return true + } + } + return false +} diff --git a/core/services/failover/classify_test.go b/core/services/failover/classify_test.go new file mode 100644 index 000000000..455a9f993 --- /dev/null +++ b/core/services/failover/classify_test.go @@ -0,0 +1,39 @@ +package failover + +import ( + "context" + "errors" + "fmt" + "net/http" + + "github.com/labstack/echo/v4" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +var _ = DescribeTable("IsRetryable", + func(err error, status int, want bool) { + Expect(IsRetryable(err, status)).To(Equal(want)) + }, + Entry("nil error, no status", nil, 0, false), + Entry("held 503", nil, http.StatusServiceUnavailable, true), + Entry("held 500", nil, http.StatusInternalServerError, true), + Entry("held 501", nil, http.StatusNotImplemented, false), + Entry("client cancel", context.Canceled, 0, false), + Entry("wrapped client cancel", fmt.Errorf("predict: %w", context.Canceled), 0, false), + Entry("deadline", context.DeadlineExceeded, 0, true), + Entry("echo 502", echo.NewHTTPError(http.StatusBadGateway, "x"), 0, true), + Entry("echo 400", echo.NewHTTPError(http.StatusBadRequest, "x"), 0, false), + Entry("echo 404", echo.NewHTTPError(http.StatusNotFound, "x"), 0, false), + Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), 0, true), + Entry("grpc internal", grpcstatus.Error(codes.Internal, "x"), 0, true), + Entry("grpc deadline", grpcstatus.Error(codes.DeadlineExceeded, "x"), 0, true), + Entry("grpc unknown", grpcstatus.Error(codes.Unknown, "x"), 0, true), + Entry("grpc invalid argument", grpcstatus.Error(codes.InvalidArgument, "x"), 0, false), + Entry("cloud-proxy upstream 503", errors.New("cloud-proxy: upstream 503: no healthy nodes"), 0, true), + Entry("cloud-proxy upstream 429 stays 4xx", errors.New("cloud-proxy: upstream 429: slow down"), 0, false), + Entry("context overflow", errors.New("the request exceeds the available context size"), 0, false), + Entry("dial error", errors.New("dial tcp 10.0.0.1:8080: connect: connection refused"), 0, true), +) diff --git a/core/services/failover/failover_suite_test.go b/core/services/failover/failover_suite_test.go new file mode 100644 index 000000000..1a7bce9d5 --- /dev/null +++ b/core/services/failover/failover_suite_test.go @@ -0,0 +1,13 @@ +package failover + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestFailover(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Failover test suite") +} diff --git a/core/services/failover/types.go b/core/services/failover/types.go new file mode 100644 index 000000000..dc4ac2002 --- /dev/null +++ b/core/services/failover/types.go @@ -0,0 +1,107 @@ +// Package failover serves a model name from an ordered chain of target +// models, moving to the next target when one fails and back when it +// recovers. +package failover + +import ( + "slices" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +type TargetState string + +const ( + StateHealthy TargetState = "healthy" + StateDown TargetState = "down" + StateRecovering TargetState = "recovering" + StateMissing TargetState = "missing" +) + +type ChainState string + +const ( + ChainPrimary ChainState = "primary" + ChainFallback ChainState = "fallback" + ChainDegraded ChainState = "degraded" +) + +type Kind string + +const ( + KindLocal Kind = "local" + KindRemote Kind = "remote" +) + +type Reason string + +const ( + ReasonTrip Reason = "trip" + ReasonRecovery Reason = "recovery" + ReasonManual Reason = "manual" + ReasonDegraded Reason = "degraded" + ReasonMissing Reason = "missing" + ReasonInitial Reason = "initial" +) + +type EventType string + +const ( + EventChainSwitched EventType = "chain.switched" + EventTargetState EventType = "target.state" +) + +// Event is one change of a target state or of a chain's active target. +type Event struct { + Type EventType `json:"type"` + Chain string `json:"chain,omitempty"` + Target string `json:"target,omitempty"` + From string `json:"from"` + To string `json:"to"` + State string `json:"state,omitempty"` + Reason Reason `json:"reason"` + Error string `json:"error,omitempty"` + At time.Time `json:"at"` +} + +type TargetStatus struct { + Model string `json:"model"` + Kind Kind `json:"kind"` + Warm bool `json:"warm"` + State TargetState `json:"state"` + ConsecutiveOK int `json:"consecutive_ok"` + LastProbe *time.Time `json:"last_probe,omitempty"` + LastError string `json:"last_error,omitempty"` +} + +type ChainStatus struct { + Name string `json:"name"` + State ChainState `json:"state"` + Active string `json:"active"` + ActiveSince time.Time `json:"active_since"` + Pinned *string `json:"pinned"` + Targets []TargetStatus `json:"targets"` +} + +// KindOf decides how a target is probed: proxy backends forward to another +// server and are checked over HTTP, everything else runs in this instance. +func KindOf(cfg config.ModelConfig) Kind { + switch cfg.Backend { + case "cloud-proxy", "localai-proxy": + return KindRemote + } + return KindLocal +} + +// MergePinned adds warm failover targets to the config-pinned model list, so +// the watchdog never evicts them. +func MergePinned(pinned, warm []string) []string { + out := slices.Clone(pinned) + for _, w := range warm { + if !slices.Contains(out, w) { + out = append(out, w) + } + } + return out +} diff --git a/core/services/failover/types_test.go b/core/services/failover/types_test.go new file mode 100644 index 000000000..abbe0bd09 --- /dev/null +++ b/core/services/failover/types_test.go @@ -0,0 +1,22 @@ +package failover + +import ( + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("KindOf", func() { + It("treats proxy backends as remote", func() { + Expect(KindOf(config.ModelConfig{Backend: "cloud-proxy"})).To(Equal(KindRemote)) + Expect(KindOf(config.ModelConfig{Backend: "localai-proxy"})).To(Equal(KindRemote)) + Expect(KindOf(config.ModelConfig{Backend: "llama-cpp"})).To(Equal(KindLocal)) + }) +}) + +var _ = Describe("MergePinned", func() { + It("adds warm targets without duplicates", func() { + Expect(MergePinned([]string{"a", "b"}, []string{"b", "c"})).To(Equal([]string{"a", "b", "c"})) + Expect(MergePinned(nil, nil)).To(BeEmpty()) + }) +}) From 8cffc87c8b92ac4029b69ab106f9a9391d36a3ce Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:40:35 +0000 Subject: [PATCH 07/79] feat(failover): add chain manager with trip, fail-back and pins Health is tracked per target and the active target per chain. Fail-back waits for recovery probes and a minimum time on the fallback. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/fakes_test.go | 95 ++++ core/services/failover/manager.go | 619 +++++++++++++++++++++++++ core/services/failover/manager_test.go | 232 +++++++++ 3 files changed, 946 insertions(+) create mode 100644 core/services/failover/fakes_test.go create mode 100644 core/services/failover/manager.go create mode 100644 core/services/failover/manager_test.go diff --git a/core/services/failover/fakes_test.go b/core/services/failover/fakes_test.go new file mode 100644 index 000000000..cad6b3552 --- /dev/null +++ b/core/services/failover/fakes_test.go @@ -0,0 +1,95 @@ +package failover + +import ( + "sort" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +type fakeClock struct { + mu sync.Mutex + now time.Time +} + +func newFakeClock() *fakeClock { return &fakeClock{now: time.Date(2026, 9, 26, 10, 0, 0, 0, time.UTC)} } +func (c *fakeClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.now +} +func (c *fakeClock) Advance(d time.Duration) { + c.mu.Lock() + c.now = c.now.Add(d) + c.mu.Unlock() +} + +type fakeSource struct { + mu sync.Mutex + cfgs map[string]config.ModelConfig +} + +func newFakeSource(cfgs ...config.ModelConfig) *fakeSource { + s := &fakeSource{cfgs: map[string]config.ModelConfig{}} + for _, c := range cfgs { + s.cfgs[c.Name] = c + } + return s +} +func (s *fakeSource) Put(c config.ModelConfig) { s.mu.Lock(); s.cfgs[c.Name] = c; s.mu.Unlock() } +func (s *fakeSource) Delete(name string) { s.mu.Lock(); delete(s.cfgs, name); s.mu.Unlock() } +func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { + s.mu.Lock() + defer s.mu.Unlock() + c, ok := s.cfgs[n] + return c, ok +} +func (s *fakeSource) GetAllModelsConfigs() []config.ModelConfig { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]config.ModelConfig, 0, len(s.cfgs)) + for _, c := range s.cfgs { + out = append(out, c) + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out +} + +func local(name string) config.ModelConfig { + return config.ModelConfig{Name: name, Backend: "llama-cpp"} +} +func remote(name string) config.ModelConfig { + return config.ModelConfig{Name: name, Backend: "cloud-proxy"} +} + +// chainCfg builds a chain; fc may be nil for defaults. +func chainCfg(name string, fc *config.FailoverConfig, targets ...config.FailoverTarget) config.ModelConfig { + f := config.FailoverConfig{} + if fc != nil { + f = *fc + } + f.Targets = targets + return config.ModelConfig{Name: name, Failover: &f} +} + +func t(model string) config.FailoverTarget { return config.FailoverTarget{Model: model} } +func warmT(model string) config.FailoverTarget { + return config.FailoverTarget{Model: model, Warm: true} +} + +// drain returns the events buffered so far without blocking. +func drain(ch <-chan Event) []Event { + var out []Event + for { + select { + case ev, ok := <-ch: + if !ok { + return out + } + out = append(out, ev) + default: + return out + } + } +} diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go new file mode 100644 index 000000000..b38ec9a49 --- /dev/null +++ b/core/services/failover/manager.go @@ -0,0 +1,619 @@ +package failover + +import ( + "context" + "errors" + "fmt" + "slices" + "sort" + "sync" + "sync/atomic" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/xlog" +) + +var ( + ErrChainNotFound = errors.New("failover chain not found") + ErrTargetNotInChain = errors.New("target is not in this failover chain") + ErrNoTarget = errors.New("failover chain has no usable target") +) + +// ConfigSource is the part of ModelConfigLoader the manager reads. +type ConfigSource interface { + GetModelConfig(name string) (config.ModelConfig, bool) + GetAllModelsConfigs() []config.ModelConfig +} + +type Clock interface{ Now() time.Time } + +type realClock struct{} + +func (realClock) Now() time.Time { return time.Now() } + +type Option func(*Manager) + +func WithClock(c Clock) Option { return func(m *Manager) { m.clock = c } } +func WithProber(p Prober) Option { return func(m *Manager) { m.prober = p } } + +// WithOnWarmChanged is called outside the manager lock when the set of warm +// local targets changes. The application pins and preloads them. +func WithOnWarmChanged(fn func(warm []string)) Option { return func(m *Manager) { m.onWarm = fn } } + +// Manager tracks health per target and the active target per chain. +type Manager struct { + mu sync.Mutex + src ConfigSource + clock Clock + prober Prober + onWarm func([]string) + targets map[string]*targetState + chains map[string]*chainState + subs map[int]chan Event + nextSub int + warm []string + warmPending bool + closed bool +} + +type targetState struct { + name string + kind Kind + warm bool + state TargetState + failures []time.Time + consecutiveOK int + downSince time.Time + lastProbe time.Time + lastActivity time.Time + lastError string + // params come from the first chain, in name order, that lists the target. + params config.FailoverConfig +} + +func (ts *targetState) cold() bool { return ts.kind == KindLocal && !ts.warm } + +type chainState struct { + name string + cfg config.FailoverConfig + targets []string + active int + activeSince time.Time + pinned string + state ChainState +} + +func New(src ConfigSource, opts ...Option) *Manager { + m := &Manager{ + src: src, + clock: realClock{}, + targets: map[string]*targetState{}, + chains: map[string]*chainState{}, + subs: map[int]chan Event{}, + } + for _, o := range opts { + o(m) + } + return m +} + +// Sync reconciles chains with the config source. There is no config-change +// hook in the loader, so this runs on every tick and on a lookup miss. +func (m *Manager) Sync() { + m.mu.Lock() + m.syncLocked() + warm, deliver := m.takeWarmLocked() + m.mu.Unlock() + if deliver && m.onWarm != nil { + m.onWarm(warm) + } +} + +func (m *Manager) syncLocked() { + now := m.clock.Now() + seenChains := map[string]bool{} + claimed := map[string]bool{} + for _, c := range m.src.GetAllModelsConfigs() { + if !c.IsFailover() { + continue + } + seenChains[c.Name] = true + names := make([]string, 0, len(c.Failover.Targets)) + for _, t := range c.Failover.Targets { + names = append(names, t.Model) + } + ch := m.chains[c.Name] + if ch == nil || !slices.Equal(ch.targets, names) { + pinned := "" + if ch != nil && slices.Contains(names, ch.pinned) { + pinned = ch.pinned + } + ch = &chainState{name: c.Name, targets: names, activeSince: now, state: ChainPrimary, pinned: pinned} + m.chains[c.Name] = ch + } + ch.cfg = *c.Failover + for _, t := range c.Failover.Targets { + ts := m.targets[t.Model] + if ts == nil { + ts = &targetState{name: t.Model, state: StateHealthy} + m.targets[t.Model] = ts + } + if !claimed[t.Model] { + claimed[t.Model] = true + ts.params = *c.Failover + ts.warm = false + } + tc, ok := m.lookupTarget(t.Model) + if !ok { + m.setTargetLocked(ts, StateMissing, ReasonMissing, "target config not found") + continue + } + ts.kind = KindOf(tc) + if t.Warm && ts.kind == KindLocal { + ts.warm = true + } + if ts.state == StateMissing { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } + } + } + for name := range m.chains { + if !seenChains[name] { + delete(m.chains, name) + } + } + for name := range m.targets { + if !claimed[name] { + delete(m.targets, name) + } + } + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } + var warm []string + for name, ts := range m.targets { + if ts.warm { + warm = append(warm, name) + } + } + sort.Strings(warm) + if !slices.Equal(warm, m.warm) { + m.warm = warm + m.warmPending = true + } +} + +func (m *Manager) takeWarmLocked() ([]string, bool) { + if !m.warmPending { + return nil, false + } + m.warmPending = false + return slices.Clone(m.warm), true +} + +// lookupTarget returns the config that serves a target, one alias hop deep. +func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { + c, ok := m.src.GetModelConfig(name) + if ok && c.IsAlias() { + return m.src.GetModelConfig(c.Alias) + } + return c, ok +} + +func (m *Manager) chainLocked(name string) *chainState { + if ch := m.chains[name]; ch != nil { + return ch + } + m.syncLocked() + return m.chains[name] +} + +// targetLocked looks up a target's health state, syncing lazily on a miss so +// ReportFailure/ReportSuccess work before the first Plan or Sync call. +func (m *Manager) targetLocked(name string) *targetState { + if ts := m.targets[name]; ts != nil { + return ts + } + m.syncLocked() + return m.targets[name] +} + +// WarmTargets returns the warm local targets, sorted. +func (m *Manager) WarmTargets() []string { + m.mu.Lock() + defer m.mu.Unlock() + return slices.Clone(m.warm) +} + +// Reevaluate recomputes every chain. Dwell-based fail-back needs no event, so +// the scheduler calls this on every tick. +func (m *Manager) Reevaluate() { + m.mu.Lock() + defer m.mu.Unlock() + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } +} + +func (m *Manager) setTargetLocked(ts *targetState, to TargetState, reason Reason, errMsg string) { + if ts.state == to { + return + } + from := ts.state + now := m.clock.Now() + ts.state = to + switch to { + case StateDown: + ts.downSince = now + ts.consecutiveOK = 0 + ts.failures = nil + case StateRecovering, StateHealthy: + ts.consecutiveOK = 0 + ts.failures = nil + } + m.emitLocked(Event{Type: EventTargetState, Target: ts.name, From: string(from), To: string(to), Reason: reason, Error: errMsg, At: now}) +} + +// recomputeLocked picks the active target. override replaces the reason of a +// resulting switch (pin and unpin are always "manual"). +func (m *Manager) recomputeLocked(ch *chainState, override Reason) { + now := m.clock.Now() + prev := ch.active + next := prev + reason := ReasonTrip + best := -1 + for i, name := range ch.targets { + if ts := m.targets[name]; ts != nil && ts.state == StateHealthy { + best = i + break + } + } + switch { + case ch.pinned != "": + next = slices.Index(ch.targets, ch.pinned) + reason = ReasonManual + case best == -1: + // Nothing is healthy: keep the active target, Plan tries all of them. + case best > prev: + next = best // the active target is not healthy + case best < prev: + cur := m.targets[ch.targets[prev]] + curHealthy := cur != nil && cur.state == StateHealthy + if !curHealthy { + next = best + } else if now.Sub(ch.activeSince) >= ch.cfg.MinDwell() { + next = best + reason = ReasonRecovery + } + } + if override != "" { + reason = override + } + var state ChainState + switch { + case ch.pinned == "" && best == -1: + state = ChainDegraded + case next == 0: + state = ChainPrimary + default: + state = ChainFallback + } + switch { + case next != prev: + ch.active = next + ch.activeSince = now + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: reason, At: now}) + case state == ChainDegraded && ch.state != ChainDegraded: + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonDegraded, At: now}) + } + ch.state = state +} + +func (m *Manager) recomputeForLocked(target string) { + for _, ch := range m.chains { + if slices.Contains(ch.targets, target) { + m.recomputeLocked(ch, "") + } + } +} + +// Attempt walks the targets of one request in order. +type Attempt struct { + m *Manager + chain string + primary string + degraded bool + targets []string + i int +} + +// Plan returns the attempt order for one request: the active target, then the +// other healthy targets. A degraded chain tries every target in priority +// order; a pinned chain only the pinned target. +func (m *Manager) Plan(chain string) (*Attempt, error) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return nil, fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + att := &Attempt{m: m, chain: ch.name, primary: ch.targets[0], degraded: ch.state == ChainDegraded} + usable := func(name string) bool { + ts := m.targets[name] + if ts == nil || ts.state == StateMissing { + return false + } + return att.degraded || ts.state == StateHealthy + } + switch { + case ch.pinned != "": + att.targets = []string{ch.pinned} + case att.degraded: + for _, name := range ch.targets { + if usable(name) { + att.targets = append(att.targets, name) + } + } + default: + active := ch.targets[ch.active] + if usable(active) { + att.targets = append(att.targets, active) + } + for _, name := range ch.targets { + if name != active && usable(name) { + att.targets = append(att.targets, name) + } + } + } + if len(att.targets) == 0 { + return nil, fmt.Errorf("%w: %q", ErrNoTarget, chain) + } + return att, nil +} + +func (a *Attempt) Chain() string { return a.chain } +func (a *Attempt) Target() string { return a.targets[a.i] } +func (a *Attempt) Primary() string { return a.primary } +func (a *Attempt) Degraded() bool { return a.degraded } + +// Fail records err against the current target and moves to the next one. It +// returns false when no target is left. +func (a *Attempt) Fail(err error) bool { + a.m.ReportFailure(a.Target(), err) + if a.i+1 >= len(a.targets) { + return false + } + a.i++ + return true +} + +// Report records err against the current target without moving on: the +// response was already committed, so nothing is left to retry. +func (a *Attempt) Report(err error) { a.m.ReportFailure(a.Target(), err) } + +func (a *Attempt) Succeed() { a.m.ReportSuccess(a.Target()) } + +func (m *Manager) ReportFailure(target string, err error) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targetLocked(target) + if ts == nil { + return + } + msg := "" + if err != nil { + msg = err.Error() + } + m.recordFailureLocked(ts, msg) + m.recomputeForLocked(target) +} + +func (m *Manager) recordFailureLocked(ts *targetState, msg string) { + now := m.clock.Now() + ts.lastError = msg + switch ts.state { + case StateRecovering: + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + case StateHealthy: + cut := now.Add(-ts.params.TripWindow()) + kept := ts.failures[:0] + for _, f := range ts.failures { + if f.After(cut) { + kept = append(kept, f) + } + } + ts.failures = append(kept, now) + if len(ts.failures) >= ts.params.TripErrors() { + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + } + } +} + +func (m *Manager) ReportSuccess(target string) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targetLocked(target) + if ts == nil { + return + } + ts.lastActivity = m.clock.Now() + m.recordPassLocked(ts) + m.recomputeForLocked(target) +} + +// recordPassLocked counts a served request or a passed inference probe. +func (m *Manager) recordPassLocked(ts *targetState) { + switch ts.state { + case StateHealthy: + ts.failures = nil + return + case StateMissing: + return + case StateDown: + if ts.cold() { + // Cold targets are never probed; a served request is proof enough. + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + return + } + m.setTargetLocked(ts, StateRecovering, ReasonRecovery, "") + } + ts.consecutiveOK++ + if ts.consecutiveOK >= ts.params.RecoveryProbes() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } +} + +func (m *Manager) Pin(chain, target string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + if !slices.Contains(ch.targets, target) { + return fmt.Errorf("%w: %q", ErrTargetNotInChain, target) + } + ch.pinned = target + m.recomputeLocked(ch, ReasonManual) + return nil +} + +func (m *Manager) Unpin(chain string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + ch.pinned = "" + m.recomputeLocked(ch, ReasonManual) + return nil +} + +// Status returns every chain, sorted by name. +func (m *Manager) Status() []ChainStatus { + m.mu.Lock() + defer m.mu.Unlock() + m.syncLocked() + names := make([]string, 0, len(m.chains)) + for name := range m.chains { + names = append(names, name) + } + sort.Strings(names) + out := make([]ChainStatus, 0, len(names)) + for _, name := range names { + out = append(out, m.statusLocked(m.chains[name])) + } + return out +} + +func (m *Manager) ChainStatus(name string) (ChainStatus, bool) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(name) + if ch == nil { + return ChainStatus{}, false + } + return m.statusLocked(ch), true +} + +func (m *Manager) statusLocked(ch *chainState) ChainStatus { + cs := ChainStatus{Name: ch.name, State: ch.state, Active: ch.targets[ch.active], ActiveSince: ch.activeSince} + if ch.pinned != "" { + p := ch.pinned + cs.Pinned = &p + } + for _, name := range ch.targets { + st := TargetStatus{Model: name} + if ts := m.targets[name]; ts != nil { + st.Kind, st.Warm, st.State = ts.kind, ts.warm, ts.state + st.ConsecutiveOK, st.LastError = ts.consecutiveOK, ts.lastError + if !ts.lastProbe.IsZero() { + lp := ts.lastProbe + st.LastProbe = &lp + } + } + cs.Targets = append(cs.Targets, st) + } + return cs +} + +// Subscribe returns a buffered event channel and a cancel func. A subscriber +// that does not keep up loses events rather than blocking the manager. +func (m *Manager) Subscribe(buffer int) (<-chan Event, func()) { + m.mu.Lock() + defer m.mu.Unlock() + ch := make(chan Event, buffer) + if m.closed { + close(ch) + return ch, func() {} + } + id := m.nextSub + m.nextSub++ + m.subs[id] = ch + var once sync.Once + return ch, func() { + once.Do(func() { + m.mu.Lock() + defer m.mu.Unlock() + if c, ok := m.subs[id]; ok { + delete(m.subs, id) + close(c) + } + }) + } +} + +func (m *Manager) emitLocked(ev Event) { + for _, c := range m.subs { + select { + case c <- ev: + default: + xlog.Warn("failover: dropping event for a slow subscriber", "type", ev.Type, "chain", ev.Chain, "target", ev.Target) + } + } +} + +func (m *Manager) close() { + m.mu.Lock() + defer m.mu.Unlock() + m.closed = true + for id, c := range m.subs { + close(c) + delete(m.subs, id) + } +} + +// Do runs fn against the chain's targets in plan order. fn calls commit once +// output has reached the client; after that a failure is not retried. +func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Context, target string, commit func()) error) error { + att, err := m.Plan(chain) + if err != nil { + return err + } + for { + var committed atomic.Bool + err := fn(ctx, att.Target(), func() { committed.Store(true) }) + switch { + case err == nil: + att.Succeed() + return nil + case ctx.Err() != nil || !IsRetryable(err, 0): + return err + case committed.Load(): + att.Report(err) + return err + case !att.Fail(err): + return err + } + } +} + +// Prober checks targets. Implemented by DefaultProber (prober.go). +type Prober interface { + // Liveness is the cheap steady-state check. + Liveness(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error + // Inference sends one minimal real request to confirm recovery. + Inference(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error +} diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go new file mode 100644 index 000000000..0ea98c2a0 --- /dev/null +++ b/core/services/failover/manager_test.go @@ -0,0 +1,232 @@ +package failover + +import ( + "context" + "errors" + "time" + + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + . "github.com/onsi/gomega/gstruct" +) + +var errBoom = errors.New("dial tcp: connection refused") + +var _ = Describe("Manager", func() { + var ( + clock *fakeClock + src *fakeSource + m *Manager + ) + + BeforeEach(func() { + clock = newFakeClock() + src = newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m = New(src, WithClock(clock)) + }) + + switched := func(evs []Event) []Event { + var out []Event + for _, e := range evs { + if e.Type == EventChainSwitched { + out = append(out, e) + } + } + return out + } + + It("plans the primary first on a fresh chain", func() { + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Primary()).To(Equal("a")) + Expect(att.Degraded()).To(BeFalse()) + st, ok := m.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(st.State).To(Equal(ChainPrimary)) + Expect(st.Targets[0].Kind).To(Equal(KindRemote)) + Expect(st.Targets[1].Kind).To(Equal(KindLocal)) + }) + + It("returns ErrChainNotFound for an unknown chain", func() { + _, err := m.Plan("nope") + Expect(errors.Is(err, ErrChainNotFound)).To(BeTrue()) + }) + + It("trips on the first failure by default and switches with an event", func() { + events, cancel := m.Subscribe(16) + defer cancel() + att, _ := m.Plan("chain") + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(st.State).To(Equal(ChainFallback)) + Expect(st.Targets[0].State).To(Equal(StateDown)) + Expect(st.Targets[0].LastError).To(ContainSubstring("connection refused")) + sw := switched(drain(events)) + Expect(sw).To(HaveLen(1)) + Expect(sw[0]).To(MatchFields(IgnoreExtras, Fields{ + "Chain": Equal("chain"), "From": Equal("a"), "To": Equal("b"), + "State": Equal("fallback"), "Reason": Equal(ReasonTrip), + })) + }) + + It("counts failures inside the trip window only", func() { + src.Put(chainCfg("chain", &config.FailoverConfig{Trip: config.FailoverTrip{Errors: 2, Window: "30s"}}, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + clock.Advance(31 * time.Second) + m.ReportFailure("a", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + clock.Advance(time.Second) + m.ReportFailure("a", errBoom) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("fails back only after recovery probes and min_dwell", func() { + m.Plan("chain") + m.ReportFailure("a", errBoom) + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a real success counts like a passed inference probe + } + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Active).To(Equal("b"), "min_dwell has not passed") + clock.Advance(61 * time.Second) + events, cancel := m.Subscribe(16) + defer cancel() + m.Reevaluate() + st, _ = m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + Expect(switched(drain(events))[0].Reason).To(Equal(ReasonRecovery)) + }) + + It("moves up at once when the active target itself goes down", func() { + src.Put(chainCfg("chain", nil, t("a"), t("b"), t("c"))) + src.Put(local("c")) + m.Sync() + m.ReportFailure("a", errBoom) // active: b + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a healthy again, but dwell not passed + } + m.ReportFailure("b", errBoom) // b down: go to a now, not c + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + }) + + It("goes degraded when all targets are down and plans all of them in priority order", func() { + m.ReportFailure("a", errBoom) + m.ReportFailure("b", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.State).To(Equal(ChainDegraded)) + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Degraded()).To(BeTrue()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("pins a target regardless of health", func() { + Expect(m.Pin("chain", "b")).To(Succeed()) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse(), "a pin allows only the pinned target") + st, _ := m.ChainStatus("chain") + Expect(*st.Pinned).To(Equal("b")) + Expect(st.Active).To(Equal("b")) + Expect(m.Unpin("chain")).To(Succeed()) + st, _ = m.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + Expect(errors.Is(m.Pin("chain", "zzz"), ErrTargetNotInChain)).To(BeTrue()) + Expect(errors.Is(m.Pin("nope", "a"), ErrChainNotFound)).To(BeTrue()) + }) + + It("shares target health across chains", func() { + src.Put(chainCfg("chain2", nil, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + s1, _ := m.ChainStatus("chain") + s2, _ := m.ChainStatus("chain2") + Expect(s1.Active).To(Equal("b")) + Expect(s2.Active).To(Equal("b")) + }) + + It("marks a removed target missing and leaves it out of plans", func() { + m.Plan("chain") + src.Delete("a") + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateMissing)) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("resets a chain whose target list changed", func() { + m.ReportFailure("a", errBoom) + src.Put(local("c")) + src.Put(chainCfg("chain", nil, t("c"), t("b"))) + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("c")) + }) + + It("reports warm local targets and ignores warm on remote ones", func() { + var got []string + m = New(src, WithClock(clock), WithOnWarmChanged(func(w []string) { got = w })) + src.Put(chainCfg("chain", nil, warmT("a"), warmT("b"))) + m.Sync() + Expect(got).To(Equal([]string{"b"})) + Expect(m.WarmTargets()).To(Equal([]string{"b"})) + }) + + It("closes a subscription on cancel", func() { + events, cancel := m.Subscribe(1) + cancel() + _, ok := <-events + Expect(ok).To(BeFalse()) + cancel() // idempotent + }) + + Describe("Do", func() { + It("retries on the next target until commit", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + if target == "a" { + return errBoom + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + }) + + It("does not retry after commit but still trips the target", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + commit() + return errBoom + }) + Expect(err).To(MatchError(errBoom)) + Expect(tried).To(Equal([]string{"a"})) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("does not retry or trip on a non-retryable error", func() { + bad := errors.New("the request exceeds the available context size") + err := m.Do(context.Background(), "chain", func(_ context.Context, _ string, _ func()) error { return bad }) + Expect(err).To(MatchError(bad)) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + }) +}) From 3de62576097114b6a012ee85e6e64c4714e16733 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:41:24 +0000 Subject: [PATCH 08/79] fix(failover): satisfy lint on the manager package errcheck flagged two side-effect-only m.Plan calls in tests, and unused flagged close(), which Task 5's probe scheduler wires in. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/manager.go | 1 + core/services/failover/manager_test.go | 4 ++-- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index b38ec9a49..c52aa2b7d 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -575,6 +575,7 @@ func (m *Manager) emitLocked(ev Event) { } } +//nolint:unused // wired by the Task 5 probe scheduler's Stop, which owns the manager's lifecycle func (m *Manager) close() { m.mu.Lock() defer m.mu.Unlock() diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 0ea98c2a0..f10a04f2c 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -88,7 +88,7 @@ var _ = Describe("Manager", func() { }) It("fails back only after recovery probes and min_dwell", func() { - m.Plan("chain") + _, _ = m.Plan("chain") m.ReportFailure("a", errBoom) for i := 0; i < 3; i++ { m.ReportSuccess("a") // a real success counts like a passed inference probe @@ -158,7 +158,7 @@ var _ = Describe("Manager", func() { }) It("marks a removed target missing and leaves it out of plans", func() { - m.Plan("chain") + _, _ = m.Plan("chain") src.Delete("a") m.Sync() st, _ := m.ChainStatus("chain") From 560ffee182e8ce12b1748df91f590bdd8c1aea6f Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:45:41 +0000 Subject: [PATCH 09/79] fix(failover): emit chain.switched when leaving degraded in place recomputeLocked only fired the event on an active-target change or on entering degraded. When the active target itself recovered while every target was down, the chain silently left degraded with no event, so SSE/realtime consumers tracking chain.switched.state got stuck on "degraded". Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/manager.go | 6 ++++++ core/services/failover/manager_test.go | 19 +++++++++++++++++++ 2 files changed, 25 insertions(+) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index c52aa2b7d..1169061bf 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -305,7 +305,13 @@ func (m *Manager) recomputeLocked(ch *chainState, override Reason) { ch.activeSince = now m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: reason, At: now}) case state == ChainDegraded && ch.state != ChainDegraded: + // Entering degraded with no active-target change (every target is down). m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonDegraded, At: now}) + case state != ChainDegraded && ch.state == ChainDegraded: + // Leaving degraded with no active-target change (the active target + // itself recovered): SSE/realtime consumers watch chain.switched.state, + // so this must fire or they stay on "degraded" forever. + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonRecovery, At: now}) } ch.state = state } diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index f10a04f2c..3bef7945e 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -132,6 +132,25 @@ var _ = Describe("Manager", func() { Expect(att.Fail(errBoom)).To(BeFalse()) }) + It("emits chain.switched when leaving degraded without an active-target change", func() { + m.ReportFailure("a", errBoom) // active moves to b + m.ReportFailure("b", errBoom) // both down: degraded, active stays b + st, _ := m.ChainStatus("chain") + Expect(st.State).To(Equal(ChainDegraded)) + Expect(st.Active).To(Equal("b")) + events, cancel := m.Subscribe(16) + defer cancel() + m.ReportSuccess("b") // b is cold local: one success recovers it in place + st, _ = m.ChainStatus("chain") + Expect(st.State).To(Equal(ChainFallback)) + Expect(st.Active).To(Equal("b"), "the active target itself recovered, no switch needed") + sw := switched(drain(events)) + Expect(sw).To(HaveLen(1), "leaving degraded must still notify chain.switched listeners") + Expect(sw[0]).To(MatchFields(IgnoreExtras, Fields{ + "Chain": Equal("chain"), "State": Equal("fallback"), "Reason": Equal(ReasonRecovery), + })) + }) + It("pins a target regardless of health", func() { Expect(m.Pin("chain", "b")).To(Succeed()) att, _ := m.Plan("chain") From 4d7ced0ef602b22477639a2817696d334980c6de Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:48:05 +0000 Subject: [PATCH 10/79] feat(failover): schedule liveness and recovery probes Idle targets get a liveness probe each interval, recovering targets an inference probe. Cold local targets are never loaded to be probed. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/manager.go | 1 - core/services/failover/schedule.go | 148 ++++++++++++++++++++++ core/services/failover/schedule_test.go | 160 ++++++++++++++++++++++++ 3 files changed, 308 insertions(+), 1 deletion(-) create mode 100644 core/services/failover/schedule.go create mode 100644 core/services/failover/schedule_test.go diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 1169061bf..cfb2d8b66 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -581,7 +581,6 @@ func (m *Manager) emitLocked(ev Event) { } } -//nolint:unused // wired by the Task 5 probe scheduler's Stop, which owns the manager's lifecycle func (m *Manager) close() { m.mu.Lock() defer m.mu.Unlock() diff --git a/core/services/failover/schedule.go b/core/services/failover/schedule.go new file mode 100644 index 000000000..170de3cfa --- /dev/null +++ b/core/services/failover/schedule.go @@ -0,0 +1,148 @@ +package failover + +import ( + "context" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +// Run drives probes and dwell-based fail-back until ctx ends. Only this +// scheduler goroutine calls Sync: onWarm callbacks run after the manager's +// lock is released, and a concurrent Sync from elsewhere could reorder them. +func (m *Manager) Run(ctx context.Context) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + m.Tick(ctx) + for { + select { + case <-ctx.Done(): + m.close() + return + case <-ticker.C: + m.Tick(ctx) + } + } +} + +// Tick runs one pass: sync configs, run due probes, recompute chains. It is +// exported so tests can drive the manager without a real ticker. Like Run, it +// must only be called from the scheduler goroutine (see Run's comment on Sync). +func (m *Manager) Tick(ctx context.Context) { + m.Sync() + var wg sync.WaitGroup + for _, j := range m.dueProbes() { + wg.Add(1) + go func(j probeJob) { + defer wg.Done() + m.runProbe(ctx, j) + }(j) + } + wg.Wait() + m.Reevaluate() +} + +type probeJob struct { + target string + cfg config.ModelConfig + kind Kind + warm bool + inference bool + timeout time.Duration +} + +func (m *Manager) dueProbes() []probeJob { + m.mu.Lock() + defer m.mu.Unlock() + now := m.clock.Now() + var jobs []probeJob + for _, ts := range m.targets { + interval := ts.params.ProbeInterval() + inference := false + switch ts.state { + case StateMissing: + continue + case StateHealthy: + // A served request is as good as a liveness probe. + if now.Sub(ts.lastActivity) < interval || now.Sub(ts.lastProbe) < interval { + continue + } + case StateDown: + if ts.cold() { + // Loading a cold model only to probe it could evict others. + if now.Sub(ts.downSince) >= ts.params.MinDwell() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + m.recomputeForLocked(ts.name) + } + continue + } + if now.Sub(ts.lastProbe) < interval { + continue + } + case StateRecovering: + if ts.cold() || now.Sub(ts.lastProbe) < interval { + continue + } + inference = true + } + if m.prober == nil { + continue + } + cfg, ok := m.lookupTarget(ts.name) + if !ok { + continue + } + ts.lastProbe = now + jobs = append(jobs, probeJob{ + target: ts.name, cfg: cfg, kind: ts.kind, warm: ts.warm, + inference: inference, timeout: ts.params.ProbeTimeout(), + }) + } + return jobs +} + +func (m *Manager) runProbe(ctx context.Context, j probeJob) { + pctx, cancel := context.WithTimeout(ctx, j.timeout) + defer cancel() + var err error + if j.inference { + err = m.prober.Inference(pctx, j.cfg, j.kind, j.warm) + } else { + err = m.prober.Liveness(pctx, j.cfg, j.kind, j.warm) + } + if ctx.Err() != nil { + return // shutting down: a cancelled probe says nothing about the target + } + m.applyProbe(j, err) +} + +func (m *Manager) applyProbe(j probeJob, err error) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targets[j.target] + if ts == nil || ts.state == StateMissing { + return + } + if err != nil { + if ts.state == StateDown { + ts.lastError = err.Error() + } else { + m.recordFailureLocked(ts, err.Error()) + } + m.recomputeForLocked(ts.name) + return + } + switch ts.state { + case StateHealthy: + ts.lastActivity = m.clock.Now() + ts.failures = nil + case StateDown: + m.setTargetLocked(ts, StateRecovering, ReasonRecovery, "") + case StateRecovering: + if j.inference { + m.recordPassLocked(ts) + } + } + m.recomputeForLocked(ts.name) +} diff --git a/core/services/failover/schedule_test.go b/core/services/failover/schedule_test.go new file mode 100644 index 000000000..931f1a151 --- /dev/null +++ b/core/services/failover/schedule_test.go @@ -0,0 +1,160 @@ +package failover + +import ( + "context" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type probeCall struct { + target string + inference bool +} + +type fakeProber struct { + mu sync.Mutex + calls []probeCall + fail map[string]error // target -> error returned by every probe +} + +func (p *fakeProber) record(target string, inference bool) error { + p.mu.Lock() + defer p.mu.Unlock() + p.calls = append(p.calls, probeCall{target, inference}) + return p.fail[target] +} +func (p *fakeProber) Liveness(_ context.Context, c config.ModelConfig, _ Kind, _ bool) error { + return p.record(c.Name, false) +} +func (p *fakeProber) Inference(_ context.Context, c config.ModelConfig, _ Kind, _ bool) error { + return p.record(c.Name, true) +} +func (p *fakeProber) take() []probeCall { + p.mu.Lock() + defer p.mu.Unlock() + out := p.calls + p.calls = nil + return out +} + +var _ = Describe("Manager probes", func() { + var ( + clock *fakeClock + src *fakeSource + prober *fakeProber + m *Manager + ctx = context.Background() + ) + + BeforeEach(func() { + clock = newFakeClock() + prober = &fakeProber{fail: map[string]error{}} + src = newFakeSource(remote("a"), local("b"), local("cold"), + chainCfg("chain", nil, t("a"), warmT("b"))) + m = New(src, WithClock(clock), WithProber(prober)) + }) + + It("probes idle targets on the first tick and not again before the interval", func() { + m.Tick(ctx) + Expect(prober.take()).To(ConsistOf(probeCall{"a", false}, probeCall{"b", false})) + clock.Advance(5 * time.Second) + m.Tick(ctx) + Expect(prober.take()).To(BeEmpty()) + }) + + It("skips the liveness probe for a target with recent traffic", func() { + m.Tick(ctx) + prober.take() + clock.Advance(14 * time.Second) + m.ReportSuccess("a") + clock.Advance(2 * time.Second) + m.Tick(ctx) + Expect(prober.take()).To(ConsistOf(probeCall{"b", false})) + }) + + It("trips a target whose liveness probe fails", func() { + prober.fail["a"] = errBoom + m.Tick(ctx) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + Expect(st.Active).To(Equal("b")) + }) + + It("recovers through liveness, then inference probes, then fails back after dwell", func() { + prober.fail["a"] = errBoom + m.Tick(ctx) + delete(prober.fail, "a") + prober.take() + + clock.Advance(15 * time.Second) + m.Tick(ctx) // liveness passes: recovering + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateRecovering)) + + for i := 0; i < 3; i++ { + clock.Advance(15 * time.Second) + m.Tick(ctx) + } + calls := prober.take() + Expect(calls).To(ContainElement(probeCall{"a", true})) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Active).To(Equal("a"), "60s min_dwell passed during the 4 ticks") + }) + + It("sends a recovering target back down when an inference probe fails", func() { + m.ReportFailure("a", errBoom) + clock.Advance(15 * time.Second) + m.Tick(ctx) // liveness passes: recovering + prober.fail["a"] = errBoom + clock.Advance(15 * time.Second) + m.Tick(ctx) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("never probes a down cold target and restores it after min_dwell", func() { + src.Put(chainCfg("chain", nil, t("cold"), warmT("b"))) + m.Sync() + m.ReportFailure("cold", errBoom) + prober.take() + clock.Advance(30 * time.Second) + m.Tick(ctx) + for _, c := range prober.take() { + Expect(c.target).ToNot(Equal("cold")) + } + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + clock.Advance(31 * time.Second) + m.Tick(ctx) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + + It("probes a target shared by two chains once per tick", func() { + src.Put(chainCfg("chain2", nil, t("a"), warmT("b"))) + m.Tick(ctx) + calls := prober.take() + n := 0 + for _, c := range calls { + if c.target == "a" { + n++ + } + } + Expect(n).To(Equal(1)) + }) + + It("closes subscriptions when Run stops", func() { + events, _ := m.Subscribe(1) + rctx, cancel := context.WithCancel(ctx) + done := make(chan struct{}) + go func() { m.Run(rctx); close(done) }() + cancel() + Eventually(done).Should(BeClosed()) + Eventually(events).Should(BeClosed()) + }) +}) From a7dc8bf42c7eef0604bb07714779f64d01abd6e1 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:55:43 +0000 Subject: [PATCH 11/79] feat(failover): probe remote targets over HTTP and local ones over gRPC Remote liveness uses /v1/models, which every OpenAI-compatible upstream serves. Recovery sends one minimal request for the target's usecase. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/model_config.go | 20 ++ core/config/model_config_failover_test.go | 21 ++ core/services/failover/prober.go | 283 ++++++++++++++++++ core/services/failover/prober_test.go | 169 +++++++++++ ...2026-09-26-model-failover-chains-design.md | 2 +- 5 files changed, 494 insertions(+), 1 deletion(-) create mode 100644 core/services/failover/prober.go create mode 100644 core/services/failover/prober_test.go diff --git a/core/config/model_config.go b/core/config/model_config.go index 33fe12d21..533da24ec 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -302,6 +302,26 @@ const ( ProxyProviderAnthropic = "anthropic" ) +// ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or +// "" when neither is set. The cloud-proxy backend applies the same rules. +func (p ProxyConfig) ResolveAPIKey() (string, error) { + switch { + case p.APIKeyEnv != "": + v, ok := os.LookupEnv(p.APIKeyEnv) + if !ok { + return "", fmt.Errorf("proxy api_key_env %q is not set", p.APIKeyEnv) + } + return v, nil + case p.APIKeyFile != "": + b, err := os.ReadFile(p.APIKeyFile) + if err != nil { + return "", fmt.Errorf("proxy api_key_file: %w", 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 diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go index e79ed5f45..fe34c5c66 100644 --- a/core/config/model_config_failover_test.go +++ b/core/config/model_config_failover_test.go @@ -1,6 +1,8 @@ package config import ( + "os" + "path/filepath" "time" . "github.com/onsi/ginkgo/v2" @@ -66,3 +68,22 @@ failover: Entry("no name", func(c *ModelConfig) { c.Name = "" }, "requires a name"), ) }) + +var _ = Describe("ProxyConfig.ResolveAPIKey", func() { + It("reads the env var", func() { + GinkgoT().Setenv("FAILOVER_TEST_KEY", "k1") + Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey()).To(Equal("k1")) + }) + It("fails on an unset env var", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey() + 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()).To(Equal("k2")) + }) + It("returns empty when nothing is set", func() { + Expect(ProxyConfig{}.ResolveAPIKey()).To(Equal("")) + }) +}) diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go new file mode 100644 index 000000000..b3055c17f --- /dev/null +++ b/core/services/failover/prober.go @@ -0,0 +1,283 @@ +package failover + +import ( + "bytes" + "context" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "mime/multipart" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// LoadFunc returns the backend for a local target, loading it if needed. +type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) + +// DefaultProber probes remote targets over the upstream's OpenAI-compatible +// API and local targets through their gRPC backend. +type DefaultProber struct { + HTTP *http.Client + Load LoadFunc + ModelPath string +} + +func NewProber(load LoadFunc, modelPath string) *DefaultProber { + return &DefaultProber{HTTP: &http.Client{}, Load: load, ModelPath: modelPath} +} + +func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { + switch { + case kind == KindRemote: + return p.remoteLiveness(ctx, cfg) + case warm: + return p.localHealth(ctx, cfg) + } + return p.coldLiveness(cfg) +} + +func (p *DefaultProber) Inference(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { + if kind == KindRemote { + return p.remoteInference(ctx, cfg) + } + return p.localInference(ctx, cfg) +} + +// UpstreamBase strips the endpoint path from a cloud-proxy upstream_url: +// everything from "/v1" on, so a path prefix before it survives. +func UpstreamBase(raw string) (string, error) { + u, err := url.Parse(raw) + if err != nil || u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf("invalid upstream_url %q", raw) + } + path := u.Path + if i := strings.Index(path, "/v1"); i >= 0 { + path = path[:i] + } + return u.Scheme + "://" + u.Host + strings.TrimSuffix(path, "/"), nil +} + +// UpstreamModel is the model name the upstream knows the target by. +func UpstreamModel(cfg config.ModelConfig) string { + if cfg.Proxy.UpstreamModel != "" { + return cfg.Proxy.UpstreamModel + } + return cfg.Name +} + +func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error { + key, err := cfg.Proxy.ResolveAPIKey() + if err != nil || key == "" { + return err + } + if cfg.Proxy.Provider == config.ProxyProviderAnthropic { + req.Header.Set("x-api-key", key) + req.Header.Set("anthropic-version", "2023-06-01") + return nil + } + req.Header.Set("Authorization", "Bearer "+key) + return nil +} + +func (p *DefaultProber) do(req *http.Request, cfg config.ModelConfig) (*http.Response, error) { + if err := p.authorize(req, cfg); err != nil { + return nil, err + } + resp, err := p.HTTP.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode/100 != 2 { + resp.Body.Close() + return nil, fmt.Errorf("upstream %s: HTTP %d", req.URL.Path, resp.StatusCode) + } + return resp, nil +} + +func (p *DefaultProber) remoteLiveness(ctx context.Context, cfg config.ModelConfig) error { + base, err := UpstreamBase(cfg.Proxy.UpstreamURL) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, base+"/v1/models", nil) + if err != nil { + return err + } + resp, err := p.do(req, cfg) + if err != nil { + return err + } + defer resp.Body.Close() + var list struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&list); err != nil { + return fmt.Errorf("upstream /v1/models: %w", err) + } + want := UpstreamModel(cfg) + for _, d := range list.Data { + if d.ID == want { + return nil + } + } + return fmt.Errorf("upstream does not list model %q", want) +} + +func (p *DefaultProber) remoteInference(ctx context.Context, cfg config.ModelConfig) error { + base, err := UpstreamBase(cfg.Proxy.UpstreamURL) + if err != nil { + return err + } + model := UpstreamModel(cfg) + ping := []map[string]string{{"role": "user", "content": "ping"}} + switch { + case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION): + if cfg.Proxy.Provider == config.ProxyProviderAnthropic { + return p.postJSON(ctx, cfg, base+"/v1/messages", map[string]any{"model": model, "max_tokens": 1, "messages": ping}) + } + return p.postJSON(ctx, cfg, base+"/v1/chat/completions", map[string]any{"model": model, "max_tokens": 1, "messages": ping}) + case cfg.HasUsecases(config.FLAG_EMBEDDINGS): + return p.postJSON(ctx, cfg, base+"/v1/embeddings", map[string]any{"model": model, "input": "ping"}) + case cfg.HasUsecases(config.FLAG_TRANSCRIPT): + return p.postTranscription(ctx, cfg, base, model) + case cfg.HasUsecases(config.FLAG_TTS): + return p.postJSON(ctx, cfg, base+"/v1/audio/speech", map[string]any{"model": model, "input": "ok"}) + } + // Image, video and other costly usecases: liveness is the confirmation. + return p.remoteLiveness(ctx, cfg) +} + +func (p *DefaultProber) postJSON(ctx context.Context, cfg config.ModelConfig, endpoint string, body any) error { + b, err := json.Marshal(body) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(b)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + resp, err := p.do(req, cfg) + if err != nil { + return err + } + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + return resp.Body.Close() +} + +func (p *DefaultProber) postTranscription(ctx context.Context, cfg config.ModelConfig, base, model string) error { + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + _ = mw.WriteField("model", model) + fw, err := mw.CreateFormFile("file", "probe.wav") + if err != nil { + return err + } + _, _ = fw.Write(silenceWAV()) + if err := mw.Close(); err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, base+"/v1/audio/transcriptions", &buf) + if err != nil { + return err + } + req.Header.Set("Content-Type", mw.FormDataContentType()) + resp, err := p.do(req, cfg) + if err != nil { + return err + } + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20)) + return resp.Body.Close() +} + +// silenceWAV is 200 ms of 16 kHz mono 16-bit silence. +func silenceWAV() []byte { + const rate, samples = 16000, 3200 + data := samples * 2 + b := make([]byte, 44+data) + copy(b[0:], "RIFF") + binary.LittleEndian.PutUint32(b[4:], uint32(36+data)) + copy(b[8:], "WAVE") + copy(b[12:], "fmt ") + binary.LittleEndian.PutUint32(b[16:], 16) + binary.LittleEndian.PutUint16(b[20:], 1) // PCM + binary.LittleEndian.PutUint16(b[22:], 1) // mono + binary.LittleEndian.PutUint32(b[24:], rate) + binary.LittleEndian.PutUint32(b[28:], rate*2) + binary.LittleEndian.PutUint16(b[32:], 2) + binary.LittleEndian.PutUint16(b[34:], 16) + copy(b[36:], "data") + binary.LittleEndian.PutUint32(b[40:], uint32(data)) + return b +} + +func (p *DefaultProber) localHealth(ctx context.Context, cfg config.ModelConfig) error { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + // Load returns the running backend, or starts it again after a crash. + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + ok, err := b.HealthCheck(ctx) + if err != nil { + return err + } + if !ok { + return errors.New("backend health check failed") + } + return nil +} + +func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConfig) error { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + switch { + case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION): + _, err = b.Predict(ctx, &pb.PredictOptions{Prompt: "ping", Tokens: 1}) + return err + case cfg.HasUsecases(config.FLAG_EMBEDDINGS): + _, err = b.Embeddings(ctx, &pb.PredictOptions{Embeddings: "ping"}) + return err + } + // A backend process that answers HealthCheck rarely fails only for TTS or + // transcription, so a real request adds little here. + return p.localHealth(ctx, cfg) +} + +// coldLiveness checks the model file without loading the model. +func (p *DefaultProber) coldLiveness(cfg config.ModelConfig) error { + f := cfg.Model + if f == "" || p.ModelPath == "" || strings.Contains(f, "://") { + return nil + } + path := f + if !filepath.IsAbs(path) { + path = filepath.Join(p.ModelPath, f) + } + if _, err := os.Stat(path); err != nil { + if errors.Is(err, fs.ErrNotExist) && filepath.Ext(f) == "" { + return nil // a repository id, downloaded on demand + } + return fmt.Errorf("model file %s: %w", f, err) + } + return nil +} diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go new file mode 100644 index 000000000..2b8b7f42d --- /dev/null +++ b/core/services/failover/prober_test.go @@ -0,0 +1,169 @@ +package failover + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" +) + +type fakeUpstream struct { + mu sync.Mutex + srv *httptest.Server + models []string + status int + paths []string + auth string +} + +func newFakeUpstream() *fakeUpstream { + u := &fakeUpstream{status: http.StatusOK} + u.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + u.mu.Lock() + u.paths = append(u.paths, r.Method+" "+r.URL.Path) + u.auth = r.Header.Get("Authorization") + status, models := u.status, u.models + u.mu.Unlock() + _, _ = io.Copy(io.Discard, r.Body) + if status != http.StatusOK { + w.WriteHeader(status) + return + } + if r.URL.Path == "/v1/models" { + var data []map[string]string + for _, m := range models { + data = append(data, map[string]string{"id": m}) + } + _ = json.NewEncoder(w).Encode(map[string]any{"data": data}) + return + } + _, _ = w.Write([]byte(`{}`)) + })) + return u +} + +type fakeBackend struct { + grpc.Backend + healthy bool + predictErr error + predicted bool +} + +func (b *fakeBackend) HealthCheck(context.Context) (bool, error) { return b.healthy, nil } +func (b *fakeBackend) Predict(context.Context, *pb.PredictOptions, ...ggrpc.CallOption) (*pb.Reply, error) { + b.predicted = true + return &pb.Reply{}, b.predictErr +} + +var _ = Describe("DefaultProber", func() { + var ( + up *fakeUpstream + p *DefaultProber + ctx = context.Background() + ) + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.srv.Close) + p = NewProber(nil, "") + }) + + proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { + c := config.ModelConfig{Name: name, Backend: "cloud-proxy", KnownUsecaseStrings: usecases} + c.KnownUsecases = config.GetUsecasesFromYAML(usecases) + c.Proxy.UpstreamURL = up.srv.URL + "/v1/chat/completions" + c.Proxy.UpstreamModel = upstreamModel + return c + } + + DescribeTable("UpstreamBase", + func(in, want string) { + got, err := UpstreamBase(in) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(want)) + }, + Entry("full endpoint", "https://h:8080/v1/chat/completions", "https://h:8080"), + Entry("path prefix", "https://h/api/v1/chat/completions", "https://h/api"), + Entry("bare host", "https://h", "https://h"), + Entry("bare host slash", "https://h/", "https://h"), + ) + + It("passes liveness when the upstream lists the model", func() { + up.models = []string{"big-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", "big-llm"), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("GET /v1/models")) + }) + + It("uses the target name when upstream_model is empty", func() { + up.models = []string{"argus-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(Succeed()) + }) + + It("fails liveness when the model is not listed or the upstream errors", func() { + up.models = []string{"other"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("does not list"))) + up.status = http.StatusServiceUnavailable + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("503"))) + }) + + It("sends the API key as a bearer token", func() { + GinkgoT().Setenv("FAILOVER_PROBE_KEY", "sekret") + up.models = []string{"argus-llm"} + c := proxied("argus-llm", "") + c.Proxy.APIKeyEnv = "FAILOVER_PROBE_KEY" + Expect(p.Liveness(ctx, c, KindRemote, false)).To(Succeed()) + Expect(up.auth).To(Equal("Bearer sekret")) + }) + + DescribeTable("remote inference hits the usecase endpoint", + func(usecase, path string) { + Expect(p.Inference(ctx, proxied("m", "", usecase), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("POST " + path)) + }, + Entry("chat", "chat", "/v1/chat/completions"), + Entry("embeddings", "embeddings", "/v1/embeddings"), + Entry("transcription", "transcript", "/v1/audio/transcriptions"), + Entry("tts", "tts", "/v1/audio/speech"), + ) + + It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { + b := &fakeBackend{healthy: true} + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }, "") + c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) + Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) + b.healthy = false + Expect(p.Liveness(ctx, c, KindLocal, true)).To(HaveOccurred()) + Expect(p.Inference(ctx, c, KindLocal, true)).To(Succeed()) + Expect(b.predicted).To(BeTrue()) + b.predictErr = errors.New("boom") + Expect(p.Inference(ctx, c, KindLocal, true)).To(HaveOccurred()) + }) + + It("checks the model file for cold local liveness without loading", func() { + dir := GinkgoT().TempDir() + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { + Fail("cold liveness must not load the model") + return nil, nil + }, dir) + c := config.ModelConfig{Name: "cold", Backend: "llama-cpp"} + c.Model = "weights.gguf" + Expect(p.Liveness(ctx, c, KindLocal, false)).To(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(dir, "weights.gguf"), []byte("x"), 0o600)).To(Succeed()) + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + c.Model = "org/some-hf-repo" // no extension: downloaded on demand + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + }) +}) diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index c332082bf..11ef80e4f 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -152,7 +152,7 @@ Chain states: |---|---|---| | remote | `GET /v1/models` returns 2xx and lists the upstream model. `` is the scheme and host of `proxy.upstream_url` plus any path prefix before `/v1`. The upstream model is `proxy.upstream_model`, or the target name when it is empty. `/v1/models` works on any OpenAI-compatible upstream, and `/readyz` exists only on LocalAI. | one minimal real request, chosen by usecase | | local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. | -| local, cold | the config and model files exist and the backend is installed. The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | +| local, cold | the model file exists (skipped for URLs and repository ids). The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | Minimal requests by usecase: From b0a4c132ca64d56ed86bf1e4654bd3d629a8c389 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:59:25 +0000 Subject: [PATCH 12/79] fix(failover): treat a set-but-empty api_key_env as unset in ResolveAPIKey Matches cloud-proxy's resolveAPIKey exactly: os.Getenv + empty check rather than os.LookupEnv, so a variable that's set but empty errors instead of probing unauthenticated. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/model_config.go | 12 +++++++----- core/config/model_config_failover_test.go | 5 +++++ 2 files changed, 12 insertions(+), 5 deletions(-) diff --git a/core/config/model_config.go b/core/config/model_config.go index 533da24ec..9c38f178d 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -303,19 +303,21 @@ const ( ) // ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or -// "" when neither is set. The cloud-proxy backend applies the same rules. +// "" 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() (string, error) { switch { case p.APIKeyEnv != "": - v, ok := os.LookupEnv(p.APIKeyEnv) - if !ok { - return "", fmt.Errorf("proxy api_key_env %q is not set", p.APIKeyEnv) + v := os.Getenv(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: %w", err) + return "", fmt.Errorf("proxy api_key_file %q: %w", p.APIKeyFile, err) } return strings.TrimSpace(string(b)), nil } diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go index fe34c5c66..9e82868e9 100644 --- a/core/config/model_config_failover_test.go +++ b/core/config/model_config_failover_test.go @@ -78,6 +78,11 @@ var _ = Describe("ProxyConfig.ResolveAPIKey", func() { _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey() 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() + 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()) From 5ea35393fa9205a84e7ad58d3fa7c65dd451e908 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 16:06:15 +0000 Subject: [PATCH 13/79] feat(failover): run the chain manager and keep warm targets loaded Warm local targets are pinned in the watchdog and preloaded. Switches and target health are exported as metrics, skipped attempts as traces. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/application.go | 5 +++ core/application/failover.go | 17 ++++++++ core/application/startup.go | 19 +++++++++ core/application/watchdog.go | 4 ++ core/http/react-ui/src/pages/Traces.jsx | 1 + core/services/failover/manager.go | 14 +++++++ core/services/failover/metrics.go | 54 +++++++++++++++++++++++++ core/services/failover/metrics_test.go | 19 +++++++++ core/services/failover/trace.go | 23 +++++++++++ core/trace/backend_trace.go | 1 + 10 files changed, 157 insertions(+) create mode 100644 core/application/failover.go create mode 100644 core/services/failover/metrics.go create mode 100644 core/services/failover/metrics_test.go create mode 100644 core/services/failover/trace.go diff --git a/core/application/application.go b/core/application/application.go index b49851d47..ac66e75b6 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -16,6 +16,7 @@ import ( "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/galleryop" "github.com/mudler/LocalAI/core/services/monitoring" "github.com/mudler/LocalAI/core/services/nodes" @@ -82,6 +83,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 @@ -478,6 +480,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 diff --git a/core/application/failover.go b/core/application/failover.go new file mode 100644 index 000000000..3194e294b --- /dev/null +++ b/core/application/failover.go @@ -0,0 +1,17 @@ +package application + +import ( + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/xlog" +) + +// applyFailoverWarmTargets pins warm failover targets in the watchdog and +// loads them, so a switch does not wait for a cold load. +func (a *Application) applyFailoverWarmTargets(warm []string) { + a.SyncPinnedModelsToWatchdog() + for _, name := range warm { + if _, err := backend.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) + } + } +} diff --git a/core/application/startup.go b/core/application/startup.go index abc2f4a17..1382835f6 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -1,6 +1,7 @@ package application import ( + "context" "crypto/rand" "encoding/hex" "fmt" @@ -12,6 +13,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" @@ -27,6 +29,7 @@ import ( "github.com/mudler/LocalAI/core/trace" "github.com/mudler/LocalAI/internal" "github.com/mudler/LocalAI/pkg/downloader" + "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/modelartifacts" "github.com/mudler/LocalAI/pkg/signals" "github.com/mudler/LocalAI/pkg/vram" @@ -250,6 +253,16 @@ 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(func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) { + return application.ModelLoader().Load(backend.ModelOptions(cfg, options)...) + }, application.ModelLoader().ModelPath)), + failover.WithOnWarmChanged(application.applyFailoverWarmTargets), + ) + // 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. @@ -548,6 +561,12 @@ func New(opts ...config.AppOption) (*Application, error) { } } + // Start the failover scheduler: it syncs chains from config, runs + // liveness/recovery probes and dwell-based fail-back. Run is the only + // caller of Sync in production so onWarm callbacks stay ordered. + failover.RegisterMetrics(application.failoverManager) + go application.failoverManager.Run(options.Context) + // Watch the configuration directory startWatcher(options) diff --git a/core/application/watchdog.go b/core/application/watchdog.go index 330c95353..4674cb470 100644 --- a/core/application/watchdog.go +++ b/core/application/watchdog.go @@ -2,6 +2,7 @@ package application import ( "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/xlog" ) @@ -23,6 +24,9 @@ func (a *Application) SyncPinnedModelsToWatchdog() { pinned = append(pinned, cfg.Name) } } + if a.failoverManager != nil { + pinned = failover.MergePinned(pinned, a.failoverManager.WarmTargets()) + } wd.SetPinnedModels(pinned) xlog.Debug("Synced pinned models to watchdog", "count", len(pinned)) } diff --git a/core/http/react-ui/src/pages/Traces.jsx b/core/http/react-ui/src/pages/Traces.jsx index 6e8f3e800..3a44b9c21 100644 --- a/core/http/react-ui/src/pages/Traces.jsx +++ b/core/http/react-ui/src/pages/Traces.jsx @@ -105,6 +105,7 @@ const TYPE_COLORS = { vector_store: { bg: 'var(--color-accent-light)', color: 'var(--color-data-7)' }, token_classify: { bg: 'var(--color-info-light)', color: 'var(--color-data-3)' }, pattern_pii: { bg: 'var(--color-error-light)', color: 'var(--color-data-2)' }, + failover: { bg: 'var(--color-warning-light)', color: 'var(--color-data-2)' }, } function typeBadgeStyle(type) { diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index cfb2d8b66..408c44621 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -226,6 +226,17 @@ func (m *Manager) WarmTargets() []string { return slices.Clone(m.warm) } +// targetStates snapshots each target's health state for the metrics gauge. +func (m *Manager) targetStates() map[string]TargetState { + m.mu.Lock() + defer m.mu.Unlock() + out := make(map[string]TargetState, len(m.targets)) + for name, ts := range m.targets { + out[name] = ts.state + } + return out +} + // Reevaluate recomputes every chain. Dwell-based fail-back needs no event, so // the scheduler calls this on every tick. func (m *Manager) Reevaluate() { @@ -572,6 +583,9 @@ func (m *Manager) Subscribe(buffer int) (<-chan Event, func()) { } func (m *Manager) emitLocked(ev Event) { + if ev.Type == EventChainSwitched { + recordSwitch(ev) + } for _, c := range m.subs { select { case c <- ev: diff --git a/core/services/failover/metrics.go b/core/services/failover/metrics.go new file mode 100644 index 000000000..99d1a8212 --- /dev/null +++ b/core/services/failover/metrics.go @@ -0,0 +1,54 @@ +package failover + +import ( + "context" + "sync" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" +) + +var ( + metricsOnce sync.Once + switches metric.Int64Counter +) + +func initMetrics() { + metricsOnce.Do(func() { + meter := otel.Meter("github.com/mudler/LocalAI") + switches, _ = meter.Int64Counter("localai_failover_switches_total", + metric.WithDescription("Failover chain switches between targets")) + }) +} + +func recordSwitch(ev Event) { + initMetrics() + if switches == nil { + return + } + switches.Add(context.Background(), 1, metric.WithAttributes( + attribute.String("chain", ev.Chain), + attribute.String("from", ev.From), + attribute.String("to", ev.To), + attribute.String("reason", string(ev.Reason)), + )) +} + +// RegisterMetrics exports target health as a gauge. The application calls it +// once for its manager; tests create many managers and skip it. +func RegisterMetrics(m *Manager) { + meter := otel.Meter("github.com/mudler/LocalAI") + _, _ = meter.Int64ObservableGauge("localai_failover_target_up", + metric.WithDescription("1 when a failover target is healthy, 0 otherwise"), + metric.WithInt64Callback(func(_ context.Context, o metric.Int64Observer) error { + for name, state := range m.targetStates() { + v := int64(0) + if state == StateHealthy { + v = 1 + } + o.Observe(v, metric.WithAttributes(attribute.String("target", name))) + } + return nil + })) +} diff --git a/core/services/failover/metrics_test.go b/core/services/failover/metrics_test.go new file mode 100644 index 000000000..42dcd24b8 --- /dev/null +++ b/core/services/failover/metrics_test.go @@ -0,0 +1,19 @@ +package failover + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("metrics", func() { + It("registers and records without a meter provider", func() { + src := newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m := New(src, WithClock(newFakeClock())) + Expect(func() { RegisterMetrics(m) }).ToNot(Panic()) + Expect(func() { m.ReportFailure("a", errBoom) }).ToNot(Panic()) + }) + It("records an attempt trace only when enabled", func() { + Expect(func() { RecordAttemptTrace(false, "chain", "a", errBoom) }).ToNot(Panic()) + Expect(func() { RecordAttemptTrace(true, "chain", "a", errBoom) }).ToNot(Panic()) + }) +}) diff --git a/core/services/failover/trace.go b/core/services/failover/trace.go new file mode 100644 index 000000000..8eb6618f1 --- /dev/null +++ b/core/services/failover/trace.go @@ -0,0 +1,23 @@ +package failover + +import ( + "fmt" + "time" + + "github.com/mudler/LocalAI/core/trace" +) + +// RecordAttemptTrace shows in the Traces UI why a target was skipped. +func RecordAttemptTrace(enabled bool, chain, target string, err error) { + if !enabled || err == nil { + return + } + trace.RecordBackendTrace(trace.BackendTrace{ + Timestamp: time.Now(), + Type: trace.BackendTraceFailover, + ModelName: target, + Summary: fmt.Sprintf("failover chain %s: %s failed, trying the next target", chain, target), + Error: err.Error(), + Data: map[string]any{"chain": chain}, + }) +} diff --git a/core/trace/backend_trace.go b/core/trace/backend_trace.go index 072e41ebf..e14370c45 100644 --- a/core/trace/backend_trace.go +++ b/core/trace/backend_trace.go @@ -48,6 +48,7 @@ const ( BackendTraceTokenClassify BackendTraceType = "token_classify" BackendTracePatternPII BackendTraceType = "pattern_pii" BackendTraceVectorStore BackendTraceType = "vector_store" + BackendTraceFailover BackendTraceType = "failover" ) const ( From 6ab7a982f7fded92a778c09a1976b7270e879b4c Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 16:13:04 +0000 Subject: [PATCH 14/79] fix(failover): stop the preload loop from blocking the scheduler goroutine applyFailoverWarmTargets ran on the manager's single scheduler goroutine (Sync -> Tick), so a slow or hung PreloadModelByName call froze probing and fail-back for every chain. Keep the watchdog pin synchronous but run the preload loop in its own goroutine. Adds a seam (preloadModelByName) so a unit test can substitute a blocking loader and assert the callback still returns promptly. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/failover.go | 23 ++++++++++++++--- core/application/failover_test.go | 41 +++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 4 deletions(-) create mode 100644 core/application/failover_test.go diff --git a/core/application/failover.go b/core/application/failover.go index 3194e294b..f0b5c2c2b 100644 --- a/core/application/failover.go +++ b/core/application/failover.go @@ -5,13 +5,28 @@ import ( "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. func (a *Application) applyFailoverWarmTargets(warm []string) { a.SyncPinnedModelsToWatchdog() - for _, name := range warm { - if _, err := backend.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) + 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) + } } - } + }() } diff --git a/core/application/failover_test.go b/core/application/failover_test.go new file mode 100644 index 000000000..a9057425a --- /dev/null +++ b/core/application/failover_test.go @@ -0,0 +1,41 @@ +package application + +import ( + "context" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/model" + . "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") + }) +}) From 471bd8bf961da991d78020a8489e463198323a64 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:13:09 +0000 Subject: [PATCH 15/79] feat(failover): resolve chains per request and retry on the next target The retry wraps SetModelAndConfig, so each attempt binds the request again from a replayed body. A 5xx of a chain request is held back until the handler returns, and a streamed response is never retried. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/app.go | 1 + core/http/middleware/context_keys.go | 4 + core/http/middleware/failover.go | 282 ++++++++++++++++++++++++++ core/http/middleware/failover_test.go | 233 +++++++++++++++++++++ core/http/middleware/request.go | 22 +- 5 files changed, 540 insertions(+), 2 deletions(-) create mode 100644 core/http/middleware/failover.go create mode 100644 core/http/middleware/failover_test.go diff --git a/core/http/app.go b/core/http/app.go index f03522a45..db1317275 100644 --- a/core/http/app.go +++ b/core/http/app.go @@ -455,6 +455,7 @@ func API(application *application.Application) (*echo.Echo, error) { mcpJobsMw := auth.RequireFeature(application.AuthDB(), auth.FeatureMCPJobs) requestExtractor := httpMiddleware.NewRequestExtractor(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()) + requestExtractor.SetFailoverManager(application.FailoverManager()) // Register auth routes (login, callback, API keys, user management) routes.RegisterAuthRoutes(e, application) diff --git a/core/http/middleware/context_keys.go b/core/http/middleware/context_keys.go index d1983c882..f8b5582da 100644 --- a/core/http/middleware/context_keys.go +++ b/core/http/middleware/context_keys.go @@ -47,4 +47,8 @@ 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" ) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go new file mode 100644 index 000000000..930b5e209 --- /dev/null +++ b/core/http/middleware/failover.go @@ -0,0 +1,282 @@ +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()) + if err == nil { + 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 failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + 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 := appConfig != nil && appConfig.EnableTracing + for { + w := &failoverWriter{ResponseWriter: orig, active: active} + resp.Writer = w + 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 { + att.Succeed() + return nil + } + 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) + } + } +} + +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 +} + +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.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 +} + +func (r *replayBody) Read(p []byte) (int, error) { + n, err := r.src.Read(p) + if n > 0 && !r.overflow { + if r.buf.Len()+n > r.limit { + r.overflow = true + r.buf.Reset() + } 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 } + +// replay rewinds to the start of the body and keeps recording, so a third +// attempt can replay too. +func (r *replayBody) replay() io.ReadCloser { + data := bytes.Clone(r.buf.Bytes()) + rest := r.src + r.src = struct { + io.Reader + io.Closer + }{io.MultiReader(bytes.NewReader(data), rest), rest} + r.buf.Reset() + return r +} diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go new file mode 100644 index 000000000..94168d23f --- /dev/null +++ b/core/http/middleware/failover_test.go @@ -0,0 +1,233 @@ +package middleware + +import ( + "bytes" + "context" + "errors" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + + "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/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("failover chains in the request pipeline", func() { + var ( + app *echo.Echo + fm *failover.Manager + 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") + + 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) + 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) })) + // 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("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("leaves plain models untouched", func() { + behavior["plain"] = func(c echo.Context) error { + return c.JSON(http.StatusInternalServerError, map[string]string{"error": "boom"}) + } + rec := chat("plain") + Expect(rec.Code).To(Equal(http.StatusInternalServerError)) + Expect(rec.Body.String()).To(ContainSubstring("boom")) + Expect(rec.Header().Get(HeaderServedModel)).To(BeEmpty()) + }) +}) diff --git a/core/http/middleware/request.go b/core/http/middleware/request.go index 080a0b73c..86ab9a66c 100644 --- a/core/http/middleware/request.go +++ b/core/http/middleware/request.go @@ -12,6 +12,7 @@ import ( "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/templates" "github.com/mudler/LocalAI/pkg/distributedhdr" @@ -29,6 +30,7 @@ type RequestExtractor struct { modelConfigLoader *config.ModelConfigLoader modelLoader *model.ModelLoader applicationConfig *config.ApplicationConfig + failover *failover.Manager } func NewRequestExtractor(modelConfigLoader *config.ModelConfigLoader, modelLoader *model.ModelLoader, applicationConfig *config.ApplicationConfig) *RequestExtractor { @@ -122,7 +124,7 @@ func (re *RequestExtractor) BuildFilteredFirstAvailableDefaultModel(filterFn con // Otherwise, it's in its own method below for now func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIRequest) echo.MiddlewareFunc { return func(next echo.HandlerFunc) echo.HandlerFunc { - return func(c echo.Context) error { + return failoverRetry(re.applicationConfig, func(c echo.Context) error { input := initializer() if input == nil { return echo.NewHTTPError(http.StatusBadRequest, "unable to initialize body") @@ -194,6 +196,22 @@ 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 + } + // Check if the model is disabled if cfg != nil && cfg.IsDisabled() { return c.JSON(http.StatusForbidden, schema.ErrorResponse{ @@ -209,7 +227,7 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg) return next(c) - } + }) } } From f28d07b09db5a50e10f6f33828f9e7cd55da8088 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:24:53 +0000 Subject: [PATCH 16/79] fix(failover): spill capacity and disabled targets without tripping them An admission rejection or a disabled target moves the request to the next target through Attempt.Skip, which records no failure. A 4xx response no longer counts as a success. Requests skip body recording when no chain is configured, and stop it once the model is known not to be a chain. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/middleware/admission.go | 1 + core/http/middleware/context_keys.go | 5 ++ core/http/middleware/failover.go | 81 ++++++++++++++++++++------ core/http/middleware/failover_test.go | 79 ++++++++++++++++++++++++- core/http/middleware/request.go | 4 +- core/services/failover/manager.go | 33 +++++++++++ core/services/failover/manager_test.go | 25 ++++++++ 7 files changed, 208 insertions(+), 20 deletions(-) diff --git a/core/http/middleware/admission.go b/core/http/middleware/admission.go index c79066925..13c97c711 100644 --- a/core/http/middleware/admission.go +++ b/core/http/middleware/admission.go @@ -38,6 +38,7 @@ func AdmissionControl(limiter *admission.Limiter, events pii.EventStore) echo.Mi if !ok { retryAfter := admission.RetryAfter(cfg.Limits.RetryAfterSeconds) recordAdmissionRejection(events, cfg.Name, retryAfter) + c.Set(ContextKeyAdmissionRejected, true) c.Response().Header().Set("Retry-After", strconv.Itoa(int(retryAfter.Seconds()))) return c.JSON(http.StatusServiceUnavailable, map[string]any{ "error": map[string]any{ diff --git a/core/http/middleware/context_keys.go b/core/http/middleware/context_keys.go index f8b5582da..1e1f3886c 100644 --- a/core/http/middleware/context_keys.go +++ b/core/http/middleware/context_keys.go @@ -51,4 +51,9 @@ const ( // ContextKeyFailoverAttempt holds the *failoverState of a request whose // model is a failover chain. ContextKeyFailoverAttempt = "failover.attempt" + + // ContextKeyAdmissionRejected is set to true by AdmissionControl when it + // turns a request away because the model is at capacity. The failover + // retry sends such a request to the next target without tripping this one. + ContextKeyAdmissionRejected = "admission.rejected" ) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go index 930b5e209..f3744fdd9 100644 --- a/core/http/middleware/failover.go +++ b/core/http/middleware/failover.go @@ -49,6 +49,14 @@ func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, ch } for { cfg, err := re.loadFailoverTarget(st.attempt.Target()) + if err == nil && cfg.IsDisabled() { + // 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", chain.Name, cfg.Name) + } if err == nil { c.Set(ContextKeyRequestedModel, requested) c.Set(ContextKeyServedModel, cfg.Name) @@ -87,8 +95,13 @@ func setFailoverHeaders(h http.Header, att *failover.Attempt) { // 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 failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo.HandlerFunc { +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 { @@ -109,10 +122,11 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) return st != nil } - tracing := appConfig != nil && appConfig.EnableTracing + 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 { @@ -122,22 +136,34 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo att := st.attempt status := w.held if err == nil && status == 0 { - att.Succeed() + // A 4xx says nothing about the target's health. + if w.status < http.StatusBadRequest { + att.Succeed() + } return nil } - 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) + if rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool); rejected && !w.committed { + // The target is at capacity, not broken: spill this request to + // the next target without counting a failure. + 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 } - 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. @@ -154,6 +180,14 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo } } +// 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 @@ -186,6 +220,8 @@ type failoverWriter struct { 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) { @@ -197,6 +233,7 @@ func (w *failoverWriter) WriteHeader(code int) { return } w.committed = true + w.status = code w.ResponseWriter.WriteHeader(code) } @@ -248,14 +285,16 @@ type replayBody struct { 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 { + if n > 0 && !r.overflow && !r.stopped { if r.buf.Len()+n > r.limit { r.overflow = true - r.buf.Reset() + // A new buffer, not Reset: Reset keeps the memory. + r.buf = bytes.Buffer{} } else { r.buf.Write(p[:n]) } @@ -266,7 +305,13 @@ func (r *replayBody) Read(p []byte) (int, error) { 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 } +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. diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index 94168d23f..a16337545 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -12,11 +12,13 @@ import ( "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" @@ -27,6 +29,8 @@ 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 @@ -71,13 +75,17 @@ var _ = Describe("failover chains in the request pipeline", func() { 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") 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) + re = NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig) fm = failover.New(mcl) re.SetFailoverManager(fm) @@ -96,6 +104,10 @@ var _ = Describe("failover chains in the request pipeline", func() { 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, @@ -221,6 +233,71 @@ var _ = Describe("failover chains in the request pipeline", func() { 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("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"}) diff --git a/core/http/middleware/request.go b/core/http/middleware/request.go index 86ab9a66c..7bc702e20 100644 --- a/core/http/middleware/request.go +++ b/core/http/middleware/request.go @@ -124,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 failoverRetry(re.applicationConfig, 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") @@ -210,6 +210,8 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR }) } cfg = resolved + } else { + stopFailoverRecording(c) } // Check if the model is disabled diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 408c44621..ab130d1ee 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -201,6 +201,28 @@ func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { return c, ok } +// HasChains reports whether any failover chain is configured. The request +// path uses it to skip chain bookkeeping on installations without chains. +// Chains are synced lazily, so with none known yet the config source is +// consulted, which catches a chain added since the last sync. +func (m *Manager) HasChains() bool { + if m == nil { + return false + } + m.mu.Lock() + known := len(m.chains) > 0 + m.mu.Unlock() + if known { + return true + } + for _, c := range m.src.GetAllModelsConfigs() { + if c.IsFailover() { + return true + } + } + return false +} + func (m *Manager) chainLocked(name string) *chainState { if ch := m.chains[name]; ch != nil { return ch @@ -405,6 +427,17 @@ func (a *Attempt) Fail(err error) bool { return true } +// Skip moves to the next target without recording a failure, for a target +// that could not take this request although nothing is wrong with it (at +// capacity, disabled). It returns false when no target is left. +func (a *Attempt) Skip() bool { + if a.i+1 >= len(a.targets) { + return false + } + a.i++ + return true +} + // Report records err against the current target without moving on: the // response was already committed, so nothing is left to retry. func (a *Attempt) Report(err error) { a.m.ReportFailure(a.Target(), err) } diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 3bef7945e..d0736cf19 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -49,6 +49,31 @@ var _ = Describe("Manager", func() { Expect(st.Targets[1].Kind).To(Equal(KindLocal)) }) + It("skips to the next target without recording a failure", func() { + att, _ := m.Plan("chain") + Expect(att.Skip()).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + Expect(att.Skip()).To(BeFalse()) + Expect(att.Target()).To(Equal("b")) + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + + It("reports whether any chain is configured, including one added since the last sync", func() { + empty := New(newFakeSource(remote("a")), WithClock(clock)) + Expect(empty.HasChains()).To(BeFalse()) + lateSrc := newFakeSource(remote("a"), local("b")) + late := New(lateSrc, WithClock(clock)) + late.Sync() + Expect(late.HasChains()).To(BeFalse()) + lateSrc.Put(chainCfg("chain", nil, t("a"), t("b"))) + Expect(late.HasChains()).To(BeTrue()) + Expect(m.HasChains()).To(BeTrue()) + var nilManager *Manager + Expect(nilManager.HasChains()).To(BeFalse()) + }) + It("returns ErrChainNotFound for an unknown chain", func() { _, err := m.Plan("nope") Expect(errors.Is(err, ErrChainNotFound)).To(BeTrue()) From 2a1213f1e223fb65551d644bb0b7854c2004bf09 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:42:00 +0000 Subject: [PATCH 17/79] feat(failover): expose chain status, pins and events over REST and SSE Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/auth/helpers_test.go | 7 + core/http/auth/middleware_test.go | 35 ++++ .../endpoints/localai/api_instructions.go | 6 + .../localai/api_instructions_test.go | 3 +- core/http/endpoints/localai/failover.go | 169 ++++++++++++++++++ core/http/endpoints/localai/failover_test.go | 110 ++++++++++++ core/http/routes/localai.go | 10 ++ 7 files changed, 339 insertions(+), 1 deletion(-) create mode 100644 core/http/endpoints/localai/failover.go create mode 100644 core/http/endpoints/localai/failover_test.go diff --git a/core/http/auth/helpers_test.go b/core/http/auth/helpers_test.go index 1e31ac27f..1b52d4500 100644 --- a/core/http/auth/helpers_test.go +++ b/core/http/auth/helpers_test.go @@ -78,6 +78,9 @@ func newAuthTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo e.GET("/api/settings", ok) e.POST("/api/settings", ok) + // Failover chain reads and the event stream: standard auth, no admin gate. + e.GET("/api/failover", ok) + // Auth routes (exempt) e.GET("/api/auth/status", ok) e.GET("/api/auth/github/login", ok) @@ -139,6 +142,10 @@ func newAdminTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Ech e.POST("/backends/apply", ok, adminMw) e.GET("/api/agents", ok, adminMw) + // Failover chain pin/unpin (admin only) + e.POST("/api/failover/:chain/pin", ok, adminMw) + e.DELETE("/api/failover/:chain/pin", ok, adminMw) + // Trace/log endpoints (admin only) e.GET("/api/traces", ok, adminMw) e.POST("/api/traces/clear", ok, adminMw) diff --git a/core/http/auth/middleware_test.go b/core/http/auth/middleware_test.go index bdfadaafe..bcbba0096 100644 --- a/core/http/auth/middleware_test.go +++ b/core/http/auth/middleware_test.go @@ -223,6 +223,17 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusOK)) }) + It("allows requests to the failover status endpoint with a valid session", func() { + sessionID := createTestSession(db, user.ID) + rec := doRequest(app, http.MethodGet, "/api/failover", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("returns 401 for the failover status endpoint without credentials", func() { + rec := doRequest(app, http.MethodGet, "/api/failover") + Expect(rec.Code).To(Equal(http.StatusUnauthorized)) + }) + It("allows authenticated users to call moderation by default", func() { sessionID := createTestSession(db, user.ID) rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID)) @@ -526,6 +537,30 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusForbidden)) }) + It("allows admin to pin and unpin failover chains", func() { + admin := createTestUser(db, "admin5@example.com", auth.RoleAdmin, auth.ProviderGitHub) + sessionID := createTestSession(db, admin.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("blocks non-admin from pinning or unpinning failover chains", func() { + user := createTestUser(db, "user5@example.com", auth.RoleUser, auth.ProviderGitHub) + sessionID := createTestSession(db, user.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + }) + It("allows user to access regular inference endpoints", func() { user := createTestUser(db, "user@example.com", auth.RoleUser, auth.ProviderGitHub) sessionID := createTestSession(db, user.ID) diff --git a/core/http/endpoints/localai/api_instructions.go b/core/http/endpoints/localai/api_instructions.go index f8af73de2..8d0ea6d2f 100644 --- a/core/http/endpoints/localai/api_instructions.go +++ b/core/http/endpoints/localai/api_instructions.go @@ -129,6 +129,12 @@ var instructionDefs = []instructionDef{ Tags: []string{"middleware", "pii", "router"}, Intro: "GET /api/middleware/status is the single round-trip the /app/middleware admin page reads to render the current state: every model's resolved PII enabled state and the NER detector models it references, recent event count, and the active routing models with their classifier configurations. Admin-only (the synthetic local user is admin in no-auth mode). PII detection policy is edited on each detector model's `pii_detection:` block via the model-config tools/UI — there is no global pattern set to mutate. GET /api/router/decisions returns the routing decision log filtered by correlation_id / user_id / router_model. The same surface is exposed as MCP tools (`get_middleware_status`, `get_pii_events`, `get_router_decisions`) for agent-driven inspection.", }, + { + Name: "failover", + Description: "Model failover chains: target health, pinning and switch events", + Tags: []string{"failover"}, + Intro: "A failover chain is a model config with a failover block. Requests for the chain name are served by its highest-priority healthy target; the X-LocalAI-Served-Model response header names it. Subscribe to GET /api/failover/events (SSE) to follow switches.", + }, { Name: "intelligent-routing", Description: "Per-model `router:` configuration that classifies requests and rewrites the served model", diff --git a/core/http/endpoints/localai/api_instructions_test.go b/core/http/endpoints/localai/api_instructions_test.go index 710d4d982..f42e1c92d 100644 --- a/core/http/endpoints/localai/api_instructions_test.go +++ b/core/http/endpoints/localai/api_instructions_test.go @@ -39,7 +39,7 @@ var _ = Describe("API Instructions Endpoints", func() { instructions, ok := resp["instructions"].([]any) Expect(ok).To(BeTrue()) - Expect(instructions).To(HaveLen(19)) + Expect(instructions).To(HaveLen(20)) // Verify each instruction has required fields and correct URL format for _, s := range instructions { @@ -81,6 +81,7 @@ var _ = Describe("API Instructions Endpoints", func() { "intelligent-routing", "voice-library", "3d", + "failover", )) }) }) diff --git a/core/http/endpoints/localai/failover.go b/core/http/endpoints/localai/failover.go new file mode 100644 index 000000000..608fe69f2 --- /dev/null +++ b/core/http/endpoints/localai/failover.go @@ -0,0 +1,169 @@ +package localai + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" +) + +type FailoverChainsResponse struct { + Chains []failover.ChainStatus `json:"chains"` +} + +type FailoverPinRequest struct { + Target string `json:"target"` +} + +func failoverError(c echo.Context, code int, msg string) error { + return c.JSON(code, schema.ErrorResponse{Error: &schema.APIError{Message: msg, Code: code, Type: "failover_error"}}) +} + +// ListFailoverChainsEndpoint lists failover chains and the health of their targets +// +// @Summary List failover chains and the health of their targets +// @Tags failover +// @Produce json +// @Success 200 {object} FailoverChainsResponse +// @Router /api/failover [get] +func ListFailoverChainsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + return c.JSON(http.StatusOK, FailoverChainsResponse{Chains: fm.Status()}) + } +} + +// GetFailoverChainEndpoint returns one failover chain +// +// @Summary Get one failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain} [get] +func GetFailoverChainEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + st, ok := fm.ChainStatus(c.Param("chain")) + if !ok { + return failoverError(c, http.StatusNotFound, fmt.Sprintf("failover chain %q not found", c.Param("chain"))) + } + return c.JSON(http.StatusOK, st) + } +} + +// PinFailoverTargetEndpoint forces a chain to one target +// +// @Summary Pin a failover chain to one target +// @Tags failover +// @Accept json +// @Produce json +// @Param chain path string true "Chain name" +// @Param request body FailoverPinRequest true "Target to pin" +// @Success 200 {object} failover.ChainStatus +// @Failure 400 {object} schema.ErrorResponse +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [post] +func PinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + var req FailoverPinRequest + if err := c.Bind(&req); err != nil || req.Target == "" { + return failoverError(c, http.StatusBadRequest, "request body must set \"target\"") + } + chain := c.Param("chain") + if err := fm.Pin(chain, req.Target); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +// UnpinFailoverTargetEndpoint removes a pin +// +// @Summary Remove the pin from a failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [delete] +func UnpinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + chain := c.Param("chain") + if err := fm.Unpin(chain); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +func pinError(c echo.Context, err error) error { + switch { + case errors.Is(err, failover.ErrChainNotFound): + return failoverError(c, http.StatusNotFound, err.Error()) + case errors.Is(err, failover.ErrTargetNotInChain): + return failoverError(c, http.StatusBadRequest, err.Error()) + } + return failoverError(c, http.StatusInternalServerError, err.Error()) +} + +// FailoverEventsEndpoint streams failover events +// +// @Summary Stream failover events (server-sent events) +// @Description The first event is "snapshot" with the full state, then "chain.switched" and "target.state" events. +// @Tags failover +// @Produce text/event-stream +// @Success 200 +// @Router /api/failover/events [get] +func FailoverEventsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + // Subscribe before the snapshot so no event falls between the two. + events, cancel := fm.Subscribe(64) + defer cancel() + w := c.Response() + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + send := func(name string, v any) error { + data, err := json.Marshal(v) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", name, data); err != nil { + return err + } + w.Flush() + return nil + } + if err := send("snapshot", FailoverChainsResponse{Chains: fm.Status()}); err != nil { + return nil + } + keepalive := time.NewTicker(15 * time.Second) + defer keepalive.Stop() + for { + select { + case <-c.Request().Context().Done(): + return nil + case <-keepalive.C: + if _, err := fmt.Fprint(w, ": keepalive\n\n"); err != nil { + return nil + } + w.Flush() + case ev, ok := <-events: + if !ok { + return nil + } + if err := send(string(ev.Type), ev); err != nil { + return nil + } + } + } + } +} diff --git a/core/http/endpoints/localai/failover_test.go b/core/http/endpoints/localai/failover_test.go new file mode 100644 index 000000000..fd3b6a733 --- /dev/null +++ b/core/http/endpoints/localai/failover_test.go @@ -0,0 +1,110 @@ +package localai + +import ( + "bufio" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type mapSource map[string]config.ModelConfig + +func (s mapSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok } +func (s mapSource) GetAllModelsConfigs() []config.ModelConfig { + var out []config.ModelConfig + for _, c := range s { + out = append(out, c) + } + return out +} + +var _ = Describe("failover endpoints", func() { + var ( + e *echo.Echo + fm *failover.Manager + ) + + BeforeEach(func() { + src := mapSource{ + "a": {Name: "a", Backend: "cloud-proxy"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}}, + } + fm = failover.New(src) + e = echo.New() + e.GET("/api/failover", ListFailoverChainsEndpoint(fm)) + e.GET("/api/failover/events", FailoverEventsEndpoint(fm)) + e.GET("/api/failover/:chain", GetFailoverChainEndpoint(fm)) + e.POST("/api/failover/:chain/pin", PinFailoverTargetEndpoint(fm)) + e.DELETE("/api/failover/:chain/pin", UnpinFailoverTargetEndpoint(fm)) + }) + + do := func(method, path, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(method, path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + return rec + } + + It("lists chains", func() { + rec := do(http.MethodGet, "/api/failover", "") + Expect(rec.Code).To(Equal(http.StatusOK)) + var out FailoverChainsResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &out)).To(Succeed()) + Expect(out.Chains).To(HaveLen(1)) + Expect(out.Chains[0].Active).To(Equal("a")) + }) + + It("gets one chain or 404", func() { + Expect(do(http.MethodGet, "/api/failover/chain", "").Code).To(Equal(http.StatusOK)) + Expect(do(http.MethodGet, "/api/failover/nope", "").Code).To(Equal(http.StatusNotFound)) + }) + + It("pins and unpins", func() { + rec := do(http.MethodPost, "/api/failover/chain/pin", `{"target":"b"}`) + Expect(rec.Code).To(Equal(http.StatusOK)) + st, _ := fm.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{"target":"zzz"}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/nope/pin", `{"target":"a"}`).Code).To(Equal(http.StatusNotFound)) + Expect(do(http.MethodDelete, "/api/failover/chain/pin", "").Code).To(Equal(http.StatusOK)) + st, _ = fm.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + + It("streams a snapshot, then switch events", func() { + srv := httptest.NewServer(e) + defer srv.Close() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil) + resp, err := http.DefaultClient.Do(req) + Expect(err).ToNot(HaveOccurred()) + defer resp.Body.Close() + Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream")) + r := bufio.NewReader(resp.Body) + next := func() string { + for { + line, err := r.ReadString('\n') + Expect(err).ToNot(HaveOccurred()) + if strings.HasPrefix(line, "event: ") { + return strings.TrimSpace(strings.TrimPrefix(line, "event: ")) + } + } + } + Expect(next()).To(Equal("snapshot")) + Expect(fm.Pin("chain", "b")).To(Succeed()) + Expect(next()).To(Equal("chain.switched")) + }) +}) diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index 1d3c7adaf..7a6f901ec 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -171,6 +171,16 @@ func RegisterLocalAIRoutes(router *echo.Echo, return nil })) + // Failover chains: reads and the event stream use standard auth (any + // authenticated caller may watch chain health); pin/unpin are admin-only + // since they override the routing decision for every caller of the chain. + fm := app.FailoverManager() + router.GET("/api/failover", localai.ListFailoverChainsEndpoint(fm)) + router.GET("/api/failover/events", localai.FailoverEventsEndpoint(fm)) + router.GET("/api/failover/:chain", localai.GetFailoverChainEndpoint(fm)) + router.POST("/api/failover/:chain/pin", localai.PinFailoverTargetEndpoint(fm), adminMiddleware) + router.DELETE("/api/failover/:chain/pin", localai.UnpinFailoverTargetEndpoint(fm), adminMiddleware) + voiceProfiles := app.VoiceProfileStore() router.GET("/api/voice-profiles", localai.ListVoiceProfilesEndpoint(voiceProfiles)) router.GET("/api/voice-profiles/:id/audio", localai.ServeVoiceProfileAudioEndpoint(voiceProfiles)) From 74775a820bacbe2349d35f1b3a4ae699aa8272c2 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:50:54 +0000 Subject: [PATCH 18/79] test(failover): cover retry per endpoint family and remote fail-back Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- tests/e2e/cloud_proxy_helpers_test.go | 24 ++++ tests/e2e/e2e_failover_test.go | 172 ++++++++++++++++++++++++++ tests/e2e/e2e_suite_test.go | 33 +++++ tests/e2e/mock-backend/main.go | 4 + 4 files changed, 233 insertions(+) create mode 100644 tests/e2e/e2e_failover_test.go diff --git a/tests/e2e/cloud_proxy_helpers_test.go b/tests/e2e/cloud_proxy_helpers_test.go index 819d9aa08..f98049b4c 100644 --- a/tests/e2e/cloud_proxy_helpers_test.go +++ b/tests/e2e/cloud_proxy_helpers_test.go @@ -5,6 +5,7 @@ import ( "io" "net/http" "net/http/httptest" + "slices" "strings" "sync" "sync/atomic" @@ -46,6 +47,9 @@ type fakeOpenAIUpstreamServer struct { mu sync.Mutex script func(req []byte) (status int, body string, contentType string) + // models is what GET /v1/models lists: failover liveness probes check + // that the upstream still serves the target's model. + models []string } func newFakeOpenAIUpstream() *fakeOpenAIUpstreamServer { @@ -59,6 +63,20 @@ func newFakeOpenAIUpstream() *fakeOpenAIUpstreamServer { } func (f *fakeOpenAIUpstreamServer) serve(w http.ResponseWriter, r *http.Request) { + // Answered before recording: periodic probes must not overwrite the + // request a spec is about to assert on. + if r.Method == http.MethodGet && r.URL.Path == "/v1/models" { + f.mu.Lock() + ids := slices.Clone(f.models) + f.mu.Unlock() + data := make([]map[string]string, 0, len(ids)) + for _, id := range ids { + data = append(data, map[string]string{"id": id}) + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"data": data}) + return + } atomic.AddInt32(&f.recorder.RequestHits, 1) body, _ := io.ReadAll(r.Body) f.recorder.mu.Lock() @@ -80,6 +98,12 @@ func (f *fakeOpenAIUpstreamServer) serve(w http.ResponseWriter, r *http.Request) func (f *fakeOpenAIUpstreamServer) URL() string { return f.srv.URL } func (f *fakeOpenAIUpstreamServer) Close() { f.srv.Close() } +func (f *fakeOpenAIUpstreamServer) SetModels(ids ...string) { + f.mu.Lock() + defer f.mu.Unlock() + f.models = ids +} + func (f *fakeOpenAIUpstreamServer) SetScript(script func(req []byte) (status int, body string, contentType string)) { f.mu.Lock() defer f.mu.Unlock() diff --git a/tests/e2e/e2e_failover_test.go b/tests/e2e/e2e_failover_test.go new file mode 100644 index 000000000..540ea73c5 --- /dev/null +++ b/tests/e2e/e2e_failover_test.go @@ -0,0 +1,172 @@ +package e2e_test + +import ( + "bytes" + "encoding/json" + "io" + "mime/multipart" + "net/http" + "os" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gopkg.in/yaml.v3" +) + +var _ = Describe("Failover chains", Label("failover"), func() { + postJSON := func(path string, body map[string]any) *http.Response { + b, err := json.Marshal(body) + Expect(err).ToNot(HaveOccurred()) + resp, err := http.Post(apiURL+path, "application/json", bytes.NewReader(b)) + Expect(err).ToNot(HaveOccurred()) + return resp + } + expectServedByFallback := func(resp *http.Response) { + defer func() { _ = resp.Body.Close() }() + body, _ := io.ReadAll(resp.Body) + Expect(resp.StatusCode).To(BeNumerically("<", 300), "%s\nheaders: %v", body, resp.Header) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("failover-fallback")) + Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) + } + + // The entry name is the chain suffix: chain- is written by the suite. + DescribeTable("retries every endpoint family on the next target", + func(path string, body func(model string) map[string]any) { + expectServedByFallback(postJSON(path, body("chain-"+CurrentSpecReport().LeafNodeText))) + }, + Entry("chat", "/chat/completions", func(m string) map[string]any { + return map[string]any{"model": m, "messages": []map[string]string{{"role": "user", "content": "hi"}}} + }), + Entry("completion", "/completions", func(m string) map[string]any { + return map[string]any{"model": m, "prompt": "hi"} + }), + Entry("embeddings", "/embeddings", func(m string) map[string]any { + return map[string]any{"model": m, "input": "hi"} + }), + Entry("tts", "/audio/speech", func(m string) map[string]any { + return map[string]any{"model": m, "input": "hi", "voice": "default"} + }), + Entry("image", "/images/generations", func(m string) map[string]any { + return map[string]any{"model": m, "prompt": "a cat", "size": "256x256"} + }), + Entry("rerank", "/rerank", func(m string) map[string]any { + return map[string]any{"model": m, "query": "q", "documents": []string{"a", "b"}} + }), + Entry("vad", "/vad", func(m string) map[string]any { + return map[string]any{"model": m, "audio": []float32{0, 0, 0, 0}} + }), + ) + + It("retries transcription with the multipart body", func() { + var body bytes.Buffer + mw := multipart.NewWriter(&body) + Expect(mw.WriteField("model", "chain-transcription")).To(Succeed()) + fw, err := mw.CreateFormFile("file", "a.wav") + Expect(err).ToNot(HaveOccurred()) + // 200 ms of 16 kHz mono silence. + _, err = fw.Write(wavFromPCM(make([]byte, 6400), 16000)) + Expect(err).ToNot(HaveOccurred()) + Expect(mw.Close()).To(Succeed()) + resp, err := http.Post(apiURL+"/audio/transcriptions", mw.FormDataContentType(), &body) + Expect(err).ToNot(HaveOccurred()) + expectServedByFallback(resp) + }) + + Describe("remote targets", Ordered, func() { + var up1, up2 *fakeOpenAIUpstreamServer + + chatReply := func([]byte) (int, string, string) { + return 200, `{"id":"x","object":"chat.completion","choices":[{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}`, "application/json" + } + + BeforeAll(func() { + if cloudProxyPath == "" { + Skip("cloud-proxy backend binary not built (make build-cloud-proxy-backend)") + } + up1, up2 = newFakeOpenAIUpstream(), newFakeOpenAIUpstream() + DeferCleanup(up1.Close) + DeferCleanup(up2.Close) + up1.SetModels("up-1") + up2.SetModels("up-2") + registerFailoverRemoteModels(up1.URL(), up2.URL()) + }) + + It("fails over when the primary upstream errors and fails back when it recovers", func() { + up1.SetScript(func([]byte) (int, string, string) { + return 503, `{"error":"no healthy nodes"}`, "application/json" + }) + up2.SetScript(chatReply) + + resp := postJSON("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) + defer func() { _ = resp.Body.Close() }() + body, _ := io.ReadAll(resp.Body) + Expect(resp.StatusCode).To(Equal(200), string(body)) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-2")) + Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) + Expect(chainActive("chain-remote")).To(Equal("up-2")) + + up1.SetScript(chatReply) + Eventually(func() string { return chainActive("chain-remote") }, 30*time.Second, 500*time.Millisecond). + Should(Equal("up-1")) + + resp2 := postJSON("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) + defer func() { _ = resp2.Body.Close() }() + Expect(resp2.StatusCode).To(Equal(200)) + Expect(resp2.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-1")) + Expect(resp2.Header.Get("X-LocalAI-Failover")).To(BeEmpty()) + }) + }) +}) + +// chainActive returns the active target of a chain as the REST status reports +// it, or "" when the status cannot be read. +func chainActive(chain string) string { + r, err := http.Get(anthropicBaseURL + "/api/failover/" + chain) + if err != nil { + return "" + } + defer func() { _ = r.Body.Close() }() + var st struct { + Active string `json:"active"` + } + _ = json.NewDecoder(r.Body).Decode(&st) + return st.Active +} + +// registerFailoverRemoteModels registers two cloud-proxy passthrough models +// (up-1, up-2) and a chain over them. The upstream URLs exist only at runtime, +// so the YAMLs are written after startup and the loader re-reads the models +// directory; the failover manager picks the chain up on its next tick. +func registerFailoverRemoteModels(url1, url2 string) { + proxyModel := func(name, upstream string) map[string]any { + return map[string]any{ + "name": name, + "backend": "cloud-proxy", + "parameters": map[string]any{"model": name + ".bin"}, + "proxy": map[string]any{ + "mode": "passthrough", + "provider": "openai", + "upstream_url": upstream + "/v1/chat/completions", + "api_key_env": "CLOUD_PROXY_E2E_OPENAI_KEY", + }, + } + } + chain := map[string]any{ + "name": "chain-remote", + "failover": map[string]any{ + "targets": []map[string]any{{"model": "up-1"}, {"model": "up-2"}}, + "probe": map[string]any{"interval": "1s"}, + "recovery": map[string]any{"probes": 2, "min_dwell": "2s"}, + }, + } + for _, cfg := range []map[string]any{proxyModel("up-1", url1), proxyModel("up-2", url2), chain} { + data, err := yaml.Marshal(cfg) + Expect(err).ToNot(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) + } + Expect(localAIApp.ModelConfigLoader().LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + Eventually(func() string { return chainActive("chain-remote") }, 10*time.Second, 200*time.Millisecond). + Should(Equal("up-1")) +} diff --git a/tests/e2e/e2e_suite_test.go b/tests/e2e/e2e_suite_test.go index 31d8d0d9c..b27d25f64 100644 --- a/tests/e2e/e2e_suite_test.go +++ b/tests/e2e/e2e_suite_test.go @@ -118,6 +118,39 @@ var _ = BeforeSuite(func() { Expect(err).ToNot(HaveOccurred()) Expect(os.WriteFile(configPath, configYAML, 0644)).To(Succeed()) + // Failover chains, one per endpoint family: target 0 is a mock model whose + // load always fails (the mock rejects models named fail-load*), so every + // request exercises the retry onto failover-fallback. The fallback's model + // file must exist: the failover manager's liveness probe marks a local + // target with a missing file down, and mock-model.bin is never created. + Expect(os.WriteFile(filepath.Join(modelsPath, "failover-fallback.bin"), nil, 0644)).To(Succeed()) + fallbackData, err := yaml.Marshal(map[string]any{ + "name": "failover-fallback", + "backend": "mock-backend", + "parameters": map[string]any{"model": "failover-fallback.bin"}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(modelsPath, "failover-fallback.yaml"), fallbackData, 0644)).To(Succeed()) + for _, family := range []string{"chat", "completion", "embeddings", "transcription", "tts", "image", "rerank", "vad"} { + for _, cfg := range []map[string]any{ + { + "name": "fail-" + family, + "backend": "mock-backend", + "parameters": map[string]any{"model": "fail-load-" + family}, + }, + { + "name": "chain-" + family, + "failover": map[string]any{ + "targets": []map[string]any{{"model": "fail-" + family}, {"model": "failover-fallback"}}, + }, + }, + } { + data, err := yaml.Marshal(cfg) + Expect(err).ToNot(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) + } + } + // Create model config for autoparser tests (NoGrammar so tool calls // are driven entirely by the backend's ChatDeltas, not grammar enforcement) autoparserConfig := map[string]any{ diff --git a/tests/e2e/mock-backend/main.go b/tests/e2e/mock-backend/main.go index 36520cb55..0d842b42e 100644 --- a/tests/e2e/mock-backend/main.go +++ b/tests/e2e/mock-backend/main.go @@ -93,6 +93,10 @@ func (m *MockBackend) LoadModel(ctx context.Context, in *pb.ModelOptions) (*pb.R "draft_model", in.DraftModel, "mmproj", in.MMProj) recordLoadParams(in) + // Lets e2e specs build a failover target whose backend cannot load. + if strings.HasPrefix(in.Model, "fail-load") { + return &pb.Result{Message: "mock: load failure", Success: false}, nil + } return &pb.Result{ Message: "Model loaded successfully (mocked)", Success: true, From 3e8ca016eeb647e5cc3406baf089a8c244c5319c Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:57:17 +0000 Subject: [PATCH 19/79] fix(failover): stop judging cold local targets by their model file A missing file is not a reliable signal: models download on first use, some backends need no file, and dotted names like Phi-3.5-mini look like paths. Marking such a fallback down removed the retry a chain exists for. Cold targets are now judged only by real requests. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/startup.go | 2 +- core/services/failover/prober.go | 39 ++++++------------- core/services/failover/prober_test.go | 26 ++++++------- ...2026-09-26-model-failover-chains-design.md | 2 +- tests/e2e/e2e_failover_test.go | 8 ++-- tests/e2e/e2e_suite_test.go | 14 +------ 6 files changed, 31 insertions(+), 60 deletions(-) diff --git a/core/application/startup.go b/core/application/startup.go index 1382835f6..a700615e2 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -259,7 +259,7 @@ func New(opts ...config.AppOption) (*Application, error) { application.failoverManager = failover.New(application.ModelConfigLoader(), failover.WithProber(failover.NewProber(func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) { return application.ModelLoader().Load(backend.ModelOptions(cfg, options)...) - }, application.ModelLoader().ModelPath)), + })), failover.WithOnWarmChanged(application.applyFailoverWarmTargets), ) diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index b3055c17f..592d9ca41 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -8,12 +8,9 @@ import ( "errors" "fmt" "io" - "io/fs" "mime/multipart" "net/http" "net/url" - "os" - "path/filepath" "strings" "github.com/mudler/LocalAI/core/config" @@ -27,13 +24,12 @@ type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, e // DefaultProber probes remote targets over the upstream's OpenAI-compatible // API and local targets through their gRPC backend. type DefaultProber struct { - HTTP *http.Client - Load LoadFunc - ModelPath string + HTTP *http.Client + Load LoadFunc } -func NewProber(load LoadFunc, modelPath string) *DefaultProber { - return &DefaultProber{HTTP: &http.Client{}, Load: load, ModelPath: modelPath} +func NewProber(load LoadFunc) *DefaultProber { + return &DefaultProber{HTTP: &http.Client{}, Load: load} } func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { @@ -43,7 +39,13 @@ func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, ki case warm: return p.localHealth(ctx, cfg) } - return p.coldLiveness(cfg) + // A cold target is judged only by real requests: loading it only to probe + // it could evict other models, and no cheaper check is reliable. A missing + // model file is not one: models download on first use, some backends need + // no file, and a dotted name like "Phi-3.5-mini" looks like a file path. A + // false "down" would take away the very fallback the chain exists for, + // while a real load failure still trips the target and the request moves on. + return nil } func (p *DefaultProber) Inference(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { @@ -262,22 +264,3 @@ func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConf // transcription, so a real request adds little here. return p.localHealth(ctx, cfg) } - -// coldLiveness checks the model file without loading the model. -func (p *DefaultProber) coldLiveness(cfg config.ModelConfig) error { - f := cfg.Model - if f == "" || p.ModelPath == "" || strings.Contains(f, "://") { - return nil - } - path := f - if !filepath.IsAbs(path) { - path = filepath.Join(p.ModelPath, f) - } - if _, err := os.Stat(path); err != nil { - if errors.Is(err, fs.ErrNotExist) && filepath.Ext(f) == "" { - return nil // a repository id, downloaded on demand - } - return fmt.Errorf("model file %s: %w", f, err) - } - return nil -} diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index 2b8b7f42d..0f65f8dc4 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -7,8 +7,6 @@ import ( "io" "net/http" "net/http/httptest" - "os" - "path/filepath" "sync" "github.com/mudler/LocalAI/core/config" @@ -77,7 +75,7 @@ var _ = Describe("DefaultProber", func() { BeforeEach(func() { up = newFakeUpstream() DeferCleanup(up.srv.Close) - p = NewProber(nil, "") + p = NewProber(nil) }) proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { @@ -140,7 +138,7 @@ var _ = Describe("DefaultProber", func() { It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { b := &fakeBackend{healthy: true} - p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }, "") + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }) c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) @@ -152,18 +150,18 @@ var _ = Describe("DefaultProber", func() { Expect(p.Inference(ctx, c, KindLocal, true)).To(HaveOccurred()) }) - It("checks the model file for cold local liveness without loading", func() { - dir := GinkgoT().TempDir() + It("passes cold local liveness without a model file and without loading", func() { p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { Fail("cold liveness must not load the model") return nil, nil - }, dir) - c := config.ModelConfig{Name: "cold", Backend: "llama-cpp"} - c.Model = "weights.gguf" - Expect(p.Liveness(ctx, c, KindLocal, false)).To(HaveOccurred()) - Expect(os.WriteFile(filepath.Join(dir, "weights.gguf"), []byte("x"), 0o600)).To(Succeed()) - Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) - c.Model = "org/some-hf-repo" // no extension: downloaded on demand - Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + }) + // None of these files exist: a missing file says nothing about whether + // the target can serve (download on first use, dotted names, backends + // that need no file). Only a real request may trip a cold target. + for _, model := range []string{"weights.gguf", "Phi-3.5-mini", "org/some-hf-repo", ""} { + c := config.ModelConfig{Name: "cold", Backend: "llama-cpp"} + c.Model = model + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed(), model) + } }) }) diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 11ef80e4f..551af9ab7 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -152,7 +152,7 @@ Chain states: |---|---|---| | remote | `GET /v1/models` returns 2xx and lists the upstream model. `` is the scheme and host of `proxy.upstream_url` plus any path prefix before `/v1`. The upstream model is `proxy.upstream_model`, or the target name when it is empty. `/v1/models` works on any OpenAI-compatible upstream, and `/readyz` exists only on LocalAI. | one minimal real request, chosen by usecase | | local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. | -| local, cold | the model file exists (skipped for URLs and repository ids). The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | +| local, cold | none: a cold target is judged only by real requests; it is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | Minimal requests by usecase: diff --git a/tests/e2e/e2e_failover_test.go b/tests/e2e/e2e_failover_test.go index 540ea73c5..67e1227b4 100644 --- a/tests/e2e/e2e_failover_test.go +++ b/tests/e2e/e2e_failover_test.go @@ -23,18 +23,18 @@ var _ = Describe("Failover chains", Label("failover"), func() { Expect(err).ToNot(HaveOccurred()) return resp } - expectServedByFallback := func(resp *http.Response) { + expectServedByMock := func(resp *http.Response) { defer func() { _ = resp.Body.Close() }() body, _ := io.ReadAll(resp.Body) Expect(resp.StatusCode).To(BeNumerically("<", 300), "%s\nheaders: %v", body, resp.Header) - Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("failover-fallback")) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("mock-model")) Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) } // The entry name is the chain suffix: chain- is written by the suite. DescribeTable("retries every endpoint family on the next target", func(path string, body func(model string) map[string]any) { - expectServedByFallback(postJSON(path, body("chain-"+CurrentSpecReport().LeafNodeText))) + expectServedByMock(postJSON(path, body("chain-"+CurrentSpecReport().LeafNodeText))) }, Entry("chat", "/chat/completions", func(m string) map[string]any { return map[string]any{"model": m, "messages": []map[string]string{{"role": "user", "content": "hi"}}} @@ -71,7 +71,7 @@ var _ = Describe("Failover chains", Label("failover"), func() { Expect(mw.Close()).To(Succeed()) resp, err := http.Post(apiURL+"/audio/transcriptions", mw.FormDataContentType(), &body) Expect(err).ToNot(HaveOccurred()) - expectServedByFallback(resp) + expectServedByMock(resp) }) Describe("remote targets", Ordered, func() { diff --git a/tests/e2e/e2e_suite_test.go b/tests/e2e/e2e_suite_test.go index b27d25f64..a89318dfc 100644 --- a/tests/e2e/e2e_suite_test.go +++ b/tests/e2e/e2e_suite_test.go @@ -120,17 +120,7 @@ var _ = BeforeSuite(func() { // Failover chains, one per endpoint family: target 0 is a mock model whose // load always fails (the mock rejects models named fail-load*), so every - // request exercises the retry onto failover-fallback. The fallback's model - // file must exist: the failover manager's liveness probe marks a local - // target with a missing file down, and mock-model.bin is never created. - Expect(os.WriteFile(filepath.Join(modelsPath, "failover-fallback.bin"), nil, 0644)).To(Succeed()) - fallbackData, err := yaml.Marshal(map[string]any{ - "name": "failover-fallback", - "backend": "mock-backend", - "parameters": map[string]any{"model": "failover-fallback.bin"}, - }) - Expect(err).ToNot(HaveOccurred()) - Expect(os.WriteFile(filepath.Join(modelsPath, "failover-fallback.yaml"), fallbackData, 0644)).To(Succeed()) + // request exercises the retry onto mock-model. for _, family := range []string{"chat", "completion", "embeddings", "transcription", "tts", "image", "rerank", "vad"} { for _, cfg := range []map[string]any{ { @@ -141,7 +131,7 @@ var _ = BeforeSuite(func() { { "name": "chain-" + family, "failover": map[string]any{ - "targets": []map[string]any{{"model": "fail-" + family}, {"model": "failover-fallback"}}, + "targets": []map[string]any{{"model": "fail-" + family}, {"model": "mock-model"}}, }, }, } { From ea55e2ffc397b6eadb6955177343b6cebdcbeff0 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 18:07:56 +0000 Subject: [PATCH 20/79] feat(failover): switch realtime pipeline stages per call A stage that names a chain is resolved on every call, so a switch keeps the session and its conversation. Clients get localai.model.failover events at session start and on every switch. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/endpoints/openai/realtime.go | 8 + .../openai/realtime_classifier_test.go | 10 +- .../endpoints/openai/realtime_doubles_test.go | 21 +- .../endpoints/openai/realtime_failover.go | 70 ++++++ .../openai/realtime_failover_test.go | 99 ++++++++ core/http/endpoints/openai/realtime_model.go | 216 +++++++++++++++++- .../openai/realtime_semantic_vad_test.go | 6 +- .../openai/realtime_sound_detection_test.go | 4 +- .../endpoints/openai/realtime_stream_test.go | 12 +- .../realtime_voicegate_integration_test.go | 2 +- core/http/endpoints/openai/types/failover.go | 46 ++++ .../endpoints/openai/types/server_events.go | 6 +- tests/e2e/e2e_suite_test.go | 38 +++ tests/e2e/realtime_ws_test.go | 100 ++++++++ 14 files changed, 603 insertions(+), 35 deletions(-) create mode 100644 core/http/endpoints/openai/realtime_failover.go create mode 100644 core/http/endpoints/openai/realtime_failover_test.go create mode 100644 core/http/endpoints/openai/types/failover.go diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index 9db5f899f..fe2c15068 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -754,6 +754,14 @@ 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. + if wrapped, ok := m.(*wrappedModel); ok && wrapped.failover != nil && len(wrapped.stageChains) > 0 { + stopFailoverEvents := startFailoverEvents(t, wrapped.failover, wrapped.stageChains) + defer stopFailoverEvents() + } + var ( msg []byte wg sync.WaitGroup diff --git a/core/http/endpoints/openai/realtime_classifier_test.go b/core/http/endpoints/openai/realtime_classifier_test.go index b8630ca81..224808350 100644 --- a/core/http/endpoints/openai/realtime_classifier_test.go +++ b/core/http/endpoints/openai/realtime_classifier_test.go @@ -45,7 +45,7 @@ var classifierTestHistory = schema.Messages{ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { var out []types.ClassifierResultEvent - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ClassifierResultEvent); ok { out = append(out, ev) } @@ -57,7 +57,7 @@ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { // item — what a classifier response actually "spoke". func replyTexts(t *fakeTransport) []string { var out []string - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ResponseOutputTextDoneEvent); ok { out = append(out, ev.Text) } @@ -277,7 +277,7 @@ var _ = Describe("classifierRespond", func() { Expect(t.countEvents(types.ServerEventTypeResponseOutputTextDone)).To(Equal(1)) Expect(t.countEvents(types.ServerEventTypeResponseFunctionCallArgumentsDone)).To(Equal(1)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -656,7 +656,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(results[0].Arguments).To(MatchJSON(`{"direction":"up","distance":3,"units":"meters"}`)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -695,7 +695,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(handled).To(BeTrue()) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } diff --git a/core/http/endpoints/openai/realtime_doubles_test.go b/core/http/endpoints/openai/realtime_doubles_test.go index a2c104b3c..e82a4a130 100644 --- a/core/http/endpoints/openai/realtime_doubles_test.go +++ b/core/http/endpoints/openai/realtime_doubles_test.go @@ -17,8 +17,10 @@ import ( // so streaming behaviour can be asserted without a real WebSocket/WebRTC peer. // It is not a *WebRTCTransport, so handler code takes the WebSocket path. type fakeTransport struct { - events []types.ServerEvent - audio []fakeAudioChunk + // mu guards sent: some specs send from a background goroutine. + mu sync.Mutex + sent []types.ServerEvent + audio []fakeAudioChunk } type fakeAudioChunk struct { @@ -27,10 +29,19 @@ type fakeAudioChunk struct { } func (f *fakeTransport) SendEvent(e types.ServerEvent) error { - f.events = append(f.events, e) + f.mu.Lock() + defer f.mu.Unlock() + f.sent = append(f.sent, e) return nil } +// events returns a copy of the server events sent so far. +func (f *fakeTransport) events() []types.ServerEvent { + f.mu.Lock() + defer f.mu.Unlock() + return append([]types.ServerEvent(nil), f.sent...) +} + func (f *fakeTransport) ReadEvent() ([]byte, error) { return nil, nil } func (f *fakeTransport) SendAudio(_ context.Context, pcm []byte, sampleRate int) error { @@ -43,7 +54,7 @@ func (f *fakeTransport) Close() error { return nil } // countEvents returns how many recorded events have the given type. func (f *fakeTransport) countEvents(et types.ServerEventType) int { n := 0 - for _, e := range f.events { + for _, e := range f.events() { if e.ServerEventType() == et { n++ } @@ -55,7 +66,7 @@ func (f *fakeTransport) countEvents(et types.ServerEventType) int { // delta event — i.e. the text streamed to the client as it is generated. func (f *fakeTransport) transcriptDeltaText() string { var b strings.Builder - for _, e := range f.events { + for _, e := range f.events() { if d, ok := e.(types.ResponseOutputAudioTranscriptDeltaEvent); ok { b.WriteString(d.Delta) } diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go new file mode 100644 index 000000000..d58d52f33 --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover.go @@ -0,0 +1,70 @@ +package openai + +import ( + "context" + "sort" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" +) + +// isChainStage reports whether stage names a failover chain. +func (m *wrappedModel) isChainStage(stage string) bool { + _, ok := m.stageChains[stage] + return ok && m.failover != nil +} + +// 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 (m *wrappedModel) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { + if !m.isChainStage(stage) { + return fn(base, func() {}) + } + chain := m.stageChains[stage] + return m.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { + cfg, err := m.stageTargetConfig(target) + if err != nil { + return err + } + err = fn(cfg, commit) + if err != nil { + failover.RecordAttemptTrace(m.appTracing, chain, target, err) + } + return err + }) +} + +// startFailoverEvents tells the client which target serves each chain stage +// now, and again whenever a chain switches. The returned func stops it. +func startFailoverEvents(t Transport, fm *failover.Manager, stageChains map[string]string) func() { + // Subscribe before reading the status, so a switch that lands in + // between is still delivered. + events, cancel := fm.Subscribe(16) + stages := make([]string, 0, len(stageChains)) + for s := range stageChains { + stages = append(stages, s) + } + sort.Strings(stages) + for _, stage := range stages { + chain := stageChains[stage] + if st, ok := fm.ChainStatus(chain); ok { + sendEvent(t, types.ModelFailoverEvent{Chain: chain, Stage: stage, To: st.Active, State: string(st.State), Reason: string(failover.ReasonInitial)}) + } + } + go func() { + for ev := range events { + if ev.Type != failover.EventChainSwitched { + continue + } + for _, stage := range stages { + if stageChains[stage] == ev.Chain { + sendEvent(t, types.ModelFailoverEvent{Chain: ev.Chain, Stage: stage, From: ev.From, To: ev.To, State: ev.State, Reason: string(ev.Reason)}) + } + } + } + }() + return cancel +} diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go new file mode 100644 index 000000000..48737630c --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -0,0 +1,99 @@ +package openai + +import ( + "context" + "errors" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" + . "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{failover: fm, stageChains: map[string]string{"tts": "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(), "tts", 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(), "tts", 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(), "tts", 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{"llm": "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() + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index b525eee26..a90a0e4e5 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -19,6 +19,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/http/middleware" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" "github.com/mudler/LocalAI/core/templates" @@ -84,6 +85,18 @@ type wrappedModel struct { routerStore router.DecisionStore routerSessionID string routerUserID string + + // failover and stageChains route pipeline stages that name a failover + // chain; stageChains maps a stage ("llm", "tts", ...) to its chain. + // The *Config fields above 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) + // tuneLLM applies the pipeline's LLM overrides (reasoning effort, + // disable_thinking) to a chain target loaded per call. + tuneLLM func(cfg *config.ModelConfig) + appTracing bool } // anyToAnyModel represent a model which supports Any-to-Any operations @@ -165,15 +178,33 @@ func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { } 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, "vad", 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, "transcription", 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, "sound_detection", 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 +212,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 +235,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("llm") { + return infer(turnCfg, tokenCallback) + } + + return func() (backend.LLMResponse, error) { + var resp backend.LLMResponse + err := m.stageCall(ctx, "llm", 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 +383,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 +434,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 +481,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, "tts", 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 +498,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, "tts", 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 +533,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, "transcription", 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, "transcription", 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 { @@ -693,8 +812,32 @@ func (m *wrappedModel) Warmup(ctx context.Context) error { 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 + // 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. + var ( + plain []backend.PreloadStage + wg sync.WaitGroup + mu sync.Mutex + errs []error + ) + for _, s := range stages { + if !m.isChainStage(s.Role) { + plain = append(plain, s) + continue + } + wg.Go(func() { + err := m.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error { + _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}}) + return err + }) + mu.Lock() + errs = append(errs, err) + mu.Unlock() + }) + } + _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, plain) + wg.Wait() + return errors.Join(append(errs, err)...) } // wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream @@ -854,6 +997,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 +1020,7 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * Store: a.RouterDecisions(), SessionID: sessionID, UserID: userID, + Failover: a.FailoverManager(), } } @@ -882,7 +1028,29 @@ 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. + stageChains := map[string]string{} + resolveStage := func(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { + if cfg == nil || !cfg.IsFailover() { + return cfg, nil + } + if routing == nil || routing.Failover == nil { + return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) + } + st, ok := routing.Failover.ChainStatus(cfg.Name) + if !ok { + return nil, fmt.Errorf("failover chain %q not found", cfg.Name) + } + stageChains[stage] = cfg.Name + return cl.LoadResolvedModelConfig(st.Active, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + } + cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgVAD, err = resolveStage("vad", cfgVAD) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -894,6 +1062,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("transcription", cfgSST) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -926,6 +1097,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("llm", cfgLLM) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -937,10 +1111,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("tts", cfgTTS) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -951,6 +1132,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model } cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = resolveStage("sound_detection", cfgSound) + } if err != nil { return nil, err } @@ -1000,12 +1184,20 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model modelLoader: ml, appConfig: appConfig, evaluator: evaluator, + + stageChains: stageChains, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { + return cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + }, + tuneLLM: tuneLLM, + appTracing: appConfig.EnableTracing, } if routing != nil { wm.routerDeps = routing.Deps wm.routerStore = routing.Store wm.routerSessionID = routing.SessionID wm.routerUserID = routing.UserID + wm.failover = routing.Failover } return wm, nil } diff --git a/core/http/endpoints/openai/realtime_semantic_vad_test.go b/core/http/endpoints/openai/realtime_semantic_vad_test.go index c36e13563..b1107c1f6 100644 --- a/core/http/endpoints/openai/realtime_semantic_vad_test.go +++ b/core/http/endpoints/openai/realtime_semantic_vad_test.go @@ -269,7 +269,7 @@ var _ = Describe("liveTurnState", func() { lts.drainEvents(1.0) var got []types.ConversationItemInputAudioTranscriptionDeltaEvent - for _, e := range ftr.events { + for _, e := range ftr.events() { if d, ok := e.(types.ConversationItemInputAudioTranscriptionDeltaEvent); ok { got = append(got, d) } @@ -335,7 +335,7 @@ var _ = Describe("commitUtteranceWithTranscript", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) var completed types.ConversationItemInputAudioTranscriptionCompletedEvent - for _, e := range tr.events { + for _, e := range tr.events() { if c, ok := e.(types.ConversationItemInputAudioTranscriptionCompletedEvent); ok { completed = c } @@ -407,7 +407,7 @@ var _ = Describe("emitPrecomputedTranscription", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionDelta)).To(Equal(2), "empty deltas skipped") Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) - for _, e := range tr.events { + for _, e := range tr.events() { switch ev := e.(type) { case types.ConversationItemInputAudioTranscriptionDeltaEvent: Expect(ev.ItemID).To(Equal("item42")) diff --git a/core/http/endpoints/openai/realtime_sound_detection_test.go b/core/http/endpoints/openai/realtime_sound_detection_test.go index e440e80c3..058c74076 100644 --- a/core/http/endpoints/openai/realtime_sound_detection_test.go +++ b/core/http/endpoints/openai/realtime_sound_detection_test.go @@ -38,7 +38,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.ItemID).To(Equal("item1")) Expect(ev.ContentIndex).To(Equal(0)) @@ -62,7 +62,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.Detections).To(BeEmpty()) }) diff --git a/core/http/endpoints/openai/realtime_stream_test.go b/core/http/endpoints/openai/realtime_stream_test.go index 2d5d7d7a1..5f160027e 100644 --- a/core/http/endpoints/openai/realtime_stream_test.go +++ b/core/http/endpoints/openai/realtime_stream_test.go @@ -250,8 +250,8 @@ var _ = Describe("triggerResponse", func() { // The single terminal carries the produced output item and the usage — // both empty in the legacy code. var done *types.ResponseDoneEvent - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { done = &d } } @@ -287,8 +287,8 @@ var _ = Describe("triggerResponse", func() { var created *types.ResponseCreatedEvent var done *types.ResponseDoneEvent - for i := range t.events { - switch e := t.events[i].(type) { + for i := range t.events() { + switch e := t.events()[i].(type) { case types.ResponseCreatedEvent: created = &e case types.ResponseDoneEvent: @@ -317,8 +317,8 @@ var _ = Describe("triggerResponse", func() { triggerResponse(context.Background(), session, &Conversation{}, t, nil) - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { Expect(d.Response.Metadata).To(BeEmpty()) } } diff --git a/core/http/endpoints/openai/realtime_voicegate_integration_test.go b/core/http/endpoints/openai/realtime_voicegate_integration_test.go index b0f7f0b49..4da774c76 100644 --- a/core/http/endpoints/openai/realtime_voicegate_integration_test.go +++ b/core/http/endpoints/openai/realtime_voicegate_integration_test.go @@ -67,7 +67,7 @@ func itSession(gate *voiceGate) (*Session, *fakeModel) { // hasSpeakerNotAuthorized reports whether a speaker_not_authorized error event // was emitted to the client. func hasSpeakerNotAuthorized(tr *fakeTransport) bool { - for _, e := range tr.events { + for _, e := range tr.events() { if ev, ok := e.(types.ErrorEvent); ok && ev.Error.Code == "speaker_not_authorized" { return true } diff --git a/core/http/endpoints/openai/types/failover.go b/core/http/endpoints/openai/types/failover.go new file mode 100644 index 000000000..69c048207 --- /dev/null +++ b/core/http/endpoints/openai/types/failover.go @@ -0,0 +1,46 @@ +package types + +import "encoding/json" + +// ModelFailoverEvent is a LocalAI extension server event +// (localai.model.failover). It tells a client which target serves a +// pipeline stage that names a failover chain: once per chain stage at +// session start (reason "initial"), then on every switch of that chain. +type ModelFailoverEvent struct { + ServerEventBase + + // The failover chain the stage names. + Chain string `json:"chain"` + + // The pipeline stage: vad, transcription, llm, tts or sound_detection. + Stage string `json:"stage"` + + // The target that served the stage before the switch; "" at session start. + From string `json:"from"` + + // The target that serves the stage now. + To string `json:"to"` + + // The chain state: primary, fallback or degraded. + State string `json:"state"` + + // Why the chain switched, or "initial" at session start. + Reason string `json:"reason"` +} + +func (m ModelFailoverEvent) ServerEventType() ServerEventType { + return ServerEventTypeModelFailover +} + +func (m ModelFailoverEvent) MarshalJSON() ([]byte, error) { + type typeAlias ModelFailoverEvent + type typeWrapper struct { + typeAlias + Type ServerEventType `json:"type"` + } + shadow := typeWrapper{ + typeAlias: typeAlias(m), + Type: m.ServerEventType(), + } + return json.Marshal(shadow) +} diff --git a/core/http/endpoints/openai/types/server_events.go b/core/http/endpoints/openai/types/server_events.go index b847a35a7..114a7065a 100644 --- a/core/http/endpoints/openai/types/server_events.go +++ b/core/http/endpoints/openai/types/server_events.go @@ -27,7 +27,11 @@ const ( // ServerEventTypeClassifierResult is a LocalAI extension: it carries the // classifier-mode score distribution and decision for a response. OpenAI // clients ignore it. - ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + // ServerEventTypeModelFailover is a LocalAI extension: it names the target + // that serves a pipeline stage backed by a failover chain, at session start + // and on every chain switch. OpenAI clients ignore it. + ServerEventTypeModelFailover ServerEventType = "localai.model.failover" ServerEventTypeInputAudioBufferCommitted ServerEventType = "input_audio_buffer.committed" ServerEventTypeInputAudioBufferCleared ServerEventType = "input_audio_buffer.cleared" ServerEventTypeInputAudioBufferSpeechStarted ServerEventType = "input_audio_buffer.speech_started" diff --git a/tests/e2e/e2e_suite_test.go b/tests/e2e/e2e_suite_test.go index a89318dfc..e7bb844d5 100644 --- a/tests/e2e/e2e_suite_test.go +++ b/tests/e2e/e2e_suite_test.go @@ -273,6 +273,44 @@ var _ = BeforeSuite(func() { Expect(err).ToNot(HaveOccurred()) Expect(os.WriteFile(filepath.Join(modelsPath, "realtime-pipeline.yaml"), pipelineData, 0644)).To(Succeed()) + // Realtime pipelines whose LLM is a failover chain: target 0 always fails + // to load, target 1 is mock-llm. rt-failover skips the warm-up so the + // switch happens on the first turn, mid-session; rt-failover-warm keeps it + // so the switch happens while the session starts. Each has its own chain + // because chain state is shared across sessions. + for _, rt := range []struct { + name, suffix string + disableWarmup bool + }{{"rt-failover", "rt", true}, {"rt-failover-warm", "rt-warm", false}} { + for _, cfg := range []map[string]any{ + { + "name": "fail-" + rt.suffix, + "backend": "mock-backend", + "parameters": map[string]any{"model": "fail-load-" + rt.suffix}, + }, + { + "name": "chain-" + rt.suffix, + "failover": map[string]any{ + "targets": []map[string]any{{"model": "fail-" + rt.suffix}, {"model": "mock-llm"}}, + }, + }, + { + "name": rt.name, + "pipeline": map[string]any{ + "vad": "mock-vad", + "transcription": "mock-stt", + "llm": "chain-" + rt.suffix, + "tts": "mock-tts", + "disable_warmup": rt.disableWarmup, + }, + }, + } { + data, err := yaml.Marshal(cfg) + Expect(err).ToNot(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) + } + } + // Classifier-mode pipeline (LocalAI extension): responses are // prefill-scored against the option list via the mock backend's // ROUTE_HINT-driven Score instead of being generated. Threshold 0.6: diff --git a/tests/e2e/realtime_ws_test.go b/tests/e2e/realtime_ws_test.go index 6aeb3fcd4..1aa2552a1 100644 --- a/tests/e2e/realtime_ws_test.go +++ b/tests/e2e/realtime_ws_test.go @@ -192,6 +192,106 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { }) }) + Context("Failover chain stage", Label("failover"), func() { + // userTurn adds a user text item, asks for a response and reads until + // response.done. It returns the user item id, the response.done event + // and the localai.model.failover events seen on the way. + userTurn := func(conn *websocket.Conn, text string) (string, map[string]any, []map[string]any) { + sendClientEvent(conn, map[string]any{ + "type": "conversation.item.create", + "item": map[string]any{ + "type": "message", + "role": "user", + "content": []map[string]any{{"type": "input_text", "text": text}}, + }, + }) + added := drainUntil(conn, "conversation.item.added", 10*time.Second) + item, _ := added["item"].(map[string]any) + userID, _ := item["id"].(string) + ExpectWithOffset(1, userID).ToNot(BeEmpty()) + + sendClientEvent(conn, map[string]any{"type": "response.create"}) + var failovers []map[string]any + deadline := time.Now().Add(60 * time.Second) + for time.Now().Before(deadline) { + evt := readServerEvent(conn, time.Until(deadline)) + switch evt["type"] { + case "localai.model.failover": + failovers = append(failovers, evt) + case "error": + Fail(fmt.Sprintf("unexpected error event: %v", evt)) + case "response.done": + return userID, evt, failovers + } + } + Fail("timed out waiting for response.done") + return "", nil, nil + } + + retrieveItem := func(conn *websocket.Conn, id string) map[string]any { + sendClientEvent(conn, map[string]any{"type": "conversation.item.retrieve", "item_id": id}) + evt := drainUntil(conn, "conversation.item.retrieved", 10*time.Second) + item, _ := evt["item"].(map[string]any) + return item + } + + It("switches the LLM mid-session and keeps the conversation", func() { + conn := connectWS("rt-failover") + defer conn.Close() + + Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) + initial := drainUntil(conn, "localai.model.failover", 10*time.Second) + Expect(initial).To(HaveKeyWithValue("stage", "llm")) + Expect(initial).To(HaveKeyWithValue("chain", "chain-rt")) + Expect(initial).To(HaveKeyWithValue("reason", "initial")) + Expect(initial).To(HaveKeyWithValue("to", "fail-rt")) + + sendClientEvent(conn, disableVADEvent()) + drainUntil(conn, "session.updated", 10*time.Second) + + firstID, done, failovers := userTurn(conn, "Hello, how are you?") + Expect(failovers).To(ContainElement(And( + HaveKeyWithValue("stage", "llm"), + HaveKeyWithValue("from", "fail-rt"), + HaveKeyWithValue("to", "mock-llm"), + HaveKeyWithValue("reason", "trip"), + ))) + resp, _ := done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + output, _ := resp["output"].([]any) + Expect(output).ToNot(BeEmpty()) + firstReply, _ := output[0].(map[string]any) + firstReplyID, _ := firstReply["id"].(string) + Expect(firstReplyID).ToNot(BeEmpty()) + + _, done, _ = userTurn(conn, "And now?") + resp, _ = done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + + // The switch kept the session: the first turn is still in it. + Expect(retrieveItem(conn, firstID)).To(HaveKeyWithValue("id", firstID)) + Expect(retrieveItem(conn, firstReplyID)).To(HaveKeyWithValue("id", firstReplyID)) + }) + + It("starts the session on the next target when the active one fails to warm up", func() { + conn := connectWS("rt-failover-warm") + defer conn.Close() + + Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) + initial := drainUntil(conn, "localai.model.failover", 10*time.Second) + Expect(initial).To(HaveKeyWithValue("chain", "chain-rt-warm")) + Expect(initial).To(HaveKeyWithValue("reason", "initial")) + Expect(initial).To(HaveKeyWithValue("to", "mock-llm")) + + sendClientEvent(conn, disableVADEvent()) + drainUntil(conn, "session.updated", 10*time.Second) + + _, done, _ := userTurn(conn, "Hello?") + resp, _ := done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + }) + }) + Context("Manual audio commit", func() { It("should produce a response with audio when audio is committed", func() { conn := connectWS(pipelineModel()) From 191d2d63017dce14d044bbb4a23771ecd7113e44 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 18:18:46 +0000 Subject: [PATCH 21/79] feat(failover): add MCP tools to list chains and pin targets Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .../task-12-report.md | 137 ++++++++++++++++++ core/application/application.go | 13 ++ core/application/startup.go | 6 + .../endpoints/mcp/localai_assistant_test.go | 8 + pkg/mcp/localaitools/client.go | 11 ++ pkg/mcp/localaitools/coverage_test.go | 41 +++--- pkg/mcp/localaitools/dto.go | 23 +++ pkg/mcp/localaitools/fakes_test.go | 27 ++++ pkg/mcp/localaitools/httpapi/client.go | 20 +++ pkg/mcp/localaitools/httpapi/client_test.go | 62 ++++++++ pkg/mcp/localaitools/httpapi/routes.go | 1 + pkg/mcp/localaitools/inproc/client.go | 43 ++++++ pkg/mcp/localaitools/inproc/client_test.go | 90 ++++++++++++ pkg/mcp/localaitools/prompts/10_safety.md | 2 +- pkg/mcp/localaitools/prompts/20_tools.md | 3 + pkg/mcp/localaitools/server.go | 1 + pkg/mcp/localaitools/server_test.go | 4 + pkg/mcp/localaitools/tools.go | 11 ++ pkg/mcp/localaitools/tools_failover.go | 53 +++++++ 19 files changed, 536 insertions(+), 20 deletions(-) create mode 100644 .superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md create mode 100644 pkg/mcp/localaitools/tools_failover.go diff --git a/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md b/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md new file mode 100644 index 000000000..ae2d8cb98 --- /dev/null +++ b/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md @@ -0,0 +1,137 @@ +# Task 12 report: MCP admin tools for failover chains + +## What I implemented + +Three MCP admin tools mirroring the Task 9 REST endpoints: + +- `list_failover_chains` (read-only) → `GET /api/failover` +- `pin_failover_target` (mutating) → `POST /api/failover/:chain/pin` +- `unpin_failover_target` (mutating) → `DELETE /api/failover/:chain/pin` + +Files changed, by layer: + +- **Tool identity**: `pkg/mcp/localaitools/tools.go` — added `ToolListFailoverChains` (read-only block), `ToolPinFailoverTarget`/`ToolUnpinFailoverTarget` (mutating block + `mutatingToolNames`). +- **DTOs**: `pkg/mcp/localaitools/dto.go` — added `FailoverTargetInfo` and `FailoverChainInfo` (LLM-facing subset of `failover.TargetStatus`/`failover.ChainStatus`). +- **Interface**: `pkg/mcp/localaitools/client.go` — added `ListFailoverChains`, `PinFailoverTarget`, `UnpinFailoverTarget` to `LocalAIClient`. +- **In-process impl**: `pkg/mcp/localaitools/inproc/client.go` — added `Failover *failover.Manager` field and the three methods (nil-safe: `ListFailoverChains` returns `[]`, Pin/Unpin return `"failover is not running"`). +- **HTTP impl**: `pkg/mcp/localaitools/httpapi/client.go` + `httpapi/routes.go` — added `routeFailover = "/api/failover"` and the three methods, using `c.do`. +- **Tool registration**: new file `pkg/mcp/localaitools/tools_failover.go` (`registerFailoverTools`), wired into `pkg/mcp/localaitools/server.go`. +- **Prompts**: `prompts/20_tools.md` (read-only + mutating one-liners) and `prompts/10_safety.md` (both mutating names added to the confirmation-rule list). +- **Wiring the manager into the in-process client**: `core/application/application.go` and `core/application/startup.go` — see "Adaptation: startup ordering" below; this was the one place I had to deviate from a literal reading of the brief. +- **Tests**: `coverage_test.go`, `server_test.go`, `fakes_test.go`, `inproc/client_test.go`, `httpapi/client_test.go` — all as specified, plus `core/http/endpoints/mcp/localai_assistant_test.go` (see adaptations). + +`parity_test.go` was left untouched — the brief didn't specify a new parity spec for failover, and the file's existing specs are all hand-picked equality checks for specific methods (ListGalleries, GallerySearch, ImportModelURI, SystemInfo); it isn't a generic "every method" loop, so nothing there needed updating for the build to stay green. + +## Adaptations from the brief (read the real code, deviated where it disagreed) + +1. **`jsonResult`/`errorResult` return one value, not three.** The brief's `tools_failover.go` snippet writes `return jsonResult(chains)` as if it were the handler's whole 3-tuple return. The real helpers in `pkg/mcp/localaitools/errors.go` are: + ```go + func errorResult(err error) *mcp.CallToolResult + func jsonResult(v any) *mcp.CallToolResult + ``` + Matching `tools_aliases.go`'s actual pattern, every handler returns `jsonResult(x), nil, nil` / `errorResult(err), nil, nil` — three explicit values. I used that form throughout `tools_failover.go`. + +2. **Mutating tool descriptions reference safety rule 1.** `.agents/localai-assistant-mcp.md`'s checklist says "Mutating tools must reference safety rule 1 in the description," and `tools_aliases.go`'s `set_alias` does this ("Requires user confirmation per safety rule 1."). The brief's descriptions for `pin_failover_target`/`unpin_failover_target` didn't include this phrase, so I added it to match the established convention and the checklist. + +3. **Startup ordering: `assistantClient.Failover` can't be set where the brief implies.** The brief says to set the field "where the client is constructed (from `application.FailoverManager()`)." I found that construction site (`core/application/application.go`'s `start()`, called from `core/application/startup.go`'s `New()` at line 74) — but `application.failoverManager` is only built later in the same `New()` function, at line 259, **after** `start()` (and therefore the assistant-client construction) has already returned. Calling `a.FailoverManager()` inside `start()` would have captured a permanent `nil`. + + Fix: added an `assistantClient *localaiInproc.Client` field to `Application` (set in `start()` when the assistant client is built), then in `startup.go`, right after `application.failoverManager = failover.New(...)`, added: + ```go + if application.assistantClient != nil { + application.assistantClient.Failover = application.failoverManager + } + ``` + This is safe because `assistantClient` is a pointer already captured by value inside the `LocalAIClient` interface passed to `holder.Initialize()` — mutating a field on it after the fact is visible through the interface. Verified with `go test ./core/application/...` and `go test ./core/http/endpoints/mcp/...`. + +4. **`stubClient` in `core/http/endpoints/mcp/localai_assistant_test.go`** implements `localaitools.LocalAIClient` for that package's own tests and isn't in the brief's file list, but the interface change broke its build. Added the three stub methods (empty list, nil errors) to keep it compiling — same pattern as its existing stubs for `GetRouterCorpusStats` etc. + +5. **inproc failover test fixture**: the brief says to "build a `failover.Manager` over a small in-memory source with one chain." `failover.ConfigSource` (`GetModelConfig`/`GetAllModelsConfigs`) is exported, but the concrete fake used by `core/services/failover`'s own tests (`fakeSource`, `chainCfg`, `t()`) is unexported and package-local, so I wrote a minimal `fakeFailoverSource` directly in `inproc/client_test.go` implementing the same two-method interface over a `map[string]config.ModelConfig`, seeded with a `chain` config (`config.FailoverConfig{Targets: [...]}`) plus two local targets `a`/`b`. + +## TDD evidence + +**RED** — after writing all test-file changes (`coverage_test.go`, `server_test.go`, `fakes_test.go`, plus new specs in `inproc/client_test.go` and `httpapi/client_test.go` were added later, see below), I temporarily reverted the implementation files (`git stash push` on `tools.go`, `dto.go`, `client.go`, `server.go`, `inproc/client.go`, `httpapi/client.go`, `httpapi/routes.go`, prompts, and the two `core/application` files; moved `tools_failover.go` out of the package) and ran: + +``` +$ go test ./pkg/mcp/localaitools/... 2>&1 | tail -30 +# github.com/mudler/LocalAI/pkg/mcp/localaitools [github.com/mudler/LocalAI/pkg/mcp/localaitools.test] +pkg/mcp/localaitools/fakes_test.go:62:32: undefined: FailoverChainInfo +pkg/mcp/localaitools/fakes_test.go:400:63: undefined: FailoverChainInfo +pkg/mcp/localaitools/fakes_test.go:405:11: undefined: FailoverChainInfo +pkg/mcp/localaitools/coverage_test.go:49:2: undefined: ToolListFailoverChains +pkg/mcp/localaitools/coverage_test.go:71:2: undefined: ToolPinFailoverTarget +pkg/mcp/localaitools/coverage_test.go:72:2: undefined: ToolUnpinFailoverTarget +pkg/mcp/localaitools/server_test.go:95:2: undefined: ToolListFailoverChains +pkg/mcp/localaitools/server_test.go:159:4: undefined: ToolListFailoverChains +pkg/mcp/localaitools/server_test.go:160:4: undefined: ToolPinFailoverTarget +pkg/mcp/localaitools/server_test.go:161:4: undefined: ToolUnpinFailoverTarget +pkg/mcp/localaitools/server_test.go:161:4: too many errors +FAIL github.com/mudler/LocalAI/pkg/mcp/localaitools [build failed] +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.035s +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.167s +``` +This matches the brief's Step 1 expectation ("Expected: compile failure"). I then `git stash pop` and restored `tools_failover.go` to get back to the implemented state. + +**GREEN** — after restoring the implementation and adding the remaining `inproc/client_test.go` / `httpapi/client_test.go` specs: + +``` +$ go test ./pkg/mcp/localaitools/... -count=1 -v 2>&1 | grep -E "SUCCESS|FAIL|ok " +SUCCESS! -- 53 Passed | 0 Failed | 0 Pending | 0 Skipped +ok github.com/mudler/LocalAI/pkg/mcp/localaitools 0.133s +SUCCESS! -- 24 Passed | 0 Failed | 0 Pending | 0 Skipped +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.037s +SUCCESS! -- 16 Passed | 0 Failed | 0 Pending | 0 Skipped +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.171s +``` + +Additional verification (build scope per implementer-rules, plus the two packages touched indirectly by the interface change): + +``` +$ go build ./core/... ./pkg/mcp/... ./tests/... +(clean, no output) + +$ go test ./pkg/mcp/localaitools/... ./core/application/... ./core/http/endpoints/mcp/... -count=1 +ok github.com/mudler/LocalAI/pkg/mcp/localaitools 0.155s +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.045s +ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.176s +ok github.com/mudler/LocalAI/core/application 0.265s +ok github.com/mudler/LocalAI/core/http/endpoints/mcp 0.237s + +$ gofmt -l +(empty — clean) + +$ go vet ./core/application/... ./pkg/mcp/... +(clean, no output) +``` + +## Files changed + +- `pkg/mcp/localaitools/tools.go` +- `pkg/mcp/localaitools/dto.go` +- `pkg/mcp/localaitools/client.go` +- `pkg/mcp/localaitools/server.go` +- `pkg/mcp/localaitools/tools_failover.go` (new) +- `pkg/mcp/localaitools/inproc/client.go` +- `pkg/mcp/localaitools/httpapi/client.go` +- `pkg/mcp/localaitools/httpapi/routes.go` +- `pkg/mcp/localaitools/prompts/20_tools.md` +- `pkg/mcp/localaitools/prompts/10_safety.md` +- `pkg/mcp/localaitools/coverage_test.go` +- `pkg/mcp/localaitools/server_test.go` +- `pkg/mcp/localaitools/fakes_test.go` +- `pkg/mcp/localaitools/inproc/client_test.go` +- `pkg/mcp/localaitools/httpapi/client_test.go` +- `core/application/application.go` +- `core/application/startup.go` +- `core/http/endpoints/mcp/localai_assistant_test.go` + +## Self-review + +- Completeness: all three tools registered, gated correctly (`ToolListFailoverChains` in the read-only catalog; both mutating tools skipped when `Options.DisableMutating`), both client implementations covered, prompts updated, safety-rule coverage test (`TestPromptsContainSafetyAnchors`'s "names every mutating tool" spec) passes automatically since it reads `mutatingToolNames`. +- Quality/YAGNI: DTOs intentionally drop `ConsecutiveOK`/`LastProbe`/`ActiveSince` — internal probe bookkeeping the LLM has no use for when deciding to pin/unpin; documented why in the doc comment. +- Nil-safety: both inproc failover methods and the httpapi Pin/Unpin exercise the "no failover configured" path in tests (inproc has explicit specs for it; httpapi's behavior when unconfigured is identical to any other client error — REST returns 404, `c.do` surfaces `*HTTPError`, no new code path needed there). +- Existing patterns followed: constant grouping/comments, `errorResult`/`jsonResult` triple-return, `c.do` signature, `url.PathEscape` on the chain path segment, fake-client recording pattern, Ginkgo `Describe`/`It` structure matching the alias specs. +- Pristine output: `gofmt -l` and `go vet` clean across all touched files. + +## Concerns + +- None blocking. The startup-ordering fix (adaptation 3) is the only piece that goes beyond a single-file, mechanical change — it touches two `core/application` files instead of the "one field set inline" the brief describes. I verified it with both `core/application` and `core/http/endpoints/mcp` package tests, and confirmed via read of `startup.go` that `New()` is the sole caller of `start()` and that `failoverManager` is not read anywhere between `start()` and its own assignment, so there's no other place relying on it being nil momentarily. diff --git a/core/application/application.go b/core/application/application.go index ac66e75b6..f5454aba1 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -101,6 +101,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. // @@ -598,6 +605,12 @@ func (a *Application) start() error { assistantClient.RouterEmbedder = a.Embedder assistantClient.RouterEmbedderFingerprint = a.EmbedderFingerprint assistantClient.RouterVectorStore = a.VectorStore + // Failover chains: failoverManager does not exist yet at this point + // in startup (New() in startup.go builds it after start() returns), + // so it can't be wired here like the fields above. New() sets + // assistantClient.Failover directly once the manager is built; + // stash the client so it can reach back into it. + a.assistantClient = assistantClient if err := holder.Initialize(a.applicationConfig.Context, assistantClient, localaitools.Options{}); err != nil { // Why log+continue instead of fail: the assistant is an optional // feature; a failure here must not take down the whole server. diff --git a/core/application/startup.go b/core/application/startup.go index a700615e2..938f57496 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -262,6 +262,12 @@ func New(opts ...config.AppOption) (*Application, error) { })), 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 diff --git a/core/http/endpoints/mcp/localai_assistant_test.go b/core/http/endpoints/mcp/localai_assistant_test.go index 71ce38a7a..8629bd9e1 100644 --- a/core/http/endpoints/mcp/localai_assistant_test.go +++ b/core/http/endpoints/mcp/localai_assistant_test.go @@ -214,3 +214,11 @@ func (stubClient) SeedRouterCorpus(_ context.Context, req localaitools.RouterCor func (stubClient) ClearRouterCorpus(_ context.Context, routerModel string) (*localaitools.RouterCorpusClearResult, error) { return &localaitools.RouterCorpusClearResult{Router: routerModel}, nil } + +func (stubClient) ListFailoverChains(_ context.Context) ([]localaitools.FailoverChainInfo, error) { + return []localaitools.FailoverChainInfo{}, nil +} + +func (stubClient) PinFailoverTarget(_ context.Context, _, _ string) error { return nil } + +func (stubClient) UnpinFailoverTarget(_ context.Context, _ string) error { return nil } diff --git a/pkg/mcp/localaitools/client.go b/pkg/mcp/localaitools/client.go index bf683a508..7cba8a8de 100644 --- a/pkg/mcp/localaitools/client.go +++ b/pkg/mcp/localaitools/client.go @@ -130,4 +130,15 @@ type LocalAIClient interface { // ClearRouterCorpus wipes a knn router's corpus — file and live // index. ClearRouterCorpus(ctx context.Context, routerModel string) (*RouterCorpusClearResult, error) + + // ---- Failover chains ---- + // ListFailoverChains reports every configured failover chain, its + // currently active target, and the health of each target. + ListFailoverChains(ctx context.Context) ([]FailoverChainInfo, error) + // PinFailoverTarget forces chain to serve every request from target, + // regardless of health, until unpinned. + PinFailoverTarget(ctx context.Context, chain, target string) error + // UnpinFailoverTarget removes chain's pin so health decides the + // active target again. + UnpinFailoverTarget(ctx context.Context, chain string) error } diff --git a/pkg/mcp/localaitools/coverage_test.go b/pkg/mcp/localaitools/coverage_test.go index a802226f0..75831769e 100644 --- a/pkg/mcp/localaitools/coverage_test.go +++ b/pkg/mcp/localaitools/coverage_test.go @@ -46,27 +46,30 @@ var toolToHTTPRoute = map[string]string{ ToolGetRouterCorpusStats: "GET /api/router/:name/corpus/stats", ToolListAliases: "GET /api/aliases", ToolListVoiceProfiles: "GET /api/voice-profiles", + ToolListFailoverChains: "GET /api/failover", // Mutating tools. - ToolInstallModel: "POST /models/apply", - ToolImportModelURI: "POST /models/import-uri", - ToolDeleteModel: "POST /models/delete/:name", - ToolEditModelConfig: "PATCH /api/models/config-json/:name", - ToolReloadModels: "POST /models/reload", - ToolLoadModel: "POST /backend/load", - ToolInstallBackend: "POST /backends/apply", - ToolUpgradeBackend: "POST /backends/upgrade/:name", - ToolToggleModelState: "PUT /models/toggle-state/:name/:action", - ToolToggleModelPinned: "PUT /models/toggle-pinned/:name/:action", - ToolSetBranding: "POST /api/settings (instance_name, instance_tagline)", - ToolSetAlias: "PATCH /api/models/config-json/:name (swap) or POST /models/import (create)", - ToolSeedRouterCorpus: "POST /api/router/:name/corpus", - ToolClearRouterCorpus: "DELETE /api/router/:name/corpus", - ToolCreateVoiceProfile: "POST /api/voice-profiles", - ToolDeleteVoiceProfile: "DELETE /api/voice-profiles/:id", - ToolSetNodeVRAMBudget: "PUT /api/nodes/:id/vram-budget", - ToolSetScheduling: "POST /api/nodes/scheduling", - ToolDeleteScheduling: "DELETE /api/nodes/scheduling/:model", + ToolInstallModel: "POST /models/apply", + ToolImportModelURI: "POST /models/import-uri", + ToolDeleteModel: "POST /models/delete/:name", + ToolEditModelConfig: "PATCH /api/models/config-json/:name", + ToolReloadModels: "POST /models/reload", + ToolLoadModel: "POST /backend/load", + ToolInstallBackend: "POST /backends/apply", + ToolUpgradeBackend: "POST /backends/upgrade/:name", + ToolToggleModelState: "PUT /models/toggle-state/:name/:action", + ToolToggleModelPinned: "PUT /models/toggle-pinned/:name/:action", + ToolSetBranding: "POST /api/settings (instance_name, instance_tagline)", + ToolSetAlias: "PATCH /api/models/config-json/:name (swap) or POST /models/import (create)", + ToolSeedRouterCorpus: "POST /api/router/:name/corpus", + ToolClearRouterCorpus: "DELETE /api/router/:name/corpus", + ToolCreateVoiceProfile: "POST /api/voice-profiles", + ToolDeleteVoiceProfile: "DELETE /api/voice-profiles/:id", + ToolSetNodeVRAMBudget: "PUT /api/nodes/:id/vram-budget", + ToolSetScheduling: "POST /api/nodes/scheduling", + ToolDeleteScheduling: "DELETE /api/nodes/scheduling/:model", + ToolPinFailoverTarget: "POST /api/failover/:chain/pin", + ToolUnpinFailoverTarget: "DELETE /api/failover/:chain/pin", } // allKnownTools is the union of expectedFullCatalog (defined in diff --git a/pkg/mcp/localaitools/dto.go b/pkg/mcp/localaitools/dto.go index 24da8aa18..afd19bafd 100644 --- a/pkg/mcp/localaitools/dto.go +++ b/pkg/mcp/localaitools/dto.go @@ -413,3 +413,26 @@ type VRAMEstimateRequest struct { GPULayers int `json:"gpu_layers,omitempty" jsonschema:"Number of layers to offload to GPU. -1 for all."` KVQuantBits int `json:"kv_quant_bits,omitempty" jsonschema:"KV cache quantization bits (e.g. 4, 8, 16)."` } + +// FailoverTargetInfo is the LLM-facing view of one failover chain target's +// health. It mirrors failover.TargetStatus but drops ConsecutiveOK and +// LastProbe — internal probing detail the LLM doesn't need to decide +// whether to pin or unpin a target. +type FailoverTargetInfo struct { + Model string `json:"model"` + Kind string `json:"kind"` + Warm bool `json:"warm"` + State string `json:"state"` + LastError string `json:"last_error,omitempty"` +} + +// FailoverChainInfo is the LLM-facing view of one failover chain: its +// current state, the target serving it now, an optional pin, and every +// target's health. +type FailoverChainInfo struct { + Name string `json:"name"` + State string `json:"state"` + Active string `json:"active"` + Pinned string `json:"pinned,omitempty"` + Targets []FailoverTargetInfo `json:"targets"` +} diff --git a/pkg/mcp/localaitools/fakes_test.go b/pkg/mcp/localaitools/fakes_test.go index b0697b272..cba6efa57 100644 --- a/pkg/mcp/localaitools/fakes_test.go +++ b/pkg/mcp/localaitools/fakes_test.go @@ -59,6 +59,9 @@ type fakeClient struct { getPIIEvents func(PIIEventsQuery) ([]PIIEvent, error) getMiddlewareStatus func() (*MiddlewareStatus, error) getRouterDecisions func(RouterDecisionsQuery) ([]RouterDecision, error) + listFailoverChains func() ([]FailoverChainInfo, error) + pinFailoverTarget func(string, string) error + unpinFailoverTarget func(string) error } type fakeCall struct { @@ -393,3 +396,27 @@ func (f *fakeClient) ClearRouterCorpus(_ context.Context, routerModel string) (* f.record("ClearRouterCorpus", routerModel) return &RouterCorpusClearResult{Router: routerModel}, nil } + +func (f *fakeClient) ListFailoverChains(_ context.Context) ([]FailoverChainInfo, error) { + f.record("ListFailoverChains", nil) + if f.listFailoverChains != nil { + return f.listFailoverChains() + } + return []FailoverChainInfo{}, nil +} + +func (f *fakeClient) PinFailoverTarget(_ context.Context, chain, target string) error { + f.record("PinFailoverTarget", []any{chain, target}) + if f.pinFailoverTarget != nil { + return f.pinFailoverTarget(chain, target) + } + return nil +} + +func (f *fakeClient) UnpinFailoverTarget(_ context.Context, chain string) error { + f.record("UnpinFailoverTarget", chain) + if f.unpinFailoverTarget != nil { + return f.unpinFailoverTarget(chain) + } + return nil +} diff --git a/pkg/mcp/localaitools/httpapi/client.go b/pkg/mcp/localaitools/httpapi/client.go index 923f35eba..631977c03 100644 --- a/pkg/mcp/localaitools/httpapi/client.go +++ b/pkg/mcp/localaitools/httpapi/client.go @@ -829,3 +829,23 @@ func (c *Client) ClearRouterCorpus(ctx context.Context, routerModel string) (*lo } return &out, nil } + +// ---- Failover chains ---- + +func (c *Client) ListFailoverChains(ctx context.Context) ([]localaitools.FailoverChainInfo, error) { + var out struct { + Chains []localaitools.FailoverChainInfo `json:"chains"` + } + if err := c.do(ctx, http.MethodGet, routeFailover, nil, &out); err != nil { + return nil, err + } + return out.Chains, nil +} + +func (c *Client) PinFailoverTarget(ctx context.Context, chain, target string) error { + return c.do(ctx, http.MethodPost, routeFailover+"/"+url.PathEscape(chain)+"/pin", map[string]string{"target": target}, nil) +} + +func (c *Client) UnpinFailoverTarget(ctx context.Context, chain string) error { + return c.do(ctx, http.MethodDelete, routeFailover+"/"+url.PathEscape(chain)+"/pin", nil, nil) +} diff --git a/pkg/mcp/localaitools/httpapi/client_test.go b/pkg/mcp/localaitools/httpapi/client_test.go index c4e596650..fd1cff632 100644 --- a/pkg/mcp/localaitools/httpapi/client_test.go +++ b/pkg/mcp/localaitools/httpapi/client_test.go @@ -365,6 +365,68 @@ var _ = Describe("Model aliases", func() { }) }) +var _ = Describe("Failover chains", func() { + Describe("ListFailoverChains", func() { + It("issues GET /api/failover and unwraps the chains array", func() { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + Expect(r.Method).To(Equal(http.MethodGet)) + Expect(r.URL.Path).To(Equal("/api/failover")) + _ = json.NewEncoder(w).Encode(map[string]any{ + "chains": []map[string]any{ + { + "name": "chain", + "state": "primary", + "active": "a", + "pinned": nil, + "targets": []map[string]any{ + {"model": "a", "kind": "local", "warm": false, "state": "healthy"}, + }, + }, + }, + }) + })) + DeferCleanup(srv.Close) + + out, err := New(srv.URL, "").ListFailoverChains(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(out).To(HaveLen(1)) + Expect(out[0].Name).To(Equal("chain")) + Expect(out[0].Active).To(Equal("a")) + Expect(out[0].Pinned).To(BeEmpty()) + Expect(out[0].Targets).To(ConsistOf(localaitools.FailoverTargetInfo{Model: "a", Kind: "local", State: "healthy"})) + }) + }) + + Describe("PinFailoverTarget", func() { + It("issues POST /api/failover/:chain/pin with the target in the body", func() { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + Expect(r.Method).To(Equal(http.MethodPost)) + Expect(r.URL.Path).To(Equal("/api/failover/chain/pin")) + var body map[string]string + Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed()) + Expect(body).To(HaveKeyWithValue("target", "b")) + w.WriteHeader(http.StatusOK) + })) + DeferCleanup(srv.Close) + + Expect(New(srv.URL, "").PinFailoverTarget(context.Background(), "chain", "b")).To(Succeed()) + }) + }) + + Describe("UnpinFailoverTarget", func() { + It("issues DELETE /api/failover/:chain/pin", func() { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + Expect(r.Method).To(Equal(http.MethodDelete)) + Expect(r.URL.Path).To(Equal("/api/failover/chain/pin")) + w.WriteHeader(http.StatusOK) + })) + DeferCleanup(srv.Close) + + Expect(New(srv.URL, "").UnpinFailoverTarget(context.Background(), "chain")).To(Succeed()) + }) + }) +}) + var _ = Describe("ErrHTTPNotFound", func() { Context("on a clean 404 status", func() { var ( diff --git a/pkg/mcp/localaitools/httpapi/routes.go b/pkg/mcp/localaitools/httpapi/routes.go index fca25efad..229e6d0ff 100644 --- a/pkg/mcp/localaitools/httpapi/routes.go +++ b/pkg/mcp/localaitools/httpapi/routes.go @@ -34,6 +34,7 @@ const ( routeMiddleware = "/api/middleware/status" routeRouterDecisions = "/api/router/decisions" routeVoiceProfiles = "/api/voice-profiles" + routeFailover = "/api/failover" ) func routeJobStatus(jobID string) string { diff --git a/pkg/mcp/localaitools/inproc/client.go b/pkg/mcp/localaitools/inproc/client.go index 207956c56..aa3be911c 100644 --- a/pkg/mcp/localaitools/inproc/client.go +++ b/pkg/mcp/localaitools/inproc/client.go @@ -21,6 +21,7 @@ import ( "github.com/mudler/LocalAI/core/gallery/importers" "github.com/mudler/LocalAI/core/http/auth" "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/services/modeladmin" "github.com/mudler/LocalAI/core/services/nodes" @@ -83,6 +84,12 @@ type Client struct { RouterEmbedderFingerprint func(modelName string) (string, error) RouterVectorStore func(storeName string) backend.VectorStore + // Failover backs list_failover_chains / pin_failover_target / + // unpin_failover_target. nil makes the tools report "failover is not + // running" — the same as a deployment with no failover chains + // configured. + Failover *failover.Manager + modelAdmin *modeladmin.ConfigService } @@ -1103,3 +1110,39 @@ func (c *Client) ClearRouterCorpus(ctx context.Context, routerModel string) (*lo } return &localaitools.RouterCorpusClearResult{Router: cfg.Name, Cleared: cleared}, nil } + +// ---- Failover chains ---- + +func (c *Client) ListFailoverChains(_ context.Context) ([]localaitools.FailoverChainInfo, error) { + out := []localaitools.FailoverChainInfo{} + if c.Failover == nil { + return out, nil + } + for _, ch := range c.Failover.Status() { + info := localaitools.FailoverChainInfo{Name: ch.Name, State: string(ch.State), Active: ch.Active} + if ch.Pinned != nil { + info.Pinned = *ch.Pinned + } + for _, t := range ch.Targets { + info.Targets = append(info.Targets, localaitools.FailoverTargetInfo{ + Model: t.Model, Kind: string(t.Kind), Warm: t.Warm, State: string(t.State), LastError: t.LastError, + }) + } + out = append(out, info) + } + return out, nil +} + +func (c *Client) PinFailoverTarget(_ context.Context, chain, target string) error { + if c.Failover == nil { + return errors.New("failover is not running") + } + return c.Failover.Pin(chain, target) +} + +func (c *Client) UnpinFailoverTarget(_ context.Context, chain string) error { + if c.Failover == nil { + return errors.New("failover is not running") + } + return c.Failover.Unpin(chain) +} diff --git a/pkg/mcp/localaitools/inproc/client_test.go b/pkg/mcp/localaitools/inproc/client_test.go index 38a05a8a7..0a9e6a5b7 100644 --- a/pkg/mcp/localaitools/inproc/client_test.go +++ b/pkg/mcp/localaitools/inproc/client_test.go @@ -13,6 +13,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/nodes" localaitools "github.com/mudler/LocalAI/pkg/mcp/localaitools" @@ -129,6 +130,95 @@ var _ = Describe("inproc.Client model aliases", func() { }) }) +// fakeFailoverSource is a minimal failover.ConfigSource over an in-memory +// map, so these specs don't need a real ModelConfigLoader + on-disk YAML. +type fakeFailoverSource struct { + cfgs map[string]config.ModelConfig +} + +func (s *fakeFailoverSource) GetModelConfig(name string) (config.ModelConfig, bool) { + c, ok := s.cfgs[name] + return c, ok +} + +func (s *fakeFailoverSource) GetAllModelsConfigs() []config.ModelConfig { + out := make([]config.ModelConfig, 0, len(s.cfgs)) + for _, c := range s.cfgs { + out = append(out, c) + } + return out +} + +var _ = Describe("inproc.Client failover chains", func() { + var ( + ctx context.Context + c *Client + fm *failover.Manager + ) + + BeforeEach(func() { + ctx = context.Background() + src := &fakeFailoverSource{cfgs: map[string]config.ModelConfig{ + "a": {Name: "a", Backend: "llama-cpp"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{ + Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}, + }}, + }} + fm = failover.New(src) + c = &Client{Failover: fm} + }) + + It("ListFailoverChains reports the chain, its active target, and target health", func() { + out, err := c.ListFailoverChains(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(out).To(HaveLen(1)) + Expect(out[0].Name).To(Equal("chain")) + Expect(out[0].Active).To(Equal("a")) + Expect(out[0].Pinned).To(BeEmpty()) + Expect(out[0].Targets).To(HaveLen(2)) + Expect(out[0].Targets[0].Model).To(Equal("a")) + Expect(out[0].Targets[0].Kind).To(Equal("local")) + Expect(out[0].Targets[0].State).To(Equal("healthy")) + }) + + It("returns an empty slice, not an error, when no failover manager is wired", func() { + c = &Client{} + out, err := c.ListFailoverChains(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(out).To(BeEmpty()) + }) + + It("PinFailoverTarget pins the chain and ListFailoverChains reflects it", func() { + Expect(c.PinFailoverTarget(ctx, "chain", "b")).To(Succeed()) + + out, err := c.ListFailoverChains(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(out[0].Pinned).To(Equal("b")) + }) + + It("PinFailoverTarget errors when the failover manager is unavailable", func() { + c = &Client{} + err := c.PinFailoverTarget(ctx, "chain", "b") + Expect(err).To(HaveOccurred()) + }) + + It("UnpinFailoverTarget clears a pin", func() { + Expect(c.PinFailoverTarget(ctx, "chain", "b")).To(Succeed()) + Expect(c.UnpinFailoverTarget(ctx, "chain")).To(Succeed()) + + out, err := c.ListFailoverChains(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(out[0].Pinned).To(BeEmpty()) + }) + + It("UnpinFailoverTarget errors when the failover manager is unavailable", func() { + c = &Client{} + err := c.UnpinFailoverTarget(ctx, "chain") + Expect(err).To(HaveOccurred()) + }) +}) + var _ = Describe("inproc.Client model scheduling", func() { var ( ctx context.Context diff --git a/pkg/mcp/localaitools/prompts/10_safety.md b/pkg/mcp/localaitools/prompts/10_safety.md index 15671fdc4..5f9c5f840 100644 --- a/pkg/mcp/localaitools/prompts/10_safety.md +++ b/pkg/mcp/localaitools/prompts/10_safety.md @@ -2,7 +2,7 @@ These rules are non-negotiable. The user trusts you to operate their server without unintended changes. -1. **Confirm before mutating.** Before calling any of these tools — `install_model`, `import_model_uri`, `delete_model`, `install_backend`, `upgrade_backend`, `edit_model_config`, `reload_models`, `load_model`, `toggle_model_state`, `toggle_model_pinned`, `set_branding`, `set_alias`, `seed_router_corpus`, `clear_router_corpus`, `create_voice_profile`, `delete_voice_profile`, `set_node_vram_budget`, `set_scheduling`, `delete_scheduling` — first state in plain language what you are about to do (which tool, which target, which arguments) and wait for the user's explicit confirmation in the next turn. "Yes", "do it", "go ahead", "proceed" all count as confirmation. Anything else does not. +1. **Confirm before mutating.** Before calling any of these tools — `install_model`, `import_model_uri`, `delete_model`, `install_backend`, `upgrade_backend`, `edit_model_config`, `reload_models`, `load_model`, `toggle_model_state`, `toggle_model_pinned`, `set_branding`, `set_alias`, `seed_router_corpus`, `clear_router_corpus`, `create_voice_profile`, `delete_voice_profile`, `set_node_vram_budget`, `set_scheduling`, `delete_scheduling`, `pin_failover_target`, `unpin_failover_target` — first state in plain language what you are about to do (which tool, which target, which arguments) and wait for the user's explicit confirmation in the next turn. "Yes", "do it", "go ahead", "proceed" all count as confirmation. Anything else does not. 2. **Disambiguate before mutating.** If the user's request is ambiguous (several gallery candidates match, the model name has multiple installed versions, the backend has variants), present the candidates as a numbered list and ask the user to pick before calling any mutating tool. diff --git a/pkg/mcp/localaitools/prompts/20_tools.md b/pkg/mcp/localaitools/prompts/20_tools.md index 2a2830417..88404801e 100644 --- a/pkg/mcp/localaitools/prompts/20_tools.md +++ b/pkg/mcp/localaitools/prompts/20_tools.md @@ -24,6 +24,7 @@ The MCP `tools/list` endpoint also exposes the full input schema for each of the - `get_router_decisions` — Inspect recent router decisions and classifier signals. - `get_router_corpus_stats` — Inspect a KNN router corpus by count and label only; exemplar texts are never returned. - `list_aliases` — List configured model aliases and their targets. +- `list_failover_chains` — List failover chains, their active target and target health. ## Mutating (require user confirmation per safety rule 1) @@ -46,3 +47,5 @@ The MCP `tools/list` endpoint also exposes the full input schema for each of the - `set_node_vram_budget` — Set or clear a federated node's VRAM budget override. - `set_scheduling` — Create or update a distributed per-model scheduling config. - `delete_scheduling` — Remove a distributed per-model scheduling config. +- `pin_failover_target` — Force a failover chain to one target. +- `unpin_failover_target` — Remove a failover pin. diff --git a/pkg/mcp/localaitools/server.go b/pkg/mcp/localaitools/server.go index 8a6c64b86..f47332db5 100644 --- a/pkg/mcp/localaitools/server.go +++ b/pkg/mcp/localaitools/server.go @@ -54,6 +54,7 @@ func NewServer(client LocalAIClient, opts Options) *mcp.Server { registerUsageTools(srv, client, opts) registerPIITools(srv, client, opts) registerMiddlewareTools(srv, client, opts) + registerFailoverTools(srv, client, opts) return srv } diff --git a/pkg/mcp/localaitools/server_test.go b/pkg/mcp/localaitools/server_test.go index c18a2301c..6fe104b9e 100644 --- a/pkg/mcp/localaitools/server_test.go +++ b/pkg/mcp/localaitools/server_test.go @@ -92,6 +92,7 @@ var expectedReadOnlyCatalog = sortedStrings( ToolListVoiceProfiles, ToolSystemInfo, ToolVRAMEstimate, + ToolListFailoverChains, ) // expectedFullCatalog derives from the read-only catalog plus the canonical @@ -155,6 +156,9 @@ var _ = Describe("Tool dispatch", func() { {ToolListAliases, struct{}{}, "ListAliases"}, {ToolCreateVoiceProfile, CreateVoiceProfileRequest{Name: "Narrator", Transcript: "Reference words", AudioBase64: "UklGRg==", ConsentConfirmed: true}, "CreateVoiceProfile"}, {ToolDeleteVoiceProfile, DeleteVoiceProfileRequest{ID: "00000000-0000-0000-0000-000000000001"}, "DeleteVoiceProfile"}, + {ToolListFailoverChains, map[string]any{}, "ListFailoverChains"}, + {ToolPinFailoverTarget, map[string]any{"chain": "c", "target": "b"}, "PinFailoverTarget"}, + {ToolUnpinFailoverTarget, map[string]any{"chain": "c"}, "UnpinFailoverTarget"}, } for _, c := range cases { diff --git a/pkg/mcp/localaitools/tools.go b/pkg/mcp/localaitools/tools.go index cfc7f0bfd..e2c9b32d3 100644 --- a/pkg/mcp/localaitools/tools.go +++ b/pkg/mcp/localaitools/tools.go @@ -49,10 +49,19 @@ const ( ToolSetNodeVRAMBudget = "set_node_vram_budget" ToolSetScheduling = "set_scheduling" ToolDeleteScheduling = "delete_scheduling" + // ToolPinFailoverTarget and ToolUnpinFailoverTarget live here (rather + // than grouped with ToolListFailoverChains below) so mutatingToolNames + // stays a contiguous scan of this block. + ToolPinFailoverTarget = "pin_failover_target" + ToolUnpinFailoverTarget = "unpin_failover_target" // ToolListAliases is read-only but lives here so the alias tools stay // grouped; the catalog tests assert its read-only placement. ToolListAliases = "list_aliases" + + // ToolListFailoverChains is read-only but lives here so the failover + // tools stay grouped; the catalog tests assert its read-only placement. + ToolListFailoverChains = "list_failover_chains" ) // DefaultServerName is the MCP Implementation.Name surfaced when @@ -83,4 +92,6 @@ var mutatingToolNames = []string{ ToolSetNodeVRAMBudget, ToolSetScheduling, ToolDeleteScheduling, + ToolPinFailoverTarget, + ToolUnpinFailoverTarget, } diff --git a/pkg/mcp/localaitools/tools_failover.go b/pkg/mcp/localaitools/tools_failover.go new file mode 100644 index 000000000..ca3368896 --- /dev/null +++ b/pkg/mcp/localaitools/tools_failover.go @@ -0,0 +1,53 @@ +package localaitools + +import ( + "context" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +// registerFailoverTools wires the conversational failover-chain tools. +// list_failover_chains reports the health of every chain, pin_failover_target +// forces a chain to one target, and unpin_failover_target hands control back +// to health-based selection. +func registerFailoverTools(s *mcp.Server, client LocalAIClient, opts Options) { + mcp.AddTool(s, &mcp.Tool{ + Name: ToolListFailoverChains, + Description: "List model failover chains, the target serving each one now, and the health of every target.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, any, error) { + chains, err := client.ListFailoverChains(ctx) + if err != nil { + return errorResult(err), nil, nil + } + return jsonResult(chains), nil, nil + }) + + if opts.DisableMutating { + return + } + + mcp.AddTool(s, &mcp.Tool{ + Name: ToolPinFailoverTarget, + Description: "Force a failover chain to serve every request from one target, regardless of health, until it is unpinned. Requires user confirmation per safety rule 1.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, args struct { + Chain string `json:"chain" jsonschema:"failover chain name"` + Target string `json:"target" jsonschema:"target model to pin"` + }) (*mcp.CallToolResult, any, error) { + if err := client.PinFailoverTarget(ctx, args.Chain, args.Target); err != nil { + return errorResult(err), nil, nil + } + return jsonResult(map[string]string{"chain": args.Chain, "pinned": args.Target}), nil, nil + }) + + mcp.AddTool(s, &mcp.Tool{ + Name: ToolUnpinFailoverTarget, + Description: "Remove the pin from a failover chain so health decides the target again. Requires user confirmation per safety rule 1.", + }, func(ctx context.Context, _ *mcp.CallToolRequest, args struct { + Chain string `json:"chain" jsonschema:"failover chain name"` + }) (*mcp.CallToolResult, any, error) { + if err := client.UnpinFailoverTarget(ctx, args.Chain); err != nil { + return errorResult(err), nil, nil + } + return jsonResult(map[string]string{"chain": args.Chain, "pinned": ""}), nil, nil + }) +} From dcc7bd854821db279086d524c5f8563a5916a807 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:06:00 +0000 Subject: [PATCH 22/79] docs: document model failover chains Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- docs/content/features/model-aliases.md | 3 + docs/content/features/model-failover.md | 179 +++++++++++++++++++++++ docs/content/features/openai-realtime.md | 2 + docs/content/operations/cloud-proxy.md | 3 + 4 files changed, 187 insertions(+) create mode 100644 docs/content/features/model-failover.md diff --git a/docs/content/features/model-aliases.md b/docs/content/features/model-aliases.md index ed52f0977..292c05374 100644 --- a/docs/content/features/model-aliases.md +++ b/docs/content/features/model-aliases.md @@ -38,6 +38,9 @@ That is the whole config: a `name` (the alias clients call) and an `alias` key - Usage accounting records both sides: requested `gpt-4`, served `my-llama-3`. - Aliases work for every modality (chat, embeddings, audio, images, and so on). +To serve a name from several models with automatic fallback, use a [failover +chain]({{%relref "features/model-failover" %}}). + ## Managing aliases You can create, swap, and remove aliases from any of the management surfaces. diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md new file mode 100644 index 000000000..b3da3f7d9 --- /dev/null +++ b/docs/content/features/model-failover.md @@ -0,0 +1,179 @@ + ++++ +disableToc = false +title = "Model Failover" +weight = 15 +url = "/features/model-failover/" ++++ + +A **failover chain** is a model name that is served by an ordered list of +other models. LocalAI sends each request to the first healthy target. When a +target fails, the request moves to the next target, and later requests stay +there until the first target has recovered. + +Use it to serve a model from a remote LocalAI or another OpenAI-compatible +provider, and to fall back to a local model when the remote one is down. + +## Declaring a chain + +```yaml +name: assistant-llm +failover: + targets: + - model: argus-llm # for example a cloud-proxy model + - model: gemma-local + warm: true # keep it loaded +``` + +Clients call `assistant-llm`. Each target is a normal model config. A chain +has no `backend` and no `parameters.model`. + +Optional settings, with their defaults: + +```yaml +failover: + probe: + interval: 15s # how often an idle target is checked + timeout: 5s + trip: + errors: 1 # failures within the window that mark a target down + window: 30s + recovery: + probes: 3 # test requests a target must pass before it is used again + min_dwell: 60s # minimum time on a lower target before moving back +``` + +Rules: + +- A chain needs at least 2 targets. A target can be an alias, but not another + chain. +- A chain cannot also set `alias` or `backend`. +- Responses name the chain as the model. The `X-LocalAI-Served-Model` header + names the target that served the request. + +## How the target is chosen + +- The active target is the first healthy target in the list. +- When a target fails, LocalAI marks it down and moves to the next target at + once. +- LocalAI moves back to a higher target only when that target has passed + `recovery.probes` test requests **and** the current target has been active + for at least `recovery.min_dwell`. This stops an unstable upstream from + moving traffic back and forth. +- When all targets are down, the chain is `degraded`. Each request still tries + every target in order. + +## Retry inside a request + +When a target fails before the response starts, LocalAI sends the same +request to the next target. The client does not see the failure. + +- LocalAI does not retry after the first byte of a response is sent (for + example after the first streamed token). The request fails, the target is + marked down, and the next request uses the next target. +- A target that is at its concurrency limit (an admission rejection) or that + is disabled is skipped for that request without being marked down. +- LocalAI does not retry client errors (4xx), such as a prompt that is too + long, because the next target would reject it too. A 4xx counts neither as + a success nor as a failure for the target. +- Request bodies larger than 32 MiB are not retried. + +When the primary did not serve the request, the response has the header +`X-LocalAI-Failover: fallback`, or `X-LocalAI-Failover: degraded` when all +targets were down. + +## Health checks + +| Target | Regular check | Check before moving back | +|---|---|---| +| Remote (`cloud-proxy`) | `GET /v1/models` on the upstream lists the model | one small real request, for example a 1-token completion | +| Local, `warm: true` | the backend answers a health check | one small real request | +| Local, not warm | none: judged only by real requests; it is never loaded only to check it | none: the target is used again after `min_dwell` | + +A request that succeeds counts as a check, so a busy target is almost never +probed. + +When a target is in more than one chain, its check settings come from the +first of those chains in name order. + +## Warm targets + +`warm: true` loads a local target at startup and protects it from idle and +LRU eviction, so a switch does not wait for the model to load. Warm targets +count toward the active backend limit (`--max-active-backends`) like any +pinned model: LocalAI never evicts them to make room, and if they fill the +limit, a new model still loads rather than being blocked. + +## Realtime pipelines + +A pipeline stage can name a chain: + +```yaml +name: assistant +pipeline: + vad: silero-vad + transcription: whisper-chain + llm: assistant-llm + tts: voice-chain +``` + +LocalAI resolves the chain for every call of the stage. When a chain switches, +the session stays open and keeps its conversation. The next turn uses the new +target. + +The session receives a `localai.model.failover` event for each chain stage when +it starts (`reason: initial`) and each time a chain switches: + +```json +{"type":"localai.model.failover","chain":"assistant-llm","stage":"llm", + "from":"argus-llm","to":"gemma-local","state":"fallback","reason":"trip"} +``` + +Limits: + +- Chains are resolved only in full realtime pipelines. A transcription-only or + sound-detection-only session does not resolve chains yet. +- After a `session.update` that changes the pipeline, `localai.model.failover` + events keep describing the chains from session start. +- A chain used as a router candidate, or as the classifier-mode scoring model, + is not resolved per call. + +## Watching failover + +- `GET /api/failover` lists every chain, its active target and the state of + each target. +- `GET /api/failover/{chain}` returns one chain. +- `GET /api/failover/events` is a server-sent event stream. The first event is + `snapshot` with the full state. Then `chain.switched` and `target.state` + events follow. +- Metrics: `localai_failover_switches_total{chain,from,to,reason}` and + `localai_failover_target_up{target}`. +- With tracing on, each skipped target appears in the Traces view with the + error that made LocalAI skip it. + +## Pinning a target + +An admin can force a chain to one target, for example during maintenance: + +```bash +curl -X POST http://localhost:8080/api/failover/assistant-llm/pin \ + -H 'Content-Type: application/json' -d '{"target":"gemma-local"}' +curl -X DELETE http://localhost:8080/api/failover/assistant-llm/pin +``` + +While a chain is pinned, only the pinned target serves it. Health checks +continue. A restart removes the pin. + +## Assistant and MCP + +The LocalAI Assistant and `local-ai mcp-server` offer `list_failover_chains`, +`pin_failover_target` and `unpin_failover_target`. Create and edit chains with +the model config tools, like any other model. + +## Limits + +- Failover state is kept in memory by each LocalAI instance. Several frontends + in distributed mode each keep their own view. +- Chains do not nest. +- See also [model aliases]({{%relref "features/model-aliases" %}}) and the + [realtime API]({{%relref "features/openai-realtime" %}}). diff --git a/docs/content/features/openai-realtime.md b/docs/content/features/openai-realtime.md index 54bec58ab..ac7afb484 100644 --- a/docs/content/features/openai-realtime.md +++ b/docs/content/features/openai-realtime.md @@ -31,6 +31,8 @@ This configuration links the following components: Make sure all referenced models (`silero-vad-ggml`, `whisper-large-turbo`, `qwen3-4b`, `tts-1`) are also installed or defined in your LocalAI instance. +A pipeline stage can name a [failover chain]({{%relref "features/model-failover" %}}); the stage then switches targets without closing the session. + ### Streaming the pipeline By default each stage runs to completion before the next begins: the whole utterance is transcribed, the full LLM reply is generated, then it is synthesized. Each stage can instead be streamed incrementally, which lowers the time-to-first-audio of a turn: diff --git a/docs/content/operations/cloud-proxy.md b/docs/content/operations/cloud-proxy.md index 02af25bd0..bd1d806d4 100644 --- a/docs/content/operations/cloud-proxy.md +++ b/docs/content/operations/cloud-proxy.md @@ -29,6 +29,9 @@ egress remains subject to the same redaction rules a local model would apply. - Use the intelligent router to send small or simple prompts to a local model and complex ones to Claude or GPT-4o. +To fall back to a local model when the upstream is down, list the proxy model +in a [failover chain]({{%relref "features/model-failover" %}}). + ## How it works 1. Request hits LocalAI on `/v1/chat/completions` (OpenAI-shaped) or From b67b75dcdd2e229d1ff617ef942fe1492bd082da Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:18:20 +0000 Subject: [PATCH 23/79] fix(failover): answer HasChains from a flag set at sync The request path called HasChains on every request, and with no chains it scanned the config source each time: the loader's lock plus a copy and sort of every config, forever, on every installation without chains. Sync now keeps an atomic flag and HasChains reads only that. A chain added since the last sync is still served because Plan syncs on a miss; only in-request retry waits for the next tick (at most 1s). Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/middleware/failover_test.go | 3 +++ core/services/failover/fakes_test.go | 7 +++++-- core/services/failover/manager.go | 28 +++++++++++--------------- core/services/failover/manager_test.go | 26 +++++++++++++++++++----- 4 files changed, 41 insertions(+), 23 deletions(-) diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index a16337545..81fff5e02 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -87,6 +87,9 @@ var _ = Describe("failover chains in the request pipeline", func() { 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() diff --git a/core/services/failover/fakes_test.go b/core/services/failover/fakes_test.go index cad6b3552..232b7b4df 100644 --- a/core/services/failover/fakes_test.go +++ b/core/services/failover/fakes_test.go @@ -26,8 +26,9 @@ func (c *fakeClock) Advance(d time.Duration) { } type fakeSource struct { - mu sync.Mutex - cfgs map[string]config.ModelConfig + mu sync.Mutex + cfgs map[string]config.ModelConfig + scans int // GetAllModelsConfigs calls } func newFakeSource(cfgs ...config.ModelConfig) *fakeSource { @@ -39,6 +40,7 @@ func newFakeSource(cfgs ...config.ModelConfig) *fakeSource { } func (s *fakeSource) Put(c config.ModelConfig) { s.mu.Lock(); s.cfgs[c.Name] = c; s.mu.Unlock() } func (s *fakeSource) Delete(name string) { s.mu.Lock(); delete(s.cfgs, name); s.mu.Unlock() } +func (s *fakeSource) Scans() int { s.mu.Lock(); defer s.mu.Unlock(); return s.scans } func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { s.mu.Lock() defer s.mu.Unlock() @@ -48,6 +50,7 @@ func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { func (s *fakeSource) GetAllModelsConfigs() []config.ModelConfig { s.mu.Lock() defer s.mu.Unlock() + s.scans++ out := make([]config.ModelConfig, 0, len(s.cfgs)) for _, c := range s.cfgs { out = append(out, c) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index ab130d1ee..ea31ee999 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -55,6 +55,9 @@ type Manager struct { warm []string warmPending bool closed bool + // hasChains mirrors len(chains) > 0 as of the last sync, so the request + // path can check it without the lock or a config-source scan. + hasChains atomic.Bool } type targetState struct { @@ -168,6 +171,7 @@ func (m *Manager) syncLocked() { delete(m.targets, name) } } + m.hasChains.Store(len(m.chains) > 0) for _, ch := range m.chains { m.recomputeLocked(ch, "") } @@ -201,26 +205,18 @@ func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { return c, ok } -// HasChains reports whether any failover chain is configured. The request -// path uses it to skip chain bookkeeping on installations without chains. -// Chains are synced lazily, so with none known yet the config source is -// consulted, which catches a chain added since the last sync. +// HasChains reports whether any failover chain was configured at the last +// sync. The request path calls it on every request to skip chain bookkeeping +// on installations without chains, so it reads a flag instead of scanning +// the config source (which takes the loader's lock and copies every config). +// A chain added since the last sync is still served, because Plan syncs on a +// miss; only in-request retry is missing for it until the scheduler's next +// tick, at most one second later. func (m *Manager) HasChains() bool { if m == nil { return false } - m.mu.Lock() - known := len(m.chains) > 0 - m.mu.Unlock() - if known { - return true - } - for _, c := range m.src.GetAllModelsConfigs() { - if c.IsFailover() { - return true - } - } - return false + return m.hasChains.Load() } func (m *Manager) chainLocked(name string) *chainState { diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index d0736cf19..0d2746331 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -60,18 +60,34 @@ var _ = Describe("Manager", func() { Expect(st.Targets[0].State).To(Equal(StateHealthy)) }) - It("reports whether any chain is configured, including one added since the last sync", func() { + It("reports whether any chain was configured at the last sync", func() { + m.Sync() + Expect(m.HasChains()).To(BeTrue()) + var nilManager *Manager + Expect(nilManager.HasChains()).To(BeFalse()) + }) + + It("answers HasChains from the last sync without scanning the config source", func() { empty := New(newFakeSource(remote("a")), WithClock(clock)) Expect(empty.HasChains()).To(BeFalse()) lateSrc := newFakeSource(remote("a"), local("b")) late := New(lateSrc, WithClock(clock)) late.Sync() - Expect(late.HasChains()).To(BeFalse()) + scans := lateSrc.Scans() + for range 100 { + Expect(late.HasChains()).To(BeFalse()) + } + Expect(lateSrc.Scans()).To(Equal(scans)) + // A chain added since the last sync is seen at the next sync, or + // sooner by Plan, which syncs on a miss. lateSrc.Put(chainCfg("chain", nil, t("a"), t("b"))) + Expect(late.HasChains()).To(BeFalse()) + _, err := late.Plan("chain") + Expect(err).ToNot(HaveOccurred()) Expect(late.HasChains()).To(BeTrue()) - Expect(m.HasChains()).To(BeTrue()) - var nilManager *Manager - Expect(nilManager.HasChains()).To(BeFalse()) + lateSrc.Delete("chain") + late.Sync() + Expect(late.HasChains()).To(BeFalse()) }) It("returns ErrChainNotFound for an unknown chain", func() { From 5d2b90b8484580fc9c41c9926f68eb5d5397cab0 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:22:14 +0000 Subject: [PATCH 24/79] fix(failover): never load a warm target inside a probe A warm target's liveness probe called ModelLoader.Load, which blocked until the model finished loading (while the warm preload loaded it too). Tick waited for every probe, so all probing froze, and the probe then ran HealthCheck on an expired context and tripped the target at every startup. The prober now takes a function that returns the running backend without loading it. A target that is not loaded passes liveness; its recovery is neither confirmed nor failed and it returns to healthy after min_dwell, like a cold target. Tick no longer waits for probes: each probe applies its own result and a target whose probe is running is skipped. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/failover.go | 20 ++++ core/application/failover_test.go | 22 ++++ core/application/startup.go | 6 +- core/services/failover/manager.go | 5 + core/services/failover/prober.go | 56 +++++---- core/services/failover/prober_test.go | 22 +++- core/services/failover/schedule.go | 47 ++++++-- core/services/failover/schedule_test.go | 111 +++++++++++++++--- docs/content/features/model-failover.md | 2 +- ...2026-09-26-model-failover-chains-design.md | 8 +- 10 files changed, 242 insertions(+), 57 deletions(-) diff --git a/core/application/failover.go b/core/application/failover.go index f0b5c2c2b..a513c1f10 100644 --- a/core/application/failover.go +++ b/core/application/failover.go @@ -2,6 +2,10 @@ 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" ) @@ -30,3 +34,19 @@ func (a *Application) applyFailoverWarmTargets(warm []string) { } }() } + +// failoverLoadedBackend gives the failover prober the running backend of a +// local target without ever loading it. CheckIsLoaded may run the loader's +// own health check and drop a dead process; the target is then "not loaded" +// and the next real request loads and judges it. +func failoverLoadedBackend(ml *model.ModelLoader) failover.LoadedFunc { + return func(cfg config.ModelConfig) grpc.Backend { + m := ml.CheckIsLoaded(cfg.ModelID()) + if m == nil { + return nil + } + // Load always enables parallel requests; match it in case this is + // the first client built for the model. + return m.GRPC(true, ml.GetWatchDog()) + } +} diff --git a/core/application/failover_test.go b/core/application/failover_test.go index a9057425a..14eaef241 100644 --- a/core/application/failover_test.go +++ b/core/application/failover_test.go @@ -5,7 +5,9 @@ import ( "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" ) @@ -39,3 +41,23 @@ var _ = Describe("applyFailoverWarmTargets", func() { Eventually(started, time.Second).Should(BeClosed(), "the preload goroutine should still run in the background") }) }) + +type healthyBackend struct{ grpc.Backend } + +func (healthyBackend) HealthCheck(context.Context) (bool, error) { return true, nil } + +var _ = Describe("failoverLoadedBackend", func() { + It("returns the running backend and never loads a model that is not loaded", func() { + ml := model.NewModelLoader(&system.SystemState{Model: system.Model{ModelsPath: GinkgoT().TempDir()}}) + store := model.NewInMemoryModelStore() + ml.SetModelStore(store) + loaded := failoverLoadedBackend(ml) + + Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(BeNil()) + Expect(ml.ListLoadedModels()).To(BeEmpty(), "the lookup must not start a load") + + client := healthyBackend{} + store.Set("gemma", model.NewModelWithClient("gemma", "127.0.0.1:0", client)) + Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(Equal(client)) + }) +}) diff --git a/core/application/startup.go b/core/application/startup.go index 938f57496..4c1ddd4e9 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -1,7 +1,6 @@ package application import ( - "context" "crypto/rand" "encoding/hex" "fmt" @@ -29,7 +28,6 @@ import ( "github.com/mudler/LocalAI/core/trace" "github.com/mudler/LocalAI/internal" "github.com/mudler/LocalAI/pkg/downloader" - "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/modelartifacts" "github.com/mudler/LocalAI/pkg/signals" "github.com/mudler/LocalAI/pkg/vram" @@ -257,9 +255,7 @@ func New(opts ...config.AppOption) (*Application, error) { // 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(func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) { - return application.ModelLoader().Load(backend.ModelOptions(cfg, options)...) - })), + failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()))), failover.WithOnWarmChanged(application.applyFailoverWarmTargets), ) // The assistant client was built in start() (above), before this diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index ea31ee999..847584d62 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -58,6 +58,8 @@ type Manager struct { // hasChains mirrors len(chains) > 0 as of the last sync, so the request // path can check it without the lock or a config-source scan. hasChains atomic.Bool + // probes counts running probes; only tests wait on it. + probes sync.WaitGroup } type targetState struct { @@ -71,6 +73,9 @@ type targetState struct { lastProbe time.Time lastActivity time.Time lastError string + // probing is set while a probe runs, so the scheduler does not start a + // second one for the same target. + probing bool // params come from the first chain, in name order, that lists the target. params config.FailoverConfig } diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index 592d9ca41..03108d0c9 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -18,18 +18,26 @@ import ( pb "github.com/mudler/LocalAI/pkg/grpc/proto" ) -// LoadFunc returns the backend for a local target, loading it if needed. -type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) +// ErrNotLoaded is what Inference returns for a local target whose backend is +// not running. It neither confirms nor fails recovery: the manager judges the +// target like a cold one, by real requests once min_dwell has passed. +var ErrNotLoaded = errors.New("failover: target is not loaded") + +// LoadedFunc returns the running backend of a local target, or nil when it is +// not loaded. It must never load the model: a probe that loads blocks until +// the load ends (while the warm preload loads the same model) and then judges +// the target on an expired context. +type LoadedFunc func(cfg config.ModelConfig) grpc.Backend // DefaultProber probes remote targets over the upstream's OpenAI-compatible // API and local targets through their gRPC backend. type DefaultProber struct { - HTTP *http.Client - Load LoadFunc + HTTP *http.Client + Loaded LoadedFunc } -func NewProber(load LoadFunc) *DefaultProber { - return &DefaultProber{HTTP: &http.Client{}, Load: load} +func NewProber(loaded LoadedFunc) *DefaultProber { + return &DefaultProber{HTTP: &http.Client{}, Loaded: loaded} } func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { @@ -225,15 +233,25 @@ func silenceWAV() []byte { return b } +func (p *DefaultProber) loaded(cfg config.ModelConfig) grpc.Backend { + if p.Loaded == nil { + return nil + } + return p.Loaded(cfg) +} + +// localHealth checks a warm target's running backend. A target that is not +// loaded passes: the warm preload is loading it, or a crash removed it and +// the next real request loads it again and judges it. func (p *DefaultProber) localHealth(ctx context.Context, cfg config.ModelConfig) error { - if p.Load == nil { - return errors.New("failover: no backend loader configured") - } - // Load returns the running backend, or starts it again after a crash. - b, err := p.Load(ctx, cfg) - if err != nil { - return err + b := p.loaded(cfg) + if b == nil { + return nil } + return healthCheck(ctx, b) +} + +func healthCheck(ctx context.Context, b grpc.Backend) error { ok, err := b.HealthCheck(ctx) if err != nil { return err @@ -245,13 +263,11 @@ func (p *DefaultProber) localHealth(ctx context.Context, cfg config.ModelConfig) } func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConfig) error { - if p.Load == nil { - return errors.New("failover: no backend loader configured") - } - b, err := p.Load(ctx, cfg) - if err != nil { - return err + b := p.loaded(cfg) + if b == nil { + return ErrNotLoaded } + var err error switch { case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION): _, err = b.Predict(ctx, &pb.PredictOptions{Prompt: "ping", Tokens: 1}) @@ -262,5 +278,5 @@ func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConf } // A backend process that answers HealthCheck rarely fails only for TTS or // transcription, so a real request adds little here. - return p.localHealth(ctx, cfg) + return healthCheck(ctx, b) } diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index 0f65f8dc4..72f03887b 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -138,7 +138,7 @@ var _ = Describe("DefaultProber", func() { It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { b := &fakeBackend{healthy: true} - p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }) + p = NewProber(func(config.ModelConfig) grpc.Backend { return b }) c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) @@ -150,10 +150,24 @@ var _ = Describe("DefaultProber", func() { Expect(p.Inference(ctx, c, KindLocal, true)).To(HaveOccurred()) }) + It("passes warm liveness for a target that is not loaded, and leaves recovery unconfirmed", func() { + // The warm preload loads the model; a probe that loaded it too would + // block until the load finished and then judge it on an expired ctx. + asked := 0 + p = NewProber(func(config.ModelConfig) grpc.Backend { asked++; return nil }) + for _, uc := range []string{"chat", "tts"} { + c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{uc}} + c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) + Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) + Expect(p.Inference(ctx, c, KindLocal, true)).To(MatchError(ErrNotLoaded), uc) + } + Expect(asked).To(Equal(4)) + }) + It("passes cold local liveness without a model file and without loading", func() { - p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { - Fail("cold liveness must not load the model") - return nil, nil + p = NewProber(func(config.ModelConfig) grpc.Backend { + Fail("cold liveness must not look up the backend") + return nil }) // None of these files exist: a missing file says nothing about whether // the target can serve (download on first use, dotted names, backends diff --git a/core/services/failover/schedule.go b/core/services/failover/schedule.go index 170de3cfa..08f138d00 100644 --- a/core/services/failover/schedule.go +++ b/core/services/failover/schedule.go @@ -2,7 +2,7 @@ package failover import ( "context" - "sync" + "errors" "time" "github.com/mudler/LocalAI/core/config" @@ -26,23 +26,30 @@ func (m *Manager) Run(ctx context.Context) { } } -// Tick runs one pass: sync configs, run due probes, recompute chains. It is +// Tick runs one pass: sync configs, start due probes, recompute chains. It is // exported so tests can drive the manager without a real ticker. Like Run, it // must only be called from the scheduler goroutine (see Run's comment on Sync). +// +// Tick does not wait for the probes it starts: one slow target (a probe that +// hangs until its timeout) must not delay probing and fail-back of every +// other chain. Each probe applies its own result, and a target whose probe is +// still running is skipped until it ends. func (m *Manager) Tick(ctx context.Context) { m.Sync() - var wg sync.WaitGroup for _, j := range m.dueProbes() { - wg.Add(1) + m.probes.Add(1) go func(j probeJob) { - defer wg.Done() + defer m.probes.Done() m.runProbe(ctx, j) }(j) } - wg.Wait() m.Reevaluate() } +// waitProbes waits for the probes started so far. Tests use it to see a +// tick's results; the scheduler never waits. +func (m *Manager) waitProbes() { m.probes.Wait() } + type probeJob struct { target string cfg config.ModelConfig @@ -58,6 +65,9 @@ func (m *Manager) dueProbes() []probeJob { now := m.clock.Now() var jobs []probeJob for _, ts := range m.targets { + if ts.probing { + continue + } interval := ts.params.ProbeInterval() inference := false switch ts.state { @@ -94,6 +104,7 @@ func (m *Manager) dueProbes() []probeJob { continue } ts.lastProbe = now + ts.probing = true jobs = append(jobs, probeJob{ target: ts.name, cfg: cfg, kind: ts.kind, warm: ts.warm, inference: inference, timeout: ts.params.ProbeTimeout(), @@ -112,16 +123,38 @@ func (m *Manager) runProbe(ctx context.Context, j probeJob) { err = m.prober.Liveness(pctx, j.cfg, j.kind, j.warm) } if ctx.Err() != nil { + m.endProbe(j.target) return // shutting down: a cancelled probe says nothing about the target } m.applyProbe(j, err) } +func (m *Manager) endProbe(target string) { + m.mu.Lock() + defer m.mu.Unlock() + if ts := m.targets[target]; ts != nil { + ts.probing = false + } +} + func (m *Manager) applyProbe(j probeJob, err error) { m.mu.Lock() defer m.mu.Unlock() ts := m.targets[j.target] - if ts == nil || ts.state == StateMissing { + if ts == nil { + return + } + ts.probing = false + if ts.state == StateMissing { + return + } + if errors.Is(err, ErrNotLoaded) { + // Nothing running to confirm recovery against: judge the target like + // a cold one, by real requests once min_dwell has passed. + if ts.state == StateRecovering && m.clock.Now().Sub(ts.downSince) >= ts.params.MinDwell() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + m.recomputeForLocked(ts.name) + } return } if err != nil { diff --git a/core/services/failover/schedule_test.go b/core/services/failover/schedule_test.go index 931f1a151..a05e2f982 100644 --- a/core/services/failover/schedule_test.go +++ b/core/services/failover/schedule_test.go @@ -2,6 +2,7 @@ package failover import ( "context" + "slices" "sync" "time" @@ -18,14 +19,19 @@ type probeCall struct { type fakeProber struct { mu sync.Mutex calls []probeCall - fail map[string]error // target -> error returned by every probe + fail map[string]error // target -> error returned by every probe + block map[string]chan struct{} // target -> probes wait until it is closed } func (p *fakeProber) record(target string, inference bool) error { p.mu.Lock() - defer p.mu.Unlock() p.calls = append(p.calls, probeCall{target, inference}) - return p.fail[target] + err, block := p.fail[target], p.block[target] + p.mu.Unlock() + if block != nil { + <-block + } + return err } func (p *fakeProber) Liveness(_ context.Context, c config.ModelConfig, _ Kind, _ bool) error { return p.record(c.Name, false) @@ -52,33 +58,40 @@ var _ = Describe("Manager probes", func() { BeforeEach(func() { clock = newFakeClock() - prober = &fakeProber{fail: map[string]error{}} + prober = &fakeProber{fail: map[string]error{}, block: map[string]chan struct{}{}} src = newFakeSource(remote("a"), local("b"), local("cold"), chainCfg("chain", nil, t("a"), warmT("b"))) m = New(src, WithClock(clock), WithProber(prober)) }) - It("probes idle targets on the first tick and not again before the interval", func() { + // tick runs one scheduler pass and waits for the probes it started, so + // each spec sees their results. + tick := func() { m.Tick(ctx) + m.waitProbes() + } + + It("probes idle targets on the first tick and not again before the interval", func() { + tick() Expect(prober.take()).To(ConsistOf(probeCall{"a", false}, probeCall{"b", false})) clock.Advance(5 * time.Second) - m.Tick(ctx) + tick() Expect(prober.take()).To(BeEmpty()) }) It("skips the liveness probe for a target with recent traffic", func() { - m.Tick(ctx) + tick() prober.take() clock.Advance(14 * time.Second) m.ReportSuccess("a") clock.Advance(2 * time.Second) - m.Tick(ctx) + tick() Expect(prober.take()).To(ConsistOf(probeCall{"b", false})) }) It("trips a target whose liveness probe fails", func() { prober.fail["a"] = errBoom - m.Tick(ctx) + tick() st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateDown)) Expect(st.Active).To(Equal("b")) @@ -86,18 +99,18 @@ var _ = Describe("Manager probes", func() { It("recovers through liveness, then inference probes, then fails back after dwell", func() { prober.fail["a"] = errBoom - m.Tick(ctx) + tick() delete(prober.fail, "a") prober.take() clock.Advance(15 * time.Second) - m.Tick(ctx) // liveness passes: recovering + tick() // liveness passes: recovering st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateRecovering)) for i := 0; i < 3; i++ { clock.Advance(15 * time.Second) - m.Tick(ctx) + tick() } calls := prober.take() Expect(calls).To(ContainElement(probeCall{"a", true})) @@ -109,10 +122,10 @@ var _ = Describe("Manager probes", func() { It("sends a recovering target back down when an inference probe fails", func() { m.ReportFailure("a", errBoom) clock.Advance(15 * time.Second) - m.Tick(ctx) // liveness passes: recovering + tick() // liveness passes: recovering prober.fail["a"] = errBoom clock.Advance(15 * time.Second) - m.Tick(ctx) + tick() st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateDown)) }) @@ -123,21 +136,21 @@ var _ = Describe("Manager probes", func() { m.ReportFailure("cold", errBoom) prober.take() clock.Advance(30 * time.Second) - m.Tick(ctx) + tick() for _, c := range prober.take() { Expect(c.target).ToNot(Equal("cold")) } st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateDown)) clock.Advance(31 * time.Second) - m.Tick(ctx) + tick() st, _ = m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateHealthy)) }) It("probes a target shared by two chains once per tick", func() { src.Put(chainCfg("chain2", nil, t("a"), warmT("b"))) - m.Tick(ctx) + tick() calls := prober.take() n := 0 for _, c := range calls { @@ -148,6 +161,70 @@ var _ = Describe("Manager probes", func() { Expect(n).To(Equal(1)) }) + It("does not hold other targets' probes behind a slow one", func() { + src.Put(remote("c")) + src.Put(chainCfg("chain2", nil, t("c"), warmT("b"))) + release := make(chan struct{}) + DeferCleanup(func() { close(release); m.waitProbes() }) + prober.block["a"] = release + m.ReportFailure("c", errBoom) + prober.take() + + m.Tick(ctx) // a hangs; c's liveness still runs and starts its recovery + Eventually(func() TargetState { + st, _ := m.ChainStatus("chain2") + return st.Targets[0].State + }).Should(Equal(StateRecovering)) + + clock.Advance(15 * time.Second) + m.Tick(ctx) // a is still in flight: no second probe for it + Eventually(func() []probeCall { return prober.take() }).Should(ContainElement(probeCall{"c", true})) + }) + + It("does not probe a target again while its probe is in flight", func() { + release := make(chan struct{}) + prober.block["a"] = release + m.Tick(ctx) + Eventually(func() []probeCall { + prober.mu.Lock() + defer prober.mu.Unlock() + return slices.Clone(prober.calls) + }).Should(ContainElement(probeCall{"a", false})) + for range 3 { + clock.Advance(15 * time.Second) + m.Tick(ctx) + } + close(release) + m.waitProbes() + n := 0 + for _, c := range prober.take() { + if c.target == "a" { + n++ + } + } + Expect(n).To(Equal(1)) + }) + + It("restores a warm target that is not loaded like a cold one, after min_dwell", func() { + m.ReportFailure("b", errBoom) + prober.fail["b"] = ErrNotLoaded + prober.take() + clock.Advance(15 * time.Second) + prober.fail["b"] = nil + tick() // liveness passes: recovering + st, _ := m.ChainStatus("chain") + Expect(st.Targets[1].State).To(Equal(StateRecovering)) + prober.fail["b"] = ErrNotLoaded + clock.Advance(15 * time.Second) + tick() // nothing to confirm against yet, and no trip + st, _ = m.ChainStatus("chain") + Expect(st.Targets[1].State).To(Equal(StateRecovering)) + clock.Advance(31 * time.Second) + tick() + st, _ = m.ChainStatus("chain") + Expect(st.Targets[1].State).To(Equal(StateHealthy)) + }) + It("closes subscriptions when Run stops", func() { events, _ := m.Subscribe(1) rctx, cancel := context.WithCancel(ctx) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index b3da3f7d9..b2ea212b3 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -87,7 +87,7 @@ targets were down. | Target | Regular check | Check before moving back | |---|---|---| | Remote (`cloud-proxy`) | `GET /v1/models` on the upstream lists the model | one small real request, for example a 1-token completion | -| Local, `warm: true` | the backend answers a health check | one small real request | +| Local, `warm: true` | the backend answers a health check. A check never loads the model: while it is not loaded, the check passes and real requests judge it | one small real request. While the model is not loaded, the target is used again after `min_dwell` | | Local, not warm | none: judged only by real requests; it is never loaded only to check it | none: the target is used again after `min_dwell` | A request that succeeds counts as a check, so a busy target is almost never diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 551af9ab7..28d687ca7 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -151,7 +151,7 @@ Chain states: | Target | Liveness (steady state) | Recovery confirmation | |---|---|---| | remote | `GET /v1/models` returns 2xx and lists the upstream model. `` is the scheme and host of `proxy.upstream_url` plus any path prefix before `/v1`. The upstream model is `proxy.upstream_model`, or the target name when it is empty. `/v1/models` works on any OpenAI-compatible upstream, and `/readyz` exists only on LocalAI. | one minimal real request, chosen by usecase | -| local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. | +| local, `warm: true` | gRPC `HealthCheck` on the loaded backend, with the probe timeout. A probe never loads the model: when the backend is not loaded (the warm preload is still loading it, or a crash removed it), liveness passes and the next real request loads and judges it. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. When the backend is not loaded there is nothing to confirm against: the probe neither passes nor trips, and the target returns to `healthy` like a cold one, when `min_dwell` has passed since the trip. | | local, cold | none: a cold target is judged only by real requests; it is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | Minimal requests by usecase: @@ -165,8 +165,10 @@ Minimal requests by usecase: Probe load rules: -- Each chain has one ticker with jitter. A target shared by chains is probed - once. +- One global scheduler ticks every second, without jitter, and starts the + probes that are due. A target shared by chains is probed once. The + scheduler does not wait for a probe: a target whose probe is still running + is skipped, so one slow target does not delay the others. - A successful real request counts as a liveness pass, so a busy target is almost never probed. - Inference probes run only while a target is `recovering`. From 2ecbaab0105981b27c46fdd4c1b4202da09f1c1d Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:26:07 +0000 Subject: [PATCH 25/79] fix(failover): send remote targets their own upstream model A chain request reached a cloud-proxy target with the client's model, the chain name, whenever the target set no upstream_model: passthrough forwards the body's model and translate falls back to it. The upstream answered 404, which neither retries nor trips, while the liveness probe, which checks the target's own name, kept passing. PrepareTarget now sets the upstream model of a remote target to proxy.upstream_model or the target name, the same name the probe uses. The request pipeline and realtime chain stages both call it. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/endpoints/openai/realtime_model.go | 20 ++++++++----- core/http/middleware/failover.go | 1 + core/http/middleware/failover_test.go | 30 +++++++++++++++++++ core/services/failover/prober.go | 11 +++++++ core/services/failover/prober_test.go | 17 +++++++++++ docs/content/features/model-failover.md | 3 ++ ...2026-09-26-model-failover-chains-design.md | 6 ++++ tests/e2e/e2e_failover_test.go | 17 +++++++++++ 8 files changed, 98 insertions(+), 7 deletions(-) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index a90a0e4e5..ee42bde2d 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -1032,6 +1032,14 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // takes the chain's active target, so everything that inspects stage // configs at session start (voice, reasoning, templates) sees a real model. stageChains := map[string]string{} + loadTarget := 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 + } resolveStage := func(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { if cfg == nil || !cfg.IsFailover() { return cfg, nil @@ -1044,7 +1052,7 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model return nil, fmt.Errorf("failover chain %q not found", cfg.Name) } stageChains[stage] = cfg.Name - return cl.LoadResolvedModelConfig(st.Active, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + return loadTarget(st.Active) } cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) @@ -1185,12 +1193,10 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model appConfig: appConfig, evaluator: evaluator, - stageChains: stageChains, - stageTargetConfig: func(name string) (*config.ModelConfig, error) { - return cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) - }, - tuneLLM: tuneLLM, - appTracing: appConfig.EnableTracing, + stageChains: stageChains, + stageTargetConfig: loadTarget, + tuneLLM: tuneLLM, + appTracing: appConfig.EnableTracing, } if routing != nil { wm.routerDeps = routing.Deps diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go index f3744fdd9..853a3f22a 100644 --- a/core/http/middleware/failover.go +++ b/core/http/middleware/failover.go @@ -58,6 +58,7 @@ func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, ch return nil, fmt.Errorf("failover chain %q: target %q is disabled", 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) diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index 81fff5e02..b5e8ba5b7 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -79,6 +79,9 @@ var _ = Describe("failover chains in the request pipeline", func() { 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() @@ -130,6 +133,33 @@ var _ = Describe("failover chains in the request pipeline", func() { 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")) diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index 03108d0c9..c6d412030 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -85,6 +85,17 @@ func UpstreamModel(cfg config.ModelConfig) string { return cfg.Name } +// PrepareTarget readies a copy of a target's config to serve a chain request. +// A remote target gets its upstream model set explicitly: left empty, +// passthrough forwards the client's "model" (the chain name) and translate +// falls back to it, so the upstream would answer 404 for a model the liveness +// probe (which checks UpstreamModel) just found. +func PrepareTarget(cfg *config.ModelConfig) { + if KindOf(*cfg) == KindRemote { + cfg.Proxy.UpstreamModel = UpstreamModel(*cfg) + } +} + func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error { key, err := cfg.Proxy.ResolveAPIKey() if err != nil || key == "" { diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index 72f03887b..b557c4f96 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -179,3 +179,20 @@ var _ = Describe("DefaultProber", func() { } }) }) + +var _ = Describe("PrepareTarget", func() { + It("names the upstream model of a remote target the way the probe does", func() { + plain := config.ModelConfig{Name: "argus-llm", Backend: "cloud-proxy"} + PrepareTarget(&plain) + Expect(plain.Proxy.UpstreamModel).To(Equal("argus-llm")) + + mapped := config.ModelConfig{Name: "argus-llm", Backend: "cloud-proxy"} + mapped.Proxy.UpstreamModel = "big-llm" + PrepareTarget(&mapped) + Expect(mapped.Proxy.UpstreamModel).To(Equal("big-llm")) + + local := config.ModelConfig{Name: "gemma", Backend: "llama-cpp"} + PrepareTarget(&local) + Expect(local.Proxy.UpstreamModel).To(BeEmpty()) + }) +}) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index b2ea212b3..c7fa1039a 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -50,6 +50,9 @@ Rules: - A chain cannot also set `alias` or `backend`. - Responses name the chain as the model. The `X-LocalAI-Served-Model` header names the target that served the request. +- A remote (`cloud-proxy`) target receives its own model name, never the chain + name: `proxy.upstream_model`, or the target name when `upstream_model` is + empty. The health check looks for the same name. ## How the target is chosen diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 28d687ca7..57bb7e5b0 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -100,6 +100,12 @@ The behaviour matches aliases. Responses echo the chain name. Usage and traces record `requested=` and `served=` through the existing `ContextKeyRequestedModel` and `ContextKeyServedModel` keys. +The upstream of a remote target never sees the chain name. A request served +through a chain reaches a remote target with that target's upstream model: +`proxy.upstream_model`, or the target name when it is empty. This holds in +passthrough and translate mode, and it is the same name the liveness probe +looks for in `/v1/models` (one helper derives both). + ## Failover manager New package: `core/services/failover`. The application creates one `Manager` at diff --git a/tests/e2e/e2e_failover_test.go b/tests/e2e/e2e_failover_test.go index 67e1227b4..b8e1632d2 100644 --- a/tests/e2e/e2e_failover_test.go +++ b/tests/e2e/e2e_failover_test.go @@ -106,6 +106,11 @@ var _ = Describe("Failover chains", Label("failover"), func() { Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-2")) Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) Expect(chainActive("chain-remote")).To(Equal("up-2")) + // The targets set no upstream_model: each upstream must get its + // target's name (what the liveness probe checks), not the chain + // name the client sent. up-2 is healthy, so only this request + // posted to it. + Expect(upstreamBodyModel(up2)).To(Equal("up-2")) up1.SetScript(chatReply) Eventually(func() string { return chainActive("chain-remote") }, 30*time.Second, 500*time.Millisecond). @@ -116,10 +121,22 @@ var _ = Describe("Failover chains", Label("failover"), func() { Expect(resp2.StatusCode).To(Equal(200)) Expect(resp2.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-1")) Expect(resp2.Header.Get("X-LocalAI-Failover")).To(BeEmpty()) + Expect(upstreamBodyModel(up1)).To(Equal("up-1")) }) }) }) +// upstreamBodyModel returns the "model" field of the last request body the +// fake upstream recorded. +func upstreamBodyModel(up *fakeOpenAIUpstreamServer) string { + _, _, _, body := up.recorder.snapshot() + var req struct { + Model string `json:"model"` + } + Expect(json.Unmarshal(body, &req)).To(Succeed(), string(body)) + return req.Model +} + // chainActive returns the active target of a chain as the REST status reports // it, or "" when the status cannot be read. func chainActive(chain string) string { From 132bb216a1522773171e529b4ee43c6f3435d371 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:29:30 +0000 Subject: [PATCH 26/79] fix(failover): resolve chains in transcription and sound-only sessions Transcription-only and sound-detection-only realtime sessions passed a chain config straight to the model loader. It has no backend, so the loader fell back to greedy backend auto-detection: slow, and ending in an unhelpful error. Sound-only sessions are a main use of chains. The stage routing of the full pipeline moves into a stageRouter that both realtime model kinds embed. Every stage resolves to the chain's active target at build time and goes through the failover plan per call. The session sends failover events for any model with chain stages, and restarts them when a transcription session.update swaps the model. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/http/endpoints/openai/realtime.go | 21 ++- .../endpoints/openai/realtime_failover.go | 130 ++++++++++++++-- .../openai/realtime_failover_test.go | 111 +++++++++++++- core/http/endpoints/openai/realtime_model.go | 141 +++++++++--------- docs/content/features/model-failover.md | 5 +- ...2026-09-26-model-failover-chains-design.md | 7 +- 6 files changed, 320 insertions(+), 95 deletions(-) diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index fe2c15068..8c8988371 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -32,6 +32,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/turncoord" "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" "github.com/mudler/LocalAI/core/templates" @@ -638,6 +639,7 @@ func runRealtimeSession(application *application.Application, t Transport, model application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), + application.FailoverManager(), ) } else { m, err = newModel( @@ -756,11 +758,10 @@ func runRealtimeSession(application *application.Application, t Transport, model // 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. - if wrapped, ok := m.(*wrappedModel); ok && wrapped.failover != nil && len(wrapped.stageChains) > 0 { - stopFailoverEvents := startFailoverEvents(t, wrapped.failover, wrapped.stageChains) - defer stopFailoverEvents() - } + // 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 @@ -832,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 @@ -854,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 @@ -1151,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() @@ -1174,7 +1181,7 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config return fmt.Errorf("model is not a valid pipeline model: %s", trUpd.Model) } - m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig) + m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig, fm) if err != nil { return err } diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go index d58d52f33..c47c705ce 100644 --- a/core/http/endpoints/openai/realtime_failover.go +++ b/core/http/endpoints/openai/realtime_failover.go @@ -2,41 +2,151 @@ 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" ) -// isChainStage reports whether stage names a failover chain. -func (m *wrappedModel) isChainStage(stage string) bool { - _, ok := m.stageChains[stage] - return ok && m.failover != nil +// 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 ("llm", "tts", ...) 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 (m *wrappedModel) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { - if !m.isChainStage(stage) { +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 := m.stageChains[stage] - return m.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { - cfg, err := m.stageTargetConfig(target) + 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(m.appTracing, chain, target, err) + 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() { diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go index 48737630c..996fa1eb6 100644 --- a/core/http/endpoints/openai/realtime_failover_test.go +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -3,10 +3,15 @@ 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" ) @@ -34,8 +39,8 @@ var _ = Describe("realtime failover", func() { }) chainModel := func() *wrappedModel { - return &wrappedModel{failover: fm, stageChains: map[string]string{"tts": "chain"}, - stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }} + return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{"tts": "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() { @@ -97,3 +102,105 @@ var _ = Describe("realtime failover", func() { 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{"sound_detection": "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{"transcription": "stt-chain"})) + + var tried []string + tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { + tried = append(tried, name) + return nil, errors.New("dial tcp: refused") + } + _, err = tm.Transcribe(context.Background(), "a.wav", "", false, false, "") + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"stt-a", "stt-b"})) + }) + + It("fails before touching a backend when failover is not running", func() { + _, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + _, _, err = newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + Expect(ml.ListLoadedModels()).To(BeEmpty()) + }) + + It("sends no failover events for a session without chains", func() { + m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-a"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + t := &fakeTransport{} + startModelFailoverEvents(t, m)() + Consistently(failoverEvents(t), 200*time.Millisecond).Should(BeEmpty()) + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index ee42bde2d..f94bb510b 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -86,17 +86,10 @@ type wrappedModel struct { routerSessionID string routerUserID string - // failover and stageChains route pipeline stages that name a failover - // chain; stageChains maps a stage ("llm", "tts", ...) to its chain. - // The *Config fields above 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) + stageRouter // tuneLLM applies the pipeline's LLM overrides (reasoning effort, // disable_thinking) to a chain target loaded per call. - tuneLLM func(cfg *config.ModelConfig) - appTracing bool + tuneLLM func(cfg *config.ModelConfig) } // anyToAnyModel represent a model which supports Any-to-Any operations @@ -119,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, "vad", 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, "transcription", 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, "sound_detection", 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) { @@ -157,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, "transcription", 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, "transcription", 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 { @@ -169,12 +198,11 @@ func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { } func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ + return m.warmStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ {Role: "vad", Cfg: m.VADConfig}, {Role: "transcription", Cfg: m.TranscriptionConfig}, {Role: "sound_detection", Cfg: m.SoundDetectionConfig}, }) - return err } func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { @@ -812,32 +840,7 @@ func (m *wrappedModel) Warmup(ctx context.Context) error { if m.ScoreConfig != nil && m.ScoreConfig != m.LLMConfig { stages = append(stages, backend.PreloadStage{Role: "classifier", Cfg: m.ScoreConfig}) } - // 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. - var ( - plain []backend.PreloadStage - wg sync.WaitGroup - mu sync.Mutex - errs []error - ) - for _, s := range stages { - if !m.isChainStage(s.Role) { - plain = append(plain, s) - continue - } - wg.Go(func() { - err := m.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error { - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}}) - return err - }) - mu.Lock() - errs = append(errs, err) - mu.Unlock() - }) - } - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, plain) - wg.Wait() - return errors.Join(append(errs, err)...) + return m.warmStages(ctx, m.modelLoader, m.appConfig, stages) } // wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream @@ -930,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("vad", cfgVAD) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -942,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("transcription", cfgSST) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -952,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("sound_detection", cfgSound) + } if err != nil { return nil, nil, err } @@ -964,6 +977,7 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, cfgSST, nil } @@ -972,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("sound_detection", cfgSound) + } if err != nil { return nil, err } @@ -985,6 +1003,7 @@ func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfi confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, nil } @@ -1031,29 +1050,12 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // 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. - stageChains := map[string]string{} - loadTarget := 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 - } - resolveStage := func(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { - if cfg == nil || !cfg.IsFailover() { - return cfg, nil - } - if routing == nil || routing.Failover == nil { - return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) - } - st, ok := routing.Failover.ChainStatus(cfg.Name) - if !ok { - return nil, fmt.Errorf("failover chain %q not found", cfg.Name) - } - stageChains[stage] = cfg.Name - return loadTarget(st.Active) + 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 { @@ -1193,17 +1195,14 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model appConfig: appConfig, evaluator: evaluator, - stageChains: stageChains, - stageTargetConfig: loadTarget, - tuneLLM: tuneLLM, - appTracing: appConfig.EnableTracing, + stageRouter: sr, + tuneLLM: tuneLLM, } if routing != nil { wm.routerDeps = routing.Deps wm.routerStore = routing.Store wm.routerSessionID = routing.SessionID wm.routerUserID = routing.UserID - wm.failover = routing.Failover } return wm, nil } diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index c7fa1039a..6bc5089e7 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -120,7 +120,8 @@ pipeline: tts: voice-chain ``` -LocalAI resolves the chain for every call of the stage. When a chain switches, +LocalAI resolves the chain for every call of the stage, in full pipelines and +in transcription-only and sound-detection-only sessions. When a chain switches, the session stays open and keeps its conversation. The next turn uses the new target. @@ -134,8 +135,6 @@ it starts (`reason: initial`) and each time a chain switches: Limits: -- Chains are resolved only in full realtime pipelines. A transcription-only or - sound-detection-only session does not resolve chains yet. - After a `session.update` that changes the pipeline, `localai.model.failover` events keep describing the chains from session start. - A chain used as a router candidate, or as the classifier-mode scoring model, diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 57bb7e5b0..326cfb45a 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -278,8 +278,11 @@ Plain HTTP clients can see failover without subscribing to events. ## Request path (realtime) - In `core/http/endpoints/openai/realtime_model.go`, a pipeline stage that names - a chain is resolved **for each call** of `wrappedModel`, not once at session - start. A helper, `mgr.Do(ctx, chain, func(cfg *config.ModelConfig) error)`, + a chain is resolved **for each call**, not once at session start. This holds + for the full pipeline (`wrappedModel`) and for transcription-only and + sound-detection-only sessions (`transcriptOnlyModel`); both embed the same + stage router. A chain config never reaches the model loader: it has no + backend and would start backend auto-detection. A helper, `mgr.Do(ctx, chain, func(cfg *config.ModelConfig) error)`, goes through the plan with the same classification as HTTP. - Streaming stages (`Predict` with a token callback, `TTSStream`, `TranscribeStream`) wrap the callback. A retry is allowed only until the first From 27e5aba60e5569b3a5d740e576e31da0a2033071 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:29:57 +0000 Subject: [PATCH 27/79] fix(failover): do not follow redirects in remote probes Go resends custom headers such as x-api-key when it follows a redirect, also to another host, so a redirecting upstream could receive the target's API key elsewhere. The probe client now treats a redirect as the response, which fails the probe as a non-2xx status. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/prober.go | 7 ++++++- core/services/failover/prober_test.go | 17 +++++++++++++++++ 2 files changed, 23 insertions(+), 1 deletion(-) diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index c6d412030..9a407e8a0 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -37,7 +37,12 @@ type DefaultProber struct { } func NewProber(loaded LoadedFunc) *DefaultProber { - return &DefaultProber{HTTP: &http.Client{}, Loaded: loaded} + return &DefaultProber{HTTP: &http.Client{ + // A redirect is a failed probe, not something to follow: Go resends + // custom headers such as x-api-key to any host, and the target's + // API key must reach only the configured upstream. + CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, + }, Loaded: loaded} } func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index b557c4f96..842934876 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -125,6 +125,23 @@ var _ = Describe("DefaultProber", func() { Expect(up.auth).To(Equal("Bearer sekret")) }) + It("does not follow a redirect, so the API key never leaves the upstream", func() { + GinkgoT().Setenv("FAILOVER_PROBE_KEY", "sekret") + other := newFakeUpstream() + DeferCleanup(other.srv.Close) + other.models = []string{"argus-llm"} + redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, other.srv.URL+r.URL.Path, http.StatusFound) + })) + DeferCleanup(redirect.Close) + c := proxied("argus-llm", "") + c.Proxy.UpstreamURL = redirect.URL + "/v1/chat/completions" + c.Proxy.APIKeyEnv = "FAILOVER_PROBE_KEY" + c.Proxy.Provider = config.ProxyProviderAnthropic + Expect(p.Liveness(ctx, c, KindRemote, false)).To(MatchError(ContainSubstring("302"))) + Expect(other.paths).To(BeEmpty()) + }) + DescribeTable("remote inference hits the usecase endpoint", func(usecase, path string) { Expect(p.Inference(ctx, proxied("m", "", usecase), KindRemote, false)).To(Succeed()) From 8dbe7a0e3f8b50b9a4d68b2b40846cf2e786ac90 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:30:06 +0000 Subject: [PATCH 28/79] chore(e2e): ignore the cloud-proxy backend binary make build-cloud-proxy-backend writes it next to the mock backend, which is ignored; the cloud-proxy binary showed up as untracked. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .gitignore | 2 ++ tests/e2e/mock-backend/.gitignore | 3 ++- 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/.gitignore b/.gitignore index a4c84a0e8..a6d890185 100644 --- a/.gitignore +++ b/.gitignore @@ -48,6 +48,8 @@ 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 release/ diff --git a/tests/e2e/mock-backend/.gitignore b/tests/e2e/mock-backend/.gitignore index 32923bce6..7df9e3a51 100644 --- a/tests/e2e/mock-backend/.gitignore +++ b/tests/e2e/mock-backend/.gitignore @@ -1 +1,2 @@ -mock-backend \ No newline at end of file +mock-backend +cloud-proxy From f7f12f8203f366aa317d7b13328a75bb24d707bf Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:30:53 +0000 Subject: [PATCH 29/79] fix(failover): warn about warm on a remote target, align the spec The spec promised a load-time warning when a chain marks a remote target warm, where the flag does nothing; the loader now logs it. The remote-backend test moves into ModelConfig.IsRemoteProxy so the loader and the failover manager agree on what is remote. The spec now says what ships: a load blocked by pinned warm targets proceeds over the limit after eviction retries, without an error that names them. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/model_config_failover.go | 11 ++++++++ core/config/model_config_loader.go | 25 +++++++++++++++++++ core/config/model_config_loader_test.go | 10 ++++++++ core/services/failover/types.go | 3 +-- docs/content/features/model-failover.md | 3 +++ ...2026-09-26-model-failover-chains-design.md | 4 ++- 6 files changed, 53 insertions(+), 3 deletions(-) diff --git a/core/config/model_config_failover.go b/core/config/model_config_failover.go index e33a51ea8..f13ac2bdd 100644 --- a/core/config/model_config_failover.go +++ b/core/config/model_config_failover.go @@ -49,6 +49,17 @@ const ( // 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 "cloud-proxy", "localai-proxy": + return true + } + return false +} + func (f FailoverConfig) ProbeInterval() time.Duration { return durationOr(f.Probe.Interval, DefaultFailoverProbeInterval) } diff --git a/core/config/model_config_loader.go b/core/config/model_config_loader.go index 363fcc652..a29c95c52 100644 --- a/core/config/model_config_loader.go +++ b/core/config/model_config_loader.go @@ -543,6 +543,28 @@ func validateFailoverTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, 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 @@ -931,6 +953,9 @@ func (bcl *ModelConfigLoader) loadModelConfigsFromPath(path string, strict bool, 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) + } } return nil diff --git a/core/config/model_config_loader_test.go b/core/config/model_config_loader_test.go index 9912a365b..aed79dbb8 100644 --- a/core/config/model_config_loader_test.go +++ b/core/config/model_config_loader_test.go @@ -441,4 +441,14 @@ var _ = Describe("ModelConfigLoader failover validation", 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()) + }) }) diff --git a/core/services/failover/types.go b/core/services/failover/types.go index dc4ac2002..428123c00 100644 --- a/core/services/failover/types.go +++ b/core/services/failover/types.go @@ -87,8 +87,7 @@ type ChainStatus struct { // KindOf decides how a target is probed: proxy backends forward to another // server and are checked over HTTP, everything else runs in this instance. func KindOf(cfg config.ModelConfig) Kind { - switch cfg.Backend { - case "cloud-proxy", "localai-proxy": + if cfg.IsRemoteProxy() { return KindRemote } return KindLocal diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index 6bc5089e7..5915ff8b3 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -107,6 +107,9 @@ count toward the active backend limit (`--max-active-backends`) like any pinned model: LocalAI never evicts them to make room, and if they fill the limit, a new model still loads rather than being blocked. +`warm` applies only to local targets. On a remote (`cloud-proxy`) target it +has no effect, and LocalAI logs a warning when it loads the chain. + ## Realtime pipelines A pipeline stage can name a chain: diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index 326cfb45a..0765ddb97 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -190,7 +190,9 @@ this. The manager loads `warm: true` targets at startup and marks them pinned in the watchdog, so LRU and idle eviction skip them. They still count toward the active backend limit. When pinned warm targets leave no room for another load, -that load fails with an error that names them. The docs state this. +the loader never evicts them: it retries eviction and then loads the model +anyway, over the limit, with no error that names the warm targets. The docs +state this. ### Events From d9d1f4f2ea4e23bb90255bb1fca9dce9aae4dea6 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 19:41:22 +0000 Subject: [PATCH 30/79] chore: drop a stray task report from the branch Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .../task-12-report.md | 137 ------------------ 1 file changed, 137 deletions(-) delete mode 100644 .superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md diff --git a/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md b/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md deleted file mode 100644 index ae2d8cb98..000000000 --- a/.superpowers/sdd/2026-09-26-model-failover-chains/task-12-report.md +++ /dev/null @@ -1,137 +0,0 @@ -# Task 12 report: MCP admin tools for failover chains - -## What I implemented - -Three MCP admin tools mirroring the Task 9 REST endpoints: - -- `list_failover_chains` (read-only) → `GET /api/failover` -- `pin_failover_target` (mutating) → `POST /api/failover/:chain/pin` -- `unpin_failover_target` (mutating) → `DELETE /api/failover/:chain/pin` - -Files changed, by layer: - -- **Tool identity**: `pkg/mcp/localaitools/tools.go` — added `ToolListFailoverChains` (read-only block), `ToolPinFailoverTarget`/`ToolUnpinFailoverTarget` (mutating block + `mutatingToolNames`). -- **DTOs**: `pkg/mcp/localaitools/dto.go` — added `FailoverTargetInfo` and `FailoverChainInfo` (LLM-facing subset of `failover.TargetStatus`/`failover.ChainStatus`). -- **Interface**: `pkg/mcp/localaitools/client.go` — added `ListFailoverChains`, `PinFailoverTarget`, `UnpinFailoverTarget` to `LocalAIClient`. -- **In-process impl**: `pkg/mcp/localaitools/inproc/client.go` — added `Failover *failover.Manager` field and the three methods (nil-safe: `ListFailoverChains` returns `[]`, Pin/Unpin return `"failover is not running"`). -- **HTTP impl**: `pkg/mcp/localaitools/httpapi/client.go` + `httpapi/routes.go` — added `routeFailover = "/api/failover"` and the three methods, using `c.do`. -- **Tool registration**: new file `pkg/mcp/localaitools/tools_failover.go` (`registerFailoverTools`), wired into `pkg/mcp/localaitools/server.go`. -- **Prompts**: `prompts/20_tools.md` (read-only + mutating one-liners) and `prompts/10_safety.md` (both mutating names added to the confirmation-rule list). -- **Wiring the manager into the in-process client**: `core/application/application.go` and `core/application/startup.go` — see "Adaptation: startup ordering" below; this was the one place I had to deviate from a literal reading of the brief. -- **Tests**: `coverage_test.go`, `server_test.go`, `fakes_test.go`, `inproc/client_test.go`, `httpapi/client_test.go` — all as specified, plus `core/http/endpoints/mcp/localai_assistant_test.go` (see adaptations). - -`parity_test.go` was left untouched — the brief didn't specify a new parity spec for failover, and the file's existing specs are all hand-picked equality checks for specific methods (ListGalleries, GallerySearch, ImportModelURI, SystemInfo); it isn't a generic "every method" loop, so nothing there needed updating for the build to stay green. - -## Adaptations from the brief (read the real code, deviated where it disagreed) - -1. **`jsonResult`/`errorResult` return one value, not three.** The brief's `tools_failover.go` snippet writes `return jsonResult(chains)` as if it were the handler's whole 3-tuple return. The real helpers in `pkg/mcp/localaitools/errors.go` are: - ```go - func errorResult(err error) *mcp.CallToolResult - func jsonResult(v any) *mcp.CallToolResult - ``` - Matching `tools_aliases.go`'s actual pattern, every handler returns `jsonResult(x), nil, nil` / `errorResult(err), nil, nil` — three explicit values. I used that form throughout `tools_failover.go`. - -2. **Mutating tool descriptions reference safety rule 1.** `.agents/localai-assistant-mcp.md`'s checklist says "Mutating tools must reference safety rule 1 in the description," and `tools_aliases.go`'s `set_alias` does this ("Requires user confirmation per safety rule 1."). The brief's descriptions for `pin_failover_target`/`unpin_failover_target` didn't include this phrase, so I added it to match the established convention and the checklist. - -3. **Startup ordering: `assistantClient.Failover` can't be set where the brief implies.** The brief says to set the field "where the client is constructed (from `application.FailoverManager()`)." I found that construction site (`core/application/application.go`'s `start()`, called from `core/application/startup.go`'s `New()` at line 74) — but `application.failoverManager` is only built later in the same `New()` function, at line 259, **after** `start()` (and therefore the assistant-client construction) has already returned. Calling `a.FailoverManager()` inside `start()` would have captured a permanent `nil`. - - Fix: added an `assistantClient *localaiInproc.Client` field to `Application` (set in `start()` when the assistant client is built), then in `startup.go`, right after `application.failoverManager = failover.New(...)`, added: - ```go - if application.assistantClient != nil { - application.assistantClient.Failover = application.failoverManager - } - ``` - This is safe because `assistantClient` is a pointer already captured by value inside the `LocalAIClient` interface passed to `holder.Initialize()` — mutating a field on it after the fact is visible through the interface. Verified with `go test ./core/application/...` and `go test ./core/http/endpoints/mcp/...`. - -4. **`stubClient` in `core/http/endpoints/mcp/localai_assistant_test.go`** implements `localaitools.LocalAIClient` for that package's own tests and isn't in the brief's file list, but the interface change broke its build. Added the three stub methods (empty list, nil errors) to keep it compiling — same pattern as its existing stubs for `GetRouterCorpusStats` etc. - -5. **inproc failover test fixture**: the brief says to "build a `failover.Manager` over a small in-memory source with one chain." `failover.ConfigSource` (`GetModelConfig`/`GetAllModelsConfigs`) is exported, but the concrete fake used by `core/services/failover`'s own tests (`fakeSource`, `chainCfg`, `t()`) is unexported and package-local, so I wrote a minimal `fakeFailoverSource` directly in `inproc/client_test.go` implementing the same two-method interface over a `map[string]config.ModelConfig`, seeded with a `chain` config (`config.FailoverConfig{Targets: [...]}`) plus two local targets `a`/`b`. - -## TDD evidence - -**RED** — after writing all test-file changes (`coverage_test.go`, `server_test.go`, `fakes_test.go`, plus new specs in `inproc/client_test.go` and `httpapi/client_test.go` were added later, see below), I temporarily reverted the implementation files (`git stash push` on `tools.go`, `dto.go`, `client.go`, `server.go`, `inproc/client.go`, `httpapi/client.go`, `httpapi/routes.go`, prompts, and the two `core/application` files; moved `tools_failover.go` out of the package) and ran: - -``` -$ go test ./pkg/mcp/localaitools/... 2>&1 | tail -30 -# github.com/mudler/LocalAI/pkg/mcp/localaitools [github.com/mudler/LocalAI/pkg/mcp/localaitools.test] -pkg/mcp/localaitools/fakes_test.go:62:32: undefined: FailoverChainInfo -pkg/mcp/localaitools/fakes_test.go:400:63: undefined: FailoverChainInfo -pkg/mcp/localaitools/fakes_test.go:405:11: undefined: FailoverChainInfo -pkg/mcp/localaitools/coverage_test.go:49:2: undefined: ToolListFailoverChains -pkg/mcp/localaitools/coverage_test.go:71:2: undefined: ToolPinFailoverTarget -pkg/mcp/localaitools/coverage_test.go:72:2: undefined: ToolUnpinFailoverTarget -pkg/mcp/localaitools/server_test.go:95:2: undefined: ToolListFailoverChains -pkg/mcp/localaitools/server_test.go:159:4: undefined: ToolListFailoverChains -pkg/mcp/localaitools/server_test.go:160:4: undefined: ToolPinFailoverTarget -pkg/mcp/localaitools/server_test.go:161:4: undefined: ToolUnpinFailoverTarget -pkg/mcp/localaitools/server_test.go:161:4: too many errors -FAIL github.com/mudler/LocalAI/pkg/mcp/localaitools [build failed] -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.035s -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.167s -``` -This matches the brief's Step 1 expectation ("Expected: compile failure"). I then `git stash pop` and restored `tools_failover.go` to get back to the implemented state. - -**GREEN** — after restoring the implementation and adding the remaining `inproc/client_test.go` / `httpapi/client_test.go` specs: - -``` -$ go test ./pkg/mcp/localaitools/... -count=1 -v 2>&1 | grep -E "SUCCESS|FAIL|ok " -SUCCESS! -- 53 Passed | 0 Failed | 0 Pending | 0 Skipped -ok github.com/mudler/LocalAI/pkg/mcp/localaitools 0.133s -SUCCESS! -- 24 Passed | 0 Failed | 0 Pending | 0 Skipped -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.037s -SUCCESS! -- 16 Passed | 0 Failed | 0 Pending | 0 Skipped -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.171s -``` - -Additional verification (build scope per implementer-rules, plus the two packages touched indirectly by the interface change): - -``` -$ go build ./core/... ./pkg/mcp/... ./tests/... -(clean, no output) - -$ go test ./pkg/mcp/localaitools/... ./core/application/... ./core/http/endpoints/mcp/... -count=1 -ok github.com/mudler/LocalAI/pkg/mcp/localaitools 0.155s -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/httpapi 0.045s -ok github.com/mudler/LocalAI/pkg/mcp/localaitools/inproc 0.176s -ok github.com/mudler/LocalAI/core/application 0.265s -ok github.com/mudler/LocalAI/core/http/endpoints/mcp 0.237s - -$ gofmt -l -(empty — clean) - -$ go vet ./core/application/... ./pkg/mcp/... -(clean, no output) -``` - -## Files changed - -- `pkg/mcp/localaitools/tools.go` -- `pkg/mcp/localaitools/dto.go` -- `pkg/mcp/localaitools/client.go` -- `pkg/mcp/localaitools/server.go` -- `pkg/mcp/localaitools/tools_failover.go` (new) -- `pkg/mcp/localaitools/inproc/client.go` -- `pkg/mcp/localaitools/httpapi/client.go` -- `pkg/mcp/localaitools/httpapi/routes.go` -- `pkg/mcp/localaitools/prompts/20_tools.md` -- `pkg/mcp/localaitools/prompts/10_safety.md` -- `pkg/mcp/localaitools/coverage_test.go` -- `pkg/mcp/localaitools/server_test.go` -- `pkg/mcp/localaitools/fakes_test.go` -- `pkg/mcp/localaitools/inproc/client_test.go` -- `pkg/mcp/localaitools/httpapi/client_test.go` -- `core/application/application.go` -- `core/application/startup.go` -- `core/http/endpoints/mcp/localai_assistant_test.go` - -## Self-review - -- Completeness: all three tools registered, gated correctly (`ToolListFailoverChains` in the read-only catalog; both mutating tools skipped when `Options.DisableMutating`), both client implementations covered, prompts updated, safety-rule coverage test (`TestPromptsContainSafetyAnchors`'s "names every mutating tool" spec) passes automatically since it reads `mutatingToolNames`. -- Quality/YAGNI: DTOs intentionally drop `ConsecutiveOK`/`LastProbe`/`ActiveSince` — internal probe bookkeeping the LLM has no use for when deciding to pin/unpin; documented why in the doc comment. -- Nil-safety: both inproc failover methods and the httpapi Pin/Unpin exercise the "no failover configured" path in tests (inproc has explicit specs for it; httpapi's behavior when unconfigured is identical to any other client error — REST returns 404, `c.do` surfaces `*HTTPError`, no new code path needed there). -- Existing patterns followed: constant grouping/comments, `errorResult`/`jsonResult` triple-return, `c.do` signature, `url.PathEscape` on the chain path segment, fake-client recording pattern, Ginkgo `Describe`/`It` structure matching the alias specs. -- Pristine output: `gofmt -l` and `go vet` clean across all touched files. - -## Concerns - -- None blocking. The startup-ordering fix (adaptation 3) is the only piece that goes beyond a single-file, mechanical change — it touches two `core/application` files instead of the "one field set inline" the brief describes. I verified it with both `core/application` and `core/http/endpoints/mcp` package tests, and confirmed via read of `startup.go` that `New()` is the sole caller of `start()` and that `failoverManager` is not read anywhere between `start()` and its own assignment, so there's no other place relying on it being nil momentarily. From e62854c3403f9b39b2d7f02858ec16fbdfe5b133 Mon Sep 17 00:00:00 2001 From: localai-org-maint-bot <306269227+localai-org-maint-bot@users.noreply.github.com> Date: Sat, 26 Sep 2026 22:06:23 +0000 Subject: [PATCH 31/79] fix(failover): pass credential lookup from CLI Pass the API key environment lookup through ApplicationConfig to satisfy core configuration lint. Keep credential resolution dynamic and exclude the callback from serialization. Handle the five close results reported by errcheck. Assisted-by: Codex:gpt-6 golangci-lint Signed-off-by: Ettore Di Giacinto --- core/application/startup.go | 2 +- core/cli/run.go | 1 + core/config/application_config.go | 7 +++++++ core/config/model_config.go | 7 +++++-- core/config/model_config_failover_test.go | 21 +++++++++++++------- core/http/endpoints/localai/failover_test.go | 2 +- core/services/failover/prober.go | 15 +++++++------- core/services/failover/prober_test.go | 9 +++++---- tests/e2e/realtime_ws_test.go | 4 ++-- 9 files changed, 44 insertions(+), 24 deletions(-) diff --git a/core/application/startup.go b/core/application/startup.go index 4c1ddd4e9..db9e1caaf 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -255,7 +255,7 @@ func New(opts ...config.AppOption) (*Application, error) { // 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()))), + failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()), options.ProxyAPIKeyEnvLookup)), failover.WithOnWarmChanged(application.applyFailoverWarmTargets), ) // The assistant client was built in start() (above), before this diff --git a/core/cli/run.go b/core/cli/run.go index 3dc165ecc..de6621ba9 100644 --- a/core/cli/run.go +++ b/core/cli/run.go @@ -291,6 +291,7 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error { } opts := []config.AppOption{ + config.WithProxyAPIKeyEnvLookup(os.Getenv), config.WithContext(context.Background()), config.WithArtifactDownloadConcurrency(r.ArtifactDownloadConcurrency), config.WithModelArtifactMaterializer(modelartifacts.NewDefaultManager( diff --git a/core/config/application_config.go b/core/config/application_config.go index b83e8fe58..24e1867c4 100644 --- a/core/config/application_config.go +++ b/core/config/application_config.go @@ -14,6 +14,9 @@ import ( ) type ApplicationConfig struct { + // ProxyAPIKeyEnvLookup resolves upstream credentials at the CLI boundary. + ProxyAPIKeyEnvLookup func(string) string `json:"-" yaml:"-"` + Context context.Context ConfigFile string SystemState *system.SystemState @@ -271,6 +274,10 @@ type AgentPoolConfig struct { type AppOption func(*ApplicationConfig) +func WithProxyAPIKeyEnvLookup(lookup func(string) string) AppOption { + return func(o *ApplicationConfig) { o.ProxyAPIKeyEnvLookup = lookup } +} + func NewApplicationConfig(o ...AppOption) *ApplicationConfig { opt := &ApplicationConfig{ Context: context.Background(), diff --git a/core/config/model_config.go b/core/config/model_config.go index 9c38f178d..7b67acb2a 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -306,10 +306,13 @@ const ( // "" 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() (string, error) { +func (p ProxyConfig) ResolveAPIKey(envLookup func(string) string) (string, error) { switch { case p.APIKeyEnv != "": - v := os.Getenv(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) } diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go index 9e82868e9..6a276f056 100644 --- a/core/config/model_config_failover_test.go +++ b/core/config/model_config_failover_test.go @@ -70,25 +70,32 @@ failover: }) var _ = Describe("ProxyConfig.ResolveAPIKey", func() { - It("reads the env var", func() { - GinkgoT().Setenv("FAILOVER_TEST_KEY", "k1") - Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey()).To(Equal("k1")) + 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() + _, 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() + _, 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()).To(Equal("k2")) + Expect(ProxyConfig{APIKeyFile: f}.ResolveAPIKey(os.Getenv)).To(Equal("k2")) }) It("returns empty when nothing is set", func() { - Expect(ProxyConfig{}.ResolveAPIKey()).To(Equal("")) + Expect(ProxyConfig{}.ResolveAPIKey(os.Getenv)).To(Equal("")) }) }) diff --git a/core/http/endpoints/localai/failover_test.go b/core/http/endpoints/localai/failover_test.go index fd3b6a733..77fb48678 100644 --- a/core/http/endpoints/localai/failover_test.go +++ b/core/http/endpoints/localai/failover_test.go @@ -91,7 +91,7 @@ var _ = Describe("failover endpoints", func() { req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil) resp, err := http.DefaultClient.Do(req) Expect(err).ToNot(HaveOccurred()) - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream")) r := bufio.NewReader(resp.Body) next := func() string { diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index 9a407e8a0..6b27e4f8a 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -32,17 +32,18 @@ type LoadedFunc func(cfg config.ModelConfig) grpc.Backend // DefaultProber probes remote targets over the upstream's OpenAI-compatible // API and local targets through their gRPC backend. type DefaultProber struct { - HTTP *http.Client - Loaded LoadedFunc + HTTP *http.Client + Loaded LoadedFunc + EnvLookup func(string) string } -func NewProber(loaded LoadedFunc) *DefaultProber { +func NewProber(loaded LoadedFunc, envLookup func(string) string) *DefaultProber { return &DefaultProber{HTTP: &http.Client{ // A redirect is a failed probe, not something to follow: Go resends // custom headers such as x-api-key to any host, and the target's // API key must reach only the configured upstream. CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, - }, Loaded: loaded} + }, Loaded: loaded, EnvLookup: envLookup} } func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { @@ -102,7 +103,7 @@ func PrepareTarget(cfg *config.ModelConfig) { } func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error { - key, err := cfg.Proxy.ResolveAPIKey() + key, err := cfg.Proxy.ResolveAPIKey(p.EnvLookup) if err != nil || key == "" { return err } @@ -124,7 +125,7 @@ func (p *DefaultProber) do(req *http.Request, cfg config.ModelConfig) (*http.Res return nil, err } if resp.StatusCode/100 != 2 { - resp.Body.Close() + _ = resp.Body.Close() return nil, fmt.Errorf("upstream %s: HTTP %d", req.URL.Path, resp.StatusCode) } return resp, nil @@ -143,7 +144,7 @@ func (p *DefaultProber) remoteLiveness(ctx context.Context, cfg config.ModelConf if err != nil { return err } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() var list struct { Data []struct { ID string `json:"id"` diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index 842934876..fd22bd1e3 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "net/http/httptest" + "os" "sync" "github.com/mudler/LocalAI/core/config" @@ -75,7 +76,7 @@ var _ = Describe("DefaultProber", func() { BeforeEach(func() { up = newFakeUpstream() DeferCleanup(up.srv.Close) - p = NewProber(nil) + p = NewProber(nil, os.Getenv) }) proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { @@ -155,7 +156,7 @@ var _ = Describe("DefaultProber", func() { It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { b := &fakeBackend{healthy: true} - p = NewProber(func(config.ModelConfig) grpc.Backend { return b }) + p = NewProber(func(config.ModelConfig) grpc.Backend { return b }, nil) c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) @@ -171,7 +172,7 @@ var _ = Describe("DefaultProber", func() { // The warm preload loads the model; a probe that loaded it too would // block until the load finished and then judge it on an expired ctx. asked := 0 - p = NewProber(func(config.ModelConfig) grpc.Backend { asked++; return nil }) + p = NewProber(func(config.ModelConfig) grpc.Backend { asked++; return nil }, nil) for _, uc := range []string{"chat", "tts"} { c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{uc}} c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) @@ -185,7 +186,7 @@ var _ = Describe("DefaultProber", func() { p = NewProber(func(config.ModelConfig) grpc.Backend { Fail("cold liveness must not look up the backend") return nil - }) + }, nil) // None of these files exist: a missing file says nothing about whether // the target can serve (download on first use, dotted names, backends // that need no file). Only a real request may trip a cold target. diff --git a/tests/e2e/realtime_ws_test.go b/tests/e2e/realtime_ws_test.go index 1aa2552a1..68181ef35 100644 --- a/tests/e2e/realtime_ws_test.go +++ b/tests/e2e/realtime_ws_test.go @@ -237,7 +237,7 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { It("switches the LLM mid-session and keeps the conversation", func() { conn := connectWS("rt-failover") - defer conn.Close() + defer func() { _ = conn.Close() }() Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) initial := drainUntil(conn, "localai.model.failover", 10*time.Second) @@ -275,7 +275,7 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { It("starts the session on the next target when the active one fails to warm up", func() { conn := connectWS("rt-failover-warm") - defer conn.Close() + defer func() { _ = conn.Close() }() Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) initial := drainUntil(conn, "localai.model.failover", 10*time.Second) From b86902e78e78854acb5d60942b7cd40c2e571ccd Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 20:02:02 +0000 Subject: [PATCH 32/79] docs: design distributed failover, localai-proxy and the chain UI Failover state is per frontend, remote chain targets only cover chat, and chains have no UI. Design shared pins, health and chain state for distributed mode, a localai-proxy backend for every API including live transcription, a chain editor and health view, and a contributor rule for distributed-aware state. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- ...26-failover-distributed-proxy-ui-design.md | 335 ++++++++++++++++++ 1 file changed, 335 insertions(+) create mode 100644 docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md diff --git a/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md b/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md new file mode 100644 index 000000000..e380a2888 --- /dev/null +++ b/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md @@ -0,0 +1,335 @@ +# Failover chains: distributed mode, localai-proxy backend and WebUI + +Date: 2026-09-26 +Status: design approved in brainstorming, pending spec review +Builds on: `2026-09-26-model-failover-chains-design.md` (same PR) + +## Problem + +The failover chains in this PR work on one LocalAI instance. Three gaps +remain: + +1. **Distributed mode.** With several frontends, chain definitions converge + (config edits broadcast `cache.invalidate.models`), but runtime state does + not. Each frontend probes, trips and pins on its own. A pin applies only + on the frontend that received it and is lost on restart. `/api/failover` + and its event stream show a different view on each frontend. `warm: true` + pins only a frontend stub, so workers can evict the model, and every + frontend preloads it. +2. **Remote targets cover only chat.** `cloud-proxy` forwards chat and + completions. A remote LocalAI cannot serve transcription, TTS, VAD, sound + detection or the other modalities as a chain target, so a realtime + pipeline cannot fail over per stage between a remote and a local LocalAI. +3. **No UI.** Chains can be edited only as raw JSON in the model editor, and + their health is visible only through the API. + +LocalAI also has no rule that makes a feature state how it behaves with +several frontends. The failover feature shipped with per-instance state +because nothing asked the question. + +## Goals + +- **C. Distributed-aware failover.** Pins, target health and chain state are + the same on every frontend. One frontend probes. Warm targets stay loaded on + workers and are preloaded once. +- **A. `localai-proxy` backend.** A gRPC backend that serves every backend + method with a REST counterpart by calling an upstream LocalAI, including + live transcription through the upstream's realtime API. +- **B. WebUI.** A chain editor field, a chain template, a live health strip + per chain, a "chain" badge in the model list, and a Failover overview page. +- **D. Contributor rule.** `AGENTS.md` and a new `.agents/distributed-state.md` + require every stateful feature to choose and document a distributed mode. + +## Non-goals + +- Sharing failover state between instances that are not in one distributed + cluster. +- A bridge for backend methods that have no REST or realtime counterpart + upstream (audio encode/decode, metrics, status, fine-tune, quantization). +- Chains as router candidates or as the realtime classifier model. + +--- + +## C. Distributed-aware failover + +Standalone mode (no NATS, no PostgreSQL) keeps today's behaviour. Everything +below applies when distributed mode is on. + +### Shared state + +| State | Writers | Mechanism | Survives restart | +|---|---|---|---| +| Pins | any frontend (REST, MCP) | `syncstate.SyncedMap` named `failover.pins`, key = chain, with a gorm `Store` | yes | +| Target health: state, last error, since, consecutive passes | any frontend on a local transition, and the probe leader | `syncstate.SyncedMap` named `failover.targets`, key = target, NATS only, `Reconcile` on | no | +| Chain state: active target, active since, chain state | the probe leader only | `syncstate.SyncedMap` named `failover.chains`, key = chain, NATS only, `Reconcile` on | no | + +- The pins table is created under `advisorylock.KeySchemaMigrate`, the same + way the jobs store creates its tables. +- The manager gets a small `StateSync` dependency. The standalone + implementation is a no-op; the distributed implementation wraps the three + maps. The manager does not import NATS or gorm directly. +- A peer delta is applied through `OnApply`, which changes local state + without publishing again (no echo loops). + +### Who does what + +- **Every frontend** plans requests from the shared state. `Plan` already + leaves unhealthy targets out of the attempt order, so a target tripped on + another frontend is skipped at once. In-request retry stays local. +- **Any frontend** that sees a real request trip or pass a target publishes + the new target state. +- **The probe leader** runs probes, recovery confirmation, dwell-based + fail-back, the chain recompute and the warm preload. It publishes chain + state. Leadership uses `advisorylock.RunLeaderLoop` with a new key + `failover-prober` and the same 1 s interval as the scheduler. +- **Followers** do not recompute the active target. They adopt the leader's + chain state. When no chain state has arrived yet (start-up), a follower + uses its own recompute until the first delta. +- If the leader stops, another frontend takes the lock on its next tick. + Pins and target health are not affected. Probes and fail-back pause for at + most one tick. + +### Events + +Each frontend emits `chain.switched` and `target.state` to its own +subscribers (SSE, realtime `localai.model.failover`) when it applies a +change, whether the change is local or from a peer. Every frontend's stream +therefore shows the same events. + +### Warm targets + +- The SmartRouter's and ReplicaReconciler's pinned-model resolver includes + `WarmTargets()`, so workers never evict a warm target. +- Only the leader preloads warm targets. +- A frontend cannot tell from its local model store whether a worker has a + model loaded. In distributed mode, warm-target liveness asks the node + registry whether a healthy node serves the model. When none does, the probe + is inconclusive: it neither passes nor trips, and real requests decide. + +### Tests + +- Unit: two managers on the test fakebus (`core/services/testutil`). A pin on + one shows on the other. A trip on one is skipped by the other's plan. Only + the lock holder probes. A new leader resumes fail-back. +- A spec in `tests/e2e/distributed` when its harness supports two frontends + cheaply; otherwise the unit specs are the coverage and the PR says so. + +### Docs + +`model-failover.md` replaces the "state is per instance" limit with a +"Distributed mode" section that describes the table above. + +--- + +## A. `localai-proxy` backend + +### Shape + +- `backend/go/localai-proxy` is a separate OCI gallery backend, like + `cloud-proxy`. It is registered in the `Makefile`, `backend/index.yaml` and + `.github/backend-matrix.yml` (Linux amd64/arm64 and Darwin Metal), following + `.agents/adding-backends.md`. +- It reuses cloud-proxy's auth header, HTTP client (no redirects) and + hop-by-hop header helpers. It has no translate mode: the upstream is always + LocalAI. +- `Load` refuses a model without proxy options, so greedy backend probing + never selects it. + +### Config + +```yaml +name: argus-whisper +backend: localai-proxy +known_usecases: [transcript] +proxy: + upstream_url: https://argus:8080 # base URL; each method appends its path + upstream_model: whisper-large # optional; default: this model's name + api_key_env: ARGUS_KEY + request_timeout_seconds: 60 # applies to non-streaming calls +``` + +- `core/backend/options.go` passes `ProxyOptions` to `localai-proxy` as well + as `cloud-proxy`. +- `proxy.mode` and `proxy.provider` are ignored, with a load warning. +- A `localai-proxy` model without `known_usecases` loads with a warning, because + usecases decide default-model selection and the failover inference probe. +- The failover prober already treats `localai-proxy` as remote. + `UpstreamBase` accepts a base URL unchanged. + +### Method mapping + +| Backend method | Upstream endpoint | +|---|---| +| Predict, PredictStream | `/v1/chat/completions`, `/v1/completions` (SSE when streaming) | +| Embedding | `/v1/embeddings` | +| Rerank | `/v1/rerank` | +| TokenizeString, Detokenize | `/v1/tokenize`, `/v1/detokenize` | +| Score | `/api/score` | +| GenerateImage, UpscaleImage | `/v1/images/generations`, `/v1/images/upscale` | +| GenerateVideo | `/video` | +| Generate3D, Animate3D | `/3d/generations`, `/3d/animate` | +| TTS, TTSStream | `/tts` (streaming passes the upstream WAV header and PCM through) | +| SoundGeneration | `/v1/sound-generation` | +| AudioTranscription, AudioTranscriptionStream | `/v1/audio/transcriptions` (`stream=true` for SSE deltas) | +| AudioTranscriptionLive | upstream `/v1/realtime` transcription session (see below) | +| Diarize | `/v1/audio/diarization` | +| VAD | `/v1/vad` | +| SoundDetection | `/v1/audio/classification` | +| Detect, Depth | `/v1/detection`, `/v1/depth` | +| FaceVerify, FaceAnalyze | `/v1/face/verify`, `/v1/face/analyze` | +| VoiceVerify, VoiceAnalyze, VoiceEmbed | `/v1/voice/verify`, `/v1/voice/analyze`, `/v1/voice/embed` | +| Stores* | `/stores/set`, `/stores/get`, `/stores/delete`, `/stores/find` | +| AudioTransform | `/audio/transformations` | + +Every request uses the upstream model name (`proxy.upstream_model`, else the +model name), the same derivation as `failover.UpstreamModel`. + +Methods with no counterpart (AudioEncode, AudioDecode, AudioToAudioStream, +TokenClassify, GetMetrics, Status, ModelMetadata, fine-tune and quantization) +return gRPC `Unimplemented` with the message +`localai-proxy: has no upstream counterpart`. + +### Files + +Core passes some inputs and outputs as local paths: + +- Inputs (transcription and diarization audio, sound detection `src`, image + `src` and reference images): the proxy reads the file and uploads it as + multipart or base64, as the endpoint expects. +- Outputs (TTS, image, sound generation `dst`): the proxy writes the upstream + result (bytes, or a download of the returned URL, or decoded base64) to + `dst`. + +### Live transcription bridge + +`AudioTranscriptionLive` opens a WebSocket to the upstream `/v1/realtime` +transcription session: + +1. On the first `TranscriptLiveConfig`, send the session update with the + upstream model, language and sample rate, then answer `ready`. +2. Forward each `TranscriptLiveAudio` as `input_audio_buffer.append` + (PCM float to PCM16 base64 at the session rate). +3. Map `conversation.item.input_audio_transcription.delta` to `delta`, and + `...completed` to `delta` (any remaining text) plus `eou: true`. +4. When the gRPC send side closes, commit the buffer, wait for the final + completion, send `final_result` and close. +5. An upstream error or disconnect ends the gRPC stream with `Unavailable`. + Word timings and `eob` are not available from the upstream and stay empty. + +A realtime stage whose live session fails reopens on the next chain target at +the next utterance (behaviour from the base spec). + +### Core changes + +- **Rerank for Go backends.** `pkg/grpc` gets an optional rerank interface and + a server handler, in the same way as `Score`. +- **`Unimplemented` is a capability gap.** The failover retry path (HTTP and + `Manager.Do`) treats gRPC `Unimplemented` like an admission rejection: skip + to the next target for this request, and do not trip the target. Otherwise a + chain of a remote and a local target fails a request the local target can + serve. + +### Tests + +- Unit: a fake LocalAI `httptest` upstream per method family (request path, + body, model name, auth header, file upload and `dst` write). +- Unit: a fake WebSocket upstream for the live bridge (ready, deltas, eou, + final result, upstream disconnect). +- E2E: `localai-proxy` models that point back at the test server's own mock + models. A realtime pipeline whose stages are chains of a `localai-proxy` + target and a local target completes a turn, and switches stage when the + proxy target fails. + +### Docs + +A `localai-proxy` section in `docs/content/features/backends.md` (or the page +that documents `cloud-proxy`), and a remote-LocalAI example on +`model-failover.md`. + +--- + +## B. WebUI + +### Model editor + +- A `failover-targets` field component replaces the JSON editor for + `failover.targets` (`core/config/meta/registry.go` switches the component + name). Each row has a model picker (`SearchableModelSelect`), move up/down, + remove and a **warm** toggle. The toggle is disabled with a tooltip on remote + targets. Inline validation: at least 2 targets, no duplicates, no chain as a + target. +- The probe, trip and recovery fields stay in the Advanced group. +- A **Failover chain** template in `modelTemplates.js`, seeded with two empty + targets. `?template=failover` preselects it. + +### Health strip + +When the edited model is a chain, a `FailoverChainStatus` component above the +form shows: + +- the chain state (primary / fallback / degraded) and the active target, with + the time since it became active; +- for each target: state, kind, warm, last probe and last error; +- **Pin** and **Unpin** for admins, behind a confirm dialog. The pinned target + is marked. + +### Model list and overview + +- Installed models: a "chain" badge with the active target, in the same way + as the alias badge. +- Operate → Runtime → **Failover**: a dense table with one row per chain + (state, active target, target states, time since the last switch). Each row + links to the chain in the model editor. With no chains, an empty state links + to the Failover chain template. + +### Live data + +A `useFailoverChains` hook fetches `GET /api/failover`, then opens an +`EventSource` on `/api/failover/events`. `snapshot` replaces the state; +`chain.switched` and `target.state` patch it. The browser reconnects the +stream, and the hook polls every 15 s as a fallback (the Agent Status +pattern). A `failoverApi` group in `src/utils/api.js` holds the calls. + +### Conventions + +- Design tokens and CSS classes only; no new inline styles (inline-style + ratchet). +- `StatusPill` tones: success for healthy and primary, warning for recovering + and fallback, error for down and degraded, muted for missing. +- Strings in the `models` and `admin` i18n namespaces for all 8 locales. +- Pin controls are hidden when `useAuth().isAdmin` is false. + +### Tests + +Playwright specs with mocked APIs and a mocked `text/event-stream`: the editor +component, the template, the health strip updating on events, pin controls +for admins and not for other users, and the overview page. UI line coverage +stays at or above `core/http/react-ui/coverage-baseline.txt`. + +--- + +## D. Contributor rule + +- New guide `.agents/distributed-state.md`. A feature that keeps runtime + state (in-memory maps, caches, pins, schedulers, background loops, probes) + chooses one mode and documents it: + - **shared**: `syncstate.SyncedMap`, with a `Store` when the state must + survive a restart; + - **single-runner**: an `advisorylock` leader loop; + - **stateless per request**; + - **per-instance**: allowed only with the reason written in the feature's + docs. + + The guide gives one real example per mode (finetune jobs, the node health + monitor, open responses, failover chains). Shared and single-runner + features include a fakebus test with two instances. +- `AGENTS.md`: a Quick Reference bullet "Distributed-aware state" and a row in + the Topics table. +- `.agents/api-endpoints-and-auth.md`: a checklist line "Stateful feature: + distributed mode chosen and documented (see distributed-state.md)". + +## Order of work + +C first (it changes code already in the PR and the event contract the UI +reads), then A (it needs the `Unimplemented` classification and the rerank +handler), then B, then D. All in PR #12285. From afd4fcdf091b99e2083d2087c0a6c39c64e56e52 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 21:48:49 +0000 Subject: [PATCH 33/79] docs: refine distributed sync, warm health and the live bridge Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- ...26-failover-distributed-proxy-ui-design.md | 37 ++++++++++++++----- 1 file changed, 27 insertions(+), 10 deletions(-) diff --git a/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md b/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md index e380a2888..0390539d2 100644 --- a/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md +++ b/docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md @@ -60,8 +60,8 @@ below applies when distributed mode is on. | State | Writers | Mechanism | Survives restart | |---|---|---|---| | Pins | any frontend (REST, MCP) | `syncstate.SyncedMap` named `failover.pins`, key = chain, with a gorm `Store` | yes | -| Target health: state, last error, since, consecutive passes | any frontend on a local transition, and the probe leader | `syncstate.SyncedMap` named `failover.targets`, key = target, NATS only, `Reconcile` on | no | -| Chain state: active target, active since, chain state | the probe leader only | `syncstate.SyncedMap` named `failover.chains`, key = chain, NATS only, `Reconcile` on | no | +| Target health: state, last error, since, consecutive passes | any frontend on a local transition, and the probe leader | `syncstate.SyncedMap` named `failover.targets`, key = target, NATS only | no | +| Chain state: active target, active since, chain state | the probe leader only | `syncstate.SyncedMap` named `failover.chains`, key = chain, NATS only | no | - The pins table is created under `advisorylock.KeySchemaMigrate`, the same way the jobs store creates its tables. @@ -70,6 +70,10 @@ below applies when distributed mode is on. maps. The manager does not import NATS or gorm directly. - A peer delta is applied through `OnApply`, which changes local state without publishing again (no echo loops). +- The NATS-only maps have no `Store`, so `Reconcile` would re-hydrate them + empty. Instead the leader republishes every target and chain snapshot every + 10 s. A frontend that joins late converges within 10 s and uses its own + state until then. ### Who does what @@ -101,10 +105,10 @@ therefore shows the same events. - The SmartRouter's and ReplicaReconciler's pinned-model resolver includes `WarmTargets()`, so workers never evict a warm target. - Only the leader preloads warm targets. -- A frontend cannot tell from its local model store whether a worker has a - model loaded. In distributed mode, warm-target liveness asks the node - registry whether a healthy node serves the model. When none does, the probe - is inconclusive: it neither passes nor trips, and real requests decide. +- A frontend's loaded check sees only its own model stubs. A warm target + without a stub on the leader is treated as not loaded: its liveness passes + and its recovery is inconclusive (it heals after `min_dwell`). Worker health + is left to the node health monitor and to real requests. ### Tests @@ -202,11 +206,24 @@ Core passes some inputs and outputs as local paths: ### Live transcription bridge -`AudioTranscriptionLive` opens a WebSocket to the upstream `/v1/realtime` -transcription session: +The upstream realtime API needs a pipeline model (VAD and transcription). The +proxy takes it from the model's backend options: -1. On the first `TranscriptLiveConfig`, send the session update with the - upstream model, language and sample rate, then answer `ready`. +```yaml +options: + - realtime_pipeline:argus-transcribe # an upstream pipeline config +``` + +Without this option, `AudioTranscriptionLive` returns the standard +"live transcription unsupported" error, and realtime uses its non-live +transcription path for the stage. + +With the option, `AudioTranscriptionLive` opens a WebSocket to +`/v1/realtime?model=`: + +1. On the first `TranscriptLiveConfig`, send `session.update` with + `type: transcription`, the input rate, the language and server VAD turn + detection. Answer `ready` when `session.updated` arrives. 2. Forward each `TranscriptLiveAudio` as `input_audio_buffer.append` (PCM float to PCM16 base64 at the session rate). 3. Map `conversation.item.input_audio_transcription.delta` to `delta`, and From 6738478c8a1235b3050c89c3a9eb21b4b490c15f Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 21:51:51 +0000 Subject: [PATCH 34/79] docs: plan distributed failover, localai-proxy and the chain UI Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- ...026-09-26-failover-distributed-proxy-ui.md | 689 ++++++++++++++++++ 1 file changed, 689 insertions(+) create mode 100644 docs/superpowers/plans/2026-09-26-failover-distributed-proxy-ui.md diff --git a/docs/superpowers/plans/2026-09-26-failover-distributed-proxy-ui.md b/docs/superpowers/plans/2026-09-26-failover-distributed-proxy-ui.md new file mode 100644 index 000000000..64d13e241 --- /dev/null +++ b/docs/superpowers/plans/2026-09-26-failover-distributed-proxy-ui.md @@ -0,0 +1,689 @@ +# Failover: Distributed Mode, localai-proxy and WebUI Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Make failover chains consistent across distributed frontends, add a `localai-proxy` backend that serves every LocalAI API (including live transcription) from a remote LocalAI, add the chain UI, and add the distributed-state contributor rule — all in PR #12285. + +**Architecture:** The failover `Manager` gets a `StateSync` dependency (no-op standalone; three `syncstate.SyncedMap`s in distributed mode) and a leader gate (PostgreSQL advisory lock per tick) so one frontend probes and owns chain state. `localai-proxy` is a new Go gRPC backend mapping each backend method to the upstream LocalAI REST endpoint, with a WebSocket bridge to the upstream realtime API for live transcription. The React UI adds an editor field, a template, a live health strip, a badge and an overview page on the existing failover REST/SSE API. + +**Tech Stack:** Go (echo, gRPC, gorm, NATS via `syncstate`, gorilla/websocket), Ginkgo/Gomega, React 19 + Vite, Playwright. + +**Spec:** `docs/superpowers/specs/2026-09-26-failover-distributed-proxy-ui-design.md` (builds on `docs/superpowers/specs/2026-09-26-model-failover-chains-design.md`) + +## Global Constraints + +- Worktree `/home/mudler/_git/LocalAI/.wt/failover-chains`, branch `feat/failover-chains`, PR #12285. Push only at the end (the controller does it). +- Commit trailer exactly `Assisted-by: Claude:claude-opus-5-5`. Never `Co-Authored-By` or `Signed-off-by`. `docs/superpowers/` needs `git add -f`. +- Build scope: `go build ./core/... ./pkg/... ./tests/... ./backend/go/localai-proxy/...` — never `go build ./...` (CGo launcher/backends need X11). +- Root `core/http` tests need `LOCALAI_TEST_HTTP_PORT=19391`. +- Logging `github.com/mudler/xlog`; `any` not `interface{}`; comments explain why. +- Standalone mode (no NATS/DB) must behave exactly as before this plan. +- SyncedMap names exactly: `failover.pins`, `failover.targets`, `failover.chains`. Leader republish interval 10 s. Advisory lock key constant `KeyFailoverProber` = 108. +- Backend name exactly `localai-proxy`; backend option `realtime_pipeline:`; `Unimplemented` message format `localai-proxy: has no upstream counterpart`. +- UI: no new inline `style={{}}` (inline-style ratchet); `StatusPill` tones: healthy/primary → success, recovering/fallback → warning, down/degraded → error, missing → muted; strings in i18n for all 8 locales (en, it, es, de, zh-CN, id, ko, pt-BR); pin controls only for `isAdmin`. +- Coverage baselines (`coverage-baseline.txt` 54.2, `core/http/react-ui/coverage-baseline.txt` 40.0) must not go down; never edit them. + +## Review Focus + +1. A NATS echo of a frontend's own publish must not deadlock or double-emit events (Task 1: "echo of own publish is a no-op"). +2. A frontend that joins late (no deltas yet) must still serve requests from local state, then converge on the leader's republish (Task 2: "late joiner converges on heartbeat"). +3. Leadership moving between frontends must not reset fail-back timing (shared `activeSince`) (Task 3: "new leader keeps activeSince"). +4. A `localai-proxy` target that returns `Unimplemented` must move the request to the next chain target without tripping (Task 4: "Unimplemented skips without trip"). +5. An upstream disconnect mid live-transcription must end the gRPC stream with `Unavailable`, not hang (Task 8: "upstream disconnect ends the stream"). + +--- + +## Part C — Distributed-aware failover + +### Task 1: Manager state-sync hooks and leader gate + +**Files:** +- Create: `core/services/failover/statesync.go`, `core/services/failover/statesync_test.go` +- Modify: `core/services/failover/manager.go`, `core/services/failover/schedule.go` + +**Interfaces:** +- Consumes: existing `Manager` internals (`setTargetLocked`, `recomputeLocked`, `Pin`, `Unpin`, `emitLocked`, `chainState`, `targetState`, `takeWarmLocked`). +- Produces: + ```go + type TargetSnapshot struct { + Target string `json:"target"` + State TargetState `json:"state"` + Reason Reason `json:"reason"` + Error string `json:"error,omitempty"` + ConsecutiveOK int `json:"consecutive_ok"` + Since time.Time `json:"since"` + } + type ChainSnapshot struct { + Chain string `json:"chain"` + Active string `json:"active"` + ActiveSince time.Time `json:"active_since"` + State ChainState `json:"state"` + Reason Reason `json:"reason"` + } + type StateSync interface { + PublishTarget(TargetSnapshot) + PublishChain(ChainSnapshot) + SetPin(chain, target string) error + ClearPin(chain string) error + Pins() map[string]string + } + type LeaderGate func(ctx context.Context, fn func()) bool + func WithLeaderGate(g LeaderGate) Option + func (m *Manager) SetStateSync(s StateSync) + func (m *Manager) ApplyTarget(s TargetSnapshot) + func (m *Manager) ApplyChain(s ChainSnapshot) + func (m *Manager) ApplyPin(chain, target string) // target "" = unpinned + func (m *Manager) IsLeader() bool + func (m *Manager) Republish() // leader: publish every target and chain snapshot + ``` + +Design rules (binding): +- Publishing happens **outside `m.mu`**: a real or fake bus delivers the frontend's own publish back synchronously to `ApplyTarget`/`ApplyChain`/`ApplyPin`, which take `m.mu`. Queue publishes in `m.pending []func()` while locked; every exported mutating method drains and runs them after unlocking (helper `m.unlockAndFlush()`). +- `ApplyTarget` sets state through `setTargetLocked` with `m.applying = true`, so it emits the local `target.state` event but does not publish. An echo whose state equals the local state is a no-op (existing early return). +- Local transitions (`setTargetLocked` with `!m.applying` and `m.sync != nil`) queue `PublishTarget`. +- Chains: when `m.sync != nil && !m.leader`, `recomputeLocked` does not change `ch.active` once the chain has adopted leader state (`ch.adopted`); before adoption it computes locally. `ApplyChain` sets `active`, `activeSince`, `state`, `adopted = true` and emits `chain.switched` (with the snapshot's reason) when `active` or `state` changed. When the leader's recompute changes a chain, it queues `PublishChain`. +- Pins: `Pin`/`Unpin` set local state immediately (read-your-writes), then call `sync.SetPin`/`ClearPin` outside the lock. `ApplyPin` sets/clears `ch.pinned` and recomputes with `ReasonManual`. `SetStateSync` hydrates pins from `s.Pins()`. +- Leader gate: in `Tick`, after `Sync()`, call `gate(ctx, fn)` where `fn` runs probe scheduling + `Reevaluate()`; set `m.leader` to the returned bool. Without a gate (standalone) the manager is always leader and `Tick` behaves exactly as today. Followers still run `Reevaluate()` (local-only when not adopted). On a false→true leader transition set `m.warmPending = true` so `onWarm` fires on the new leader. +- `Republish()` queues `PublishTarget` for every target and `PublishChain` for every chain (leader only; no-op otherwise). `Run` calls it every 10 ticks when leader. + +- [ ] **Step 1: Write the failing tests** (`statesync_test.go`, package `failover`, reuse `fakeClock`, `fakeSource`, `local`, `remote`, `chainCfg`, `t`, `errBoom`, `drain` from existing test files) + +```go +package failover + +import ( + "context" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// loopSync is an in-process StateSync that delivers every publish to all +// managers synchronously, including the publisher (like NATS echo). +type loopSync struct { + mu sync.Mutex + peers []*Manager + pins map[string]string +} + +func (l *loopSync) add(m *Manager) { l.mu.Lock(); l.peers = append(l.peers, m); l.mu.Unlock() } +func (l *loopSync) each(f func(*Manager)) { + l.mu.Lock() + ps := append([]*Manager(nil), l.peers...) + l.mu.Unlock() + for _, p := range ps { + f(p) + } +} +func (l *loopSync) PublishTarget(s TargetSnapshot) { l.each(func(m *Manager) { m.ApplyTarget(s) }) } +func (l *loopSync) PublishChain(s ChainSnapshot) { l.each(func(m *Manager) { m.ApplyChain(s) }) } +func (l *loopSync) SetPin(c, t string) error { + l.mu.Lock(); l.pins[c] = t; l.mu.Unlock() + l.each(func(m *Manager) { m.ApplyPin(c, t) }) + return nil +} +func (l *loopSync) ClearPin(c string) error { + l.mu.Lock(); delete(l.pins, c); l.mu.Unlock() + l.each(func(m *Manager) { m.ApplyPin(c, "") }) + return nil +} +func (l *loopSync) Pins() map[string]string { + l.mu.Lock(); defer l.mu.Unlock() + out := map[string]string{} + for k, v := range l.pins { + out[k] = v + } + return out +} + +var _ = Describe("Manager state sync", func() { + var ( + clock *fakeClock + src *fakeSource + bus *loopSync + a, b *Manager + leaderIsA bool + gateFor func(isA bool) LeaderGate + ctx = context.Background() + ) + + BeforeEach(func() { + clock = newFakeClock() + src = newFakeSource(remote("x"), local("y"), chainCfg("chain", nil, t("x"), t("y"))) + bus = &loopSync{pins: map[string]string{}} + leaderIsA = true + gateFor = func(isA bool) LeaderGate { + return func(_ context.Context, fn func()) bool { + if isA != leaderIsA { + return false + } + fn() + return true + } + } + a = New(src, WithClock(clock), WithLeaderGate(gateFor(true))) + b = New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + bus.add(a); bus.add(b) + a.SetStateSync(bus); b.SetStateSync(bus) + a.Tick(ctx); b.Tick(ctx) + }) + + It("echo of own publish is a no-op and emits one event", func() { + events, cancel := a.Subscribe(16) + defer cancel() + a.ReportFailure("x", errBoom) + n := 0 + for _, e := range drain(events) { + if e.Type == EventTargetState && e.Target == "x" { + n++ + } + } + Expect(n).To(Equal(1)) + }) + + It("a trip on one frontend is skipped by the other's plan", func() { + b.ReportFailure("x", errBoom) + att, err := a.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("y")) + }) + + It("followers adopt the leader's chain state and emit the switch", func() { + events, cancel := b.Subscribe(16) + defer cancel() + a.ReportFailure("x", errBoom) // leader recomputes and publishes chain state + st, _ := b.ChainStatus("chain") + Expect(st.Active).To(Equal("y")) + var sw []Event + for _, e := range drain(events) { + if e.Type == EventChainSwitched { + sw = append(sw, e) + } + } + Expect(sw).ToNot(BeEmpty()) + }) + + It("a pin on one frontend applies on all", func() { + Expect(b.Pin("chain", "y")).To(Succeed()) + st, _ := a.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + Expect(a.Unpin("chain")).To(Succeed()) + st, _ = b.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + + It("hydrates pins when the sync is attached", func() { + bus.pins["chain"] = "y" + c := New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + c.SetStateSync(bus) + st, _ := c.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + }) + + It("only the leader probes", func() { + pa, pb := &fakeProber{fail: map[string]error{}}, &fakeProber{fail: map[string]error{}} + a = New(src, WithClock(clock), WithProber(pa), WithLeaderGate(gateFor(true))) + b = New(src, WithClock(clock), WithProber(pb), WithLeaderGate(gateFor(false))) + a.SetStateSync(bus); b.SetStateSync(bus) + a.Tick(ctx); b.Tick(ctx) + Eventually(func() int { return len(pa.take()) }).Should(BeNumerically(">", 0)) + Consistently(func() int { return len(pb.take()) }, 200*time.Millisecond).Should(Equal(0)) + Expect(a.IsLeader()).To(BeTrue()) + Expect(b.IsLeader()).To(BeFalse()) + }) + + It("new leader keeps activeSince across a leadership move", func() { + a.ReportFailure("x", errBoom) + before, _ := b.ChainStatus("chain") + leaderIsA = false + clock.Advance(5 * time.Second) + a.Tick(ctx); b.Tick(ctx) + after, _ := b.ChainStatus("chain") + Expect(after.ActiveSince).To(Equal(before.ActiveSince)) + Expect(b.IsLeader()).To(BeTrue()) + }) + + It("republish sends every target and chain", func() { + c := New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + bus.add(c) + c.SetStateSync(bus) + a.ReportFailure("x", errBoom) + a.Republish() + st, _ := c.ChainStatus("chain") + Expect(st.Active).To(Equal("y")) + }) + + It("standalone manager (no sync, no gate) is always leader", func() { + m := New(src, WithClock(clock)) + m.Tick(ctx) + Expect(m.IsLeader()).To(BeTrue()) + }) +}) +``` + +(`fakeProber` and its `take()` exist in `schedule_test.go`; if the in-flight probe tracking needs waiting, use `Eventually` as shown.) + +- [ ] **Step 2: Run to verify failure** + +Run: `go test ./core/services/failover/... 2>&1 | tail -5` +Expected: compile failure (`undefined: TargetSnapshot`). + +- [ ] **Step 3: Implement** `statesync.go` (types, interface, `WithLeaderGate`, `SetStateSync`, `Apply*`, `IsLeader`, `Republish`, `unlockAndFlush`) and the hooks in `manager.go`/`schedule.go` following the design rules above. Add fields to `Manager`: `sync StateSync`, `gate LeaderGate`, `leader bool` (guarded by `mu`), `applying bool`, `pending []func()`, `ticks int`; to `chainState`: `adopted bool`. Keep every existing test green. + +- [ ] **Step 4: Run tests** + +Run: `go test -race -count=3 ./core/services/failover/... 2>&1 | tail -10` +Expected: PASS, no races. + +- [ ] **Step 5: Commit** + +```bash +git add core/services/failover +git commit -m "feat(failover): share state through a sync hook and gate probes on a leader + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 2: syncstate-backed StateSync with a durable pin store + +**Files:** +- Create: `core/services/failover/distsync/distsync.go`, `core/services/failover/distsync/pinstore.go`, `core/services/failover/distsync/distsync_suite_test.go`, `core/services/failover/distsync/distsync_test.go` + +**Interfaces:** +- Consumes: Task 1 `StateSync`, `TargetSnapshot`, `ChainSnapshot`, `Manager.ApplyTarget/ApplyChain/ApplyPin`; `syncstate.New/Config/Store`; `messaging.MessagingClient`; `advisorylock.WithLockCtx`, `advisorylock.KeySchemaMigrate`; `testutil.NewFakeBus`. +- Produces: + ```go + type PinRecord struct { + Chain string `gorm:"primaryKey" json:"chain"` + Target string `json:"target"` + UpdatedAt time.Time `json:"updated_at"` + } + func (PinRecord) TableName() string { return "failover_pins" } + func NewPinStore(db *gorm.DB) (*PinStore, error) // migrates under KeySchemaMigrate + func New(ctx context.Context, nats messaging.MessagingClient, pins syncstate.Store[string, PinRecord], m *failover.Manager) (*Sync, error) + func (s *Sync) Close() error + // *Sync implements failover.StateSync + ``` + +Rules: +- Three maps: `failover.pins` (key chain, `Store: pins` when non-nil — guard against a typed-nil interface as `finetune/service.go` does), `failover.targets` (key target, NATS only), `failover.chains` (key chain, NATS only). No `Reconcile` on NATS-only maps (it would re-hydrate them empty); the leader's `Republish` covers late joiners. +- `OnApply` for targets/chains calls `m.ApplyTarget`/`m.ApplyChain`; for pins `m.ApplyPin(chain, target)` on "set" and `m.ApplyPin(chain, "")` on "delete". +- `New` starts all three maps, then calls `m.SetStateSync(s)` (which hydrates pins from `s.Pins()`). + +- [ ] **Step 1: Write the failing tests** — two managers on one `testutil.NewFakeBus()`, each with its own `distsync.New`, sharing an in-memory `syncstate.Store` fake for pins (a map with a mutex implementing `List/Upsert/Delete`). Specs: + - "a pin on A is visible on B and survives a new instance C built from the same store" (C's `ChainStatus` shows the pin right after `New`). + - "a trip on B makes A's plan skip the target". + - "the leader's chain switch reaches the follower". + - "late joiner converges on heartbeat": build C after A tripped a target; before `A.Republish()` C shows the primary active; after it, the fallback. + - "PinStore round-trips" using a sqlite gorm DB (`gorm.io/driver/sqlite`, `file::memory:`) if the repo already depends on it (check `go.mod`); otherwise skip this spec and test the adapter with the in-memory store only, and say so in the report. + +- [ ] **Step 2: Run to verify failure** — `go test ./core/services/failover/distsync/...` → compile failure. +- [ ] **Step 3: Implement** `distsync.go` and `pinstore.go` per the rules (PinStore: `List` = `db.Find`, `Upsert` = `db.Save`, `Delete` = `db.Delete(&PinRecord{Chain: k})`; migration via `advisorylock.WithLockCtx(context.Background(), db, advisorylock.KeySchemaMigrate, func() error { return db.AutoMigrate(&PinRecord{}) })`). +- [ ] **Step 4: Run tests** — `go test -race ./core/services/failover/...` → PASS. +- [ ] **Step 5: Commit** + +```bash +git add core/services/failover/distsync +git commit -m "feat(failover): sync pins, target health and chain state over NATS + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 3: Distributed wiring, warm pins on workers, docs + +**Files:** +- Modify: `core/services/advisorylock/keys.go` (add `KeyFailoverProber = 108` with a comment) +- Modify: `core/application/startup.go` (after `application.distributed = distSvc`, ~L294, and before `Run` ~L570), `core/application/distributed.go` (~L380-395, L458: pinned resolver), `core/application/failover.go` (preload only when leader) +- Create: `core/application/failover_distributed.go`, `core/application/failover_distributed_test.go` +- Modify: `docs/content/features/model-failover.md` (replace the per-instance limit with a "Distributed mode" section) + +**Interfaces:** +- Consumes: Tasks 1-2; `advisorylock.TryWithLockCtx(ctx, db, key, fn func() error) (bool, error)`; `a.IsDistributed()`, `a.Distributed().Nats`, `a.distributedDB()`; `nodes.PinnedModelResolver` (`GetPinnedModelNames() []string`); `failover.MergePinned`. +- Produces: `func failoverLeaderGate(db *gorm.DB) failover.LeaderGate`; `type failoverPinnedResolver struct{ base nodes.PinnedModelResolver; fm *failover.Manager }` implementing `GetPinnedModelNames()`. + +Rules: +- Leader gate: `func(ctx, fn) bool { ok, err := advisorylock.TryWithLockCtx(ctx, db, advisorylock.KeyFailoverProber, func() error { fn(); return nil }); if err != nil { xlog.Warn(...); return false }; return ok }`. The failover manager is constructed before distributed init (startup.go ~L257), so give it the gate through a setter `(*Manager).SetLeaderGate(LeaderGate)` added in this task (mirror of `WithLeaderGate`), then call `distsync.New(ctx, distSvc.Nats, pinStore, application.failoverManager)`; on error log and continue standalone. +- Pinned resolver: wrap `configLoader` so `GetPinnedModelNames()` returns `failover.MergePinned(configLoader.GetPinnedModelNames(), fm.WarmTargets())`, and pass it to both the SmartRouter and the ReplicaReconciler options. Adapt `initDistributed`'s signature minimally (add a `pinned nodes.PinnedModelResolver` parameter or set it after construction — choose the smaller change and explain). +- Warm preload: `applyFailoverWarmTargets` keeps the pin sync, and runs the preload goroutine only when `!a.IsDistributed() || a.failoverManager.IsLeader()`. +- `Run` calls `Republish()` every 10 ticks when leader (implemented in Task 1; verify here). + +- [ ] **Step 1: Write the failing tests** (`failover_distributed_test.go`, in the application package's existing test suite): + - `failoverPinnedResolver` merges config pins and warm targets without duplicates. + - `failoverLeaderGate` with a gorm DB that is not PostgreSQL (the advisorylock package falls back to an in-process lock for non-Postgres DBs): two gates on the same DB — while one is inside `fn`, the other returns false; afterwards the other returns true. Use the same DB setup other `core/application` or `advisorylock` tests use (read `core/services/advisorylock/*_test.go` first). +- [ ] **Step 2: Run to verify failure.** +- [ ] **Step 3: Implement** per the rules; add the `keys.go` constant; update docs: in `model-failover.md`, replace the "Failover state is kept in memory by each LocalAI instance" limit with a `## Distributed mode` section: pins are cluster-wide and persisted; target health and chain state are shared; one frontend (the probe leader) probes, fails back and preloads warm targets; warm targets are pinned on workers; a frontend that joins late converges within 10 s. +- [ ] **Step 4: Run tests** — `go test -race ./core/application/... ./core/services/failover/... ./core/services/advisorylock/...` and `go build ./core/... ./pkg/... ./tests/...` → PASS. +- [ ] **Step 5: Commit** + +```bash +git add core/services/advisorylock core/application core/services/failover docs/content/features/model-failover.md +git commit -m "feat(failover): run one prober per cluster and pin warm targets on workers + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +## Part A — localai-proxy backend + +### Task 4: Core support — Rerank for Go backends, Unimplemented skip, proxy options + +**Files:** +- Modify: `pkg/grpc/interface.go` (add `RerankModel`), `pkg/grpc/server.go` (add `Rerank` handler), `pkg/grpc/model_identity_modalities_test.go` (update the note/test that says Go servers have no Rerank) +- Modify: `core/services/failover/classify.go` (add `IsCapabilityGap`), `core/services/failover/manager.go` (`Do` uses it), `core/http/middleware/failover.go` (retry uses it) +- Modify: `core/backend/options.go:545` (proxy options for `localai-proxy`), `core/config/model_config_loader.go` (load-time warnings) +- Test: `pkg/grpc/server_rerank_test.go` (or the package's existing server test file), `core/services/failover/classify_test.go`, `core/services/failover/manager_test.go`, `core/http/middleware/failover_test.go`, `core/config/model_config_loader_test.go` + +**Interfaces:** +- Produces: + ```go + // pkg/grpc/interface.go + type RerankModel interface { + Rerank(context.Context, *pb.RerankRequest) (*pb.RerankResult, error) + } + // core/services/failover/classify.go + func IsCapabilityGap(err error) bool // gRPC Unimplemented anywhere in the chain + ``` + +- [ ] **Step 1: Write the failing tests** + - `pkg/grpc`: a fake model embedding `base.Base` and implementing `RerankModel` is served by `server.Rerank` (use the package's existing in-process server test pattern for `Score`); a model without it returns `codes.Unimplemented`. + - `classify_test.go`: `IsCapabilityGap(grpcstatus.Error(codes.Unimplemented, "x"))` true; wrapped with `fmt.Errorf("%w")` true; `codes.Unavailable` false; nil false. + - `manager_test.go` ("Unimplemented skips without trip"): `m.Do` where target `a` returns `grpcstatus.Error(codes.Unimplemented, "localai-proxy: X has no upstream counterpart")` and `b` succeeds → `tried == [a b]`, `a` stays `StateHealthy`. + - `failover_test.go` (middleware): handler for `a` returns the same Unimplemented error → served by `b`, `a` healthy. + - `model_config_loader_test.go`: a `backend: localai-proxy` config with `proxy.mode: translate` logs a warning and still loads; one without `known_usecases` loads (warning). Use the loader test's existing log-capture approach if any; otherwise assert only that loading succeeds and note it. +- [ ] **Step 2: Run to verify failure.** +- [ ] **Step 3: Implement** + ```go + // pkg/grpc/server.go — copy of the Score handler shape + func (s *server) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } + rm, ok := s.llm.(RerankModel) + if !ok { + return nil, status.Errorf(codes.Unimplemented, "method Rerank not implemented") + } + if s.llm.Locking() { + s.llm.Lock() + defer s.llm.Unlock() + } + return rm.Rerank(ctx, in) + } + ``` + (If `checkModelIdentity` does not accept `*pb.RerankRequest`, add the `GetModelIdentity()` method set it needs — `RerankRequest` has field `ModelIdentity`.) + ```go + // classify.go + // IsCapabilityGap reports a target that cannot serve this kind of request at + // all. The next target may serve it, and this target is not broken. + func IsCapabilityGap(err error) bool { + if err == nil { + return false + } + st, ok := grpcstatus.FromError(err) + return ok && st.Code() == codes.Unimplemented + } + ``` + In `Manager.Do`, before the retryable check: `if IsCapabilityGap(err) && !committed.Load() { if !att.Skip() { return err }; continue }`. In `failoverRetry`, treat `failover.IsCapabilityGap(err)` exactly like the admission-rejection branch (spill with `att.Skip()`, no trip). In `options.go` change the condition to `if c.Backend == "cloud-proxy" || c.Backend == "localai-proxy"`. In the loader's load-time pass, `xlog.Warn` for `localai-proxy` configs that set `proxy.mode`/`proxy.provider` (ignored) or have no `known_usecases`. +- [ ] **Step 4: Run tests** — `go test -race ./pkg/grpc/... ./core/services/failover/... ./core/http/middleware/... ./core/config/...` → PASS. +- [ ] **Step 5: Commit** + +```bash +git add pkg/grpc core/services/failover core/http/middleware core/backend/options.go core/config +git commit -m "feat(grpc): serve Rerank from Go backends and skip Unimplemented targets + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 5: localai-proxy backend — skeleton, packaging, text methods + +**Files:** +- Create: `backend/go/localai-proxy/{main.go,proxy.go,client.go,text.go,Makefile,package.sh,run.sh}`, `backend/go/localai-proxy/{localai_proxy_suite_test.go,fake_upstream_test.go,text_test.go}` +- Modify: root `Makefile` (6 places, mirroring cloud-proxy: `.NOTPARALLEL`, `TEST_PATHS`, `BACKEND_LOCALAI_PROXY = localai-proxy|golang|.|false|true`, `$(eval $(call generate-docker-build-target,$(BACKEND_LOCALAI_PROXY)))`, `docker-build-localai-proxy` in `docker-build-backends`, `build-localai-proxy-backend`/`clean-localai-proxy-backend` e2e helpers building to `tests/e2e/mock-backend/localai-proxy`), `backend/index.yaml` (meta + cpu/metal latest/development images, mirroring cloud-proxy), `.github/backend-matrix.yml` (amd64, arm64, darwin entries mirroring cloud-proxy), `.gitignore` (the e2e binary) + +**Interfaces:** +- Consumes: `pkg/grpc` (`StartServer`, `AIModelRich`, `RerankModel`, `ScoreModel`), `pkg/grpc/base.Base`, `pkg/httpclient.New`. +- Produces: `type LocalAIProxy struct { base.Base; cfg atomic.Pointer[proxyConfig]; client *http.Client }`, `func NewLocalAIProxy() *LocalAIProxy`, `type proxyConfig struct { base, upstreamModel, apiKey, realtimePipeline string; timeout time.Duration }`, helpers `func (p *LocalAIProxy) postJSON(ctx, path string, body, out any) error`, `func (p *LocalAIProxy) postMultipart(ctx, path string, fields map[string]string, fileField, filePath string, out any) error`, `func (p *LocalAIProxy) postStream(ctx, path string, body any) (*http.Response, error)`, `func (p *LocalAIProxy) model(req string) string` (upstream model name), `func unimplemented(method string) error` returning `status.Errorf(codes.Unimplemented, "localai-proxy: %s has no upstream counterpart", method)`. + +Rules: +- `Load`: requires `opts.GetProxy()` non-nil and a valid `upstream_url` (base URL; strip a trailing `/`); resolves the API key with the same rules as cloud-proxy's `resolveAPIKey` (copy the function); warns if `mode`/`provider` set; reads `realtime_pipeline:` from `opts.GetOptions()`; `request_timeout_seconds` becomes a per-request context timeout for non-streaming calls; model name = `upstream_model`, else `opts.GetModel()`. +- Auth: `Authorization: Bearer ` when a key is set. HTTP client from `httpclient.New()` (no redirects). +- Upstream non-2xx → a gRPC error: 5xx/transport → `codes.Unavailable`; 4xx → `codes.InvalidArgument` (so failover does not trip on client errors); body text in the message (truncated to 500 chars). +- Text methods in `text.go`: `PredictRich`/`PredictStreamRich` via `/v1/chat/completions` when `opts.GetMessages()` is non-empty, else `/v1/completions` with `opts.GetPrompt()` (map tokens, temperature, top_p, top_k, stop, seed; stream parses SSE `data:` lines and sends `pb.Reply{Message}` per delta; do not close the channel); legacy `Predict`/`PredictStream` wrap them; `Embeddings` → `/v1/embeddings` (`input` = `opts.GetEmbeddings()`, returns `data[0].embedding`); `Rerank` → `/v1/rerank`; `TokenizeString` → `/v1/tokenize`; `Score` → `/api/score` (read `core/http/endpoints` for the request shape). +- Methods with no counterpart return `unimplemented("")`: `AudioEncode`, `AudioDecode`, `AudioToAudioStream`, `TokenClassify`, `ModelMetadata`, fine-tune and quantization methods. `Status` keeps the base implementation. + +- [ ] **Step 1: Write the failing tests** — `fake_upstream_test.go`: an `httptest.Server` recording method, path, `Authorization`, and JSON/multipart body, with per-path scripted responses (JSON or SSE). `text_test.go` specs: Load rejects missing proxy options and a bad URL; Load parses `realtime_pipeline`; `PredictRich` hits `/v1/chat/completions` with the upstream model and returns the content; `PredictStreamRich` streams SSE deltas in order; `Embeddings`, `Rerank`, `TokenizeString` hit their paths and map results; a 503 upstream → `codes.Unavailable`; a 400 → `codes.InvalidArgument`; `AudioEncode` → `Unimplemented` with the exact message. +- [ ] **Step 2: Run to verify failure** — `go test ./backend/go/localai-proxy/...`. +- [ ] **Step 3: Implement** per the rules; packaging per the Files list (copy cloud-proxy's `Makefile`/`package.sh`/`run.sh` with the binary renamed). +- [ ] **Step 4: Run tests** — `go test -race ./backend/go/localai-proxy/...`; `make -C backend/go/localai-proxy build`; `make build-localai-proxy-backend` → OK. Validate YAML: `python3 -c "import yaml,sys; yaml.safe_load(open('backend/index.yaml')); yaml.safe_load(open('.github/backend-matrix.yml'))"`. +- [ ] **Step 5: Commit** + +```bash +git add backend/go/localai-proxy Makefile backend/index.yaml .github/backend-matrix.yml .gitignore +git commit -m "feat(localai-proxy): add a backend that serves text APIs from a remote LocalAI + +Assisted-by: Claude:claude-opus-5-5" +``` + +--- + +### Task 6: localai-proxy — audio methods + +**Files:** +- Create: `backend/go/localai-proxy/audio.go`, `backend/go/localai-proxy/audio_test.go` + +**Interfaces:** consumes Task 5 helpers. + +Mapping (request → upstream → result): + +| Method | Upstream | Notes | +|---|---|---| +| `TTS(req)` | `POST /tts` JSON `{model, input: req.Text, voice, language}` | write the response bytes to `req.Dst` | +| `TTSStream(req, out)` | `POST /tts` with `stream: true` | copy the chunked `audio/wav` body to `out` as it arrives (header + PCM, unchanged); close `out` per the base contract | +| `SoundGeneration(req)` | `POST /v1/sound-generation` | write bytes to `req.Dst` | +| `AudioTranscription(ctx, req)` | `POST /v1/audio/transcriptions` multipart: `file` from `req.Dst` (input audio path), `model`, `language`, `translate`, `prompt`, `diarize` | map `TranscriptionResult{text, segments, words, language, duration}` to `pb.TranscriptResult` | +| `AudioTranscriptionStream(ctx, req, out)` | same, plus `stream=true` | SSE `transcript.text.delta` → `TranscriptStreamResponse{Delta}`; `transcript.text.done` → `FinalResult`; `error` event → return `codes.Unavailable` | +| `Diarize` | `POST /v1/audio/diarization` multipart | map segments | +| `VAD(req)` | `POST /v1/vad` JSON `{model, audio: req.Audio}` | map `segments[{start,end}]` | +| `SoundDetection(ctx, req)` | `POST /v1/audio/classification` multipart `file` from `req.Src`, `top_k`, `threshold` | map `detections[{index,label,score}]` | +| `AudioTransform` | `POST /audio/transformations` | read the handler for the request shape; write output to the path the request carries | + +Read each `pb` request/response message in `backend/backend.proto` and each REST schema in `core/schema` before mapping; keep field names exact. + +- [ ] **Step 1: Write the failing tests** — one spec per row using the fake upstream: path, multipart fields (file bytes equal the input file), JSON body, and result mapping; `TTS` writes the upstream bytes to `Dst`; `TTSStream` forwards chunks in order and the first chunk starts with `RIFF`; `AudioTranscriptionStream` emits deltas then the final result; upstream disconnect mid-stream → `codes.Unavailable`. +- [ ] **Step 2: Verify failure.** **Step 3: Implement.** **Step 4:** `go test -race ./backend/go/localai-proxy/...` → PASS. +- [ ] **Step 5: Commit** — `feat(localai-proxy): serve speech, transcription and audio APIs remotely`. + +--- + +### Task 7: localai-proxy — image, video, 3D and vision methods + +**Files:** +- Create: `backend/go/localai-proxy/media.go`, `backend/go/localai-proxy/media_test.go` + +Mapping: + +| Method | Upstream | Notes | +|---|---|---| +| `GenerateImage(req)` | `POST /v1/images/generations` `{model, prompt, negative_prompt, size: "WxH", step, seed, response_format: "b64_json"}`; `req.Src`/`ref_images` sent as base64 in `files`/`ref_images` per `core/schema` | decode `data[0].b64_json` into `req.Dst` | +| `UpscaleImage` | `POST /v1/images/upscale` | same output handling | +| `GenerateVideo` | `POST /video` | write the returned file (b64 or URL download relative to the upstream base) to `Dst` | +| `Generate3D`, `Animate3D` | `POST /3d/generations`, `/3d/animate` | same output handling; implement `AnimationMetadataModel` only if the upstream response carries the metadata | +| `Detect`, `Depth` | `POST /v1/detection`, `/v1/depth` | map results | +| `FaceVerify`, `FaceAnalyze` | `POST /v1/face/verify`, `/v1/face/analyze` | map results | +| `VoiceVerify`, `VoiceAnalyze`, `VoiceEmbed` | `POST /v1/voice/verify`, `/v1/voice/analyze`, `/v1/voice/embed` | map results | +| `StoresSet/Get/Delete/Find` | `POST /stores/set`, `/stores/get`, `/stores/delete`, `/stores/find` | map keys/values | + +When the upstream returns a URL instead of b64, download it with the same client (same auth) and write it to `Dst`. + +- [ ] **Step 1: Failing tests** — one spec per row with the fake upstream (path, key request fields, `Dst` written from b64 and from a URL). **Step 2** verify failure. **Step 3** implement. **Step 4** `go test -race ./backend/go/localai-proxy/...` → PASS. +- [ ] **Step 5: Commit** — `feat(localai-proxy): serve image, video, 3D and vision APIs remotely`. + +--- + +### Task 8: localai-proxy — live transcription bridge + +**Files:** +- Create: `backend/go/localai-proxy/live.go`, `backend/go/localai-proxy/live_test.go` + +**Interfaces:** `func (p *LocalAIProxy) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest, out chan<- *pb.TranscriptLiveResponse) error`. Contract (from `pkg/grpc/server.go:477-541`): the backend closes `out`; `in` closes on client EOF; return errors immediately (callers wait for the ready ack). + +Protocol (upstream `github.com/gorilla/websocket`): +1. No `realtime_pipeline` → `close(out); return grpcerrors.LiveTranscriptionUnsupported("localai-proxy", "set the realtime_pipeline backend option")` (same helper `base.Base` uses). +2. Read the first `in` message; it must be `config` (else `codes.InvalidArgument`). Rate = `config.sample_rate` or 16000. +3. Dial `ws(s):///v1/realtime?model=` with the bearer header. Read `session.created`. +4. Send `{"type":"session.update","session":{"type":"transcription","audio":{"input":{"format":{"type":"audio/pcm","rate":},"transcription":{"model":"","language":""},"turn_detection":{"type":"server_vad"}}}}}`. On `session.updated` send `TranscriptLiveResponse{Ready: true}`; on `error` return `codes.Unavailable` with its message. +5. Writer goroutine: each `audio.pcm` (float32 in [-1,1]) → PCM16 LE → base64 → `{"type":"input_audio_buffer.append","audio":"..."}`. +6. Reader goroutine: `conversation.item.input_audio_transcription.delta` → `{Delta}`; `...completed` → `{Delta: , Eou: true}` and append to the running final text; `...failed` or `error` → end with `codes.Unavailable`. +7. When `in` closes: wait up to 5 s for an in-flight `completed` (track `speech_started` without a matching completion), then send `{FinalResult: {Text: }}`, close the socket, close `out`, return nil. +8. Socket read error before step 7 → close `out`, return `status.Error(codes.Unavailable, ...)`. + +- [ ] **Step 1: Failing tests** with an `httptest` WebSocket server (gorilla `Upgrader`) scripting the upstream: unsupported without the option; ready after `session.updated`; audio frames arrive base64-PCM16 of the right length; deltas and a completion map to `Delta`/`Eou`; closing `in` yields `FinalResult` with the concatenated text; "upstream disconnect ends the stream" (server closes mid-session → `Unavailable`, `out` closed, no goroutine left blocked — assert the call returns within 2 s). +- [ ] **Step 2** verify failure. **Step 3** implement. **Step 4** `go test -race ./backend/go/localai-proxy/...` → PASS. +- [ ] **Step 5: Commit** — `feat(localai-proxy): bridge live transcription to the upstream realtime API`. + +--- + +### Task 9: localai-proxy end to end, and docs + +**Files:** +- Modify: `tests/e2e/e2e_suite_test.go` (build/locate the `localai-proxy` binary like cloud-proxy), create `tests/e2e/e2e_localai_proxy_test.go`, extend `tests/e2e/realtime_ws_test.go` (label `failover`) +- Modify: docs — the page that documents `cloud-proxy` (find with `grep -rln "cloud-proxy" docs/content`) gets a `localai-proxy` section; `docs/content/features/model-failover.md` gets a remote-LocalAI example with per-stage chains. + +Rules: +- Point `localai-proxy` models at the test server itself (`proxy.upstream_url` = the suite's base URL, `upstream_model` = an existing mock model), registered at runtime the way the cloud-proxy/failover e2e specs register models. +- Specs: chat, embeddings, TTS and transcription through `localai-proxy` return 2xx with the upstream model's answer; a chain `[proxy-target, local mock]` where the proxy target's upstream model is `fail-load-…` fails over to the local target; a realtime pipeline whose `llm` stage is a chain `[localai-proxy → mock LLM, mock LLM]` completes a turn and, after the proxy target is made to fail (point it at a `fail-load` upstream model or stop routing), switches stage with a `localai.model.failover` event. Live transcription through the bridge is covered by Task 8's unit specs; add an e2e only if the suite has a pipeline with streaming transcription available. +- Run `make build-localai-proxy-backend build-mock-backend` first. + +- [ ] Steps: write specs → run (fail) → implement registration + docs → run `go run github.com/onsi/ginkgo/v2/ginkgo --label-filter=failover -v ./tests/e2e` and the full `!real-models` e2e → commit `test(localai-proxy): proxy APIs and realtime stages end to end`. + +--- + +## Part B — WebUI + +### Task 10: Failover data layer and health strip + +**Files:** +- Modify: `core/http/react-ui/src/utils/config.js` (endpoints), `src/utils/api.js` (`failoverApi`), `src/components/StatusPill.jsx` (tones) +- Create: `src/hooks/useFailoverChains.js`, `src/components/FailoverChainStatus.jsx` +- Modify: `src/pages/ModelEditor.jsx` (strip after the `me-head` block, before the template selector, when `!isCreateMode` and the chain exists), `src/App.css` (classes), `public/locales/*/models.json` (strings, 8 locales) +- Test: `e2e/failover-health.spec.js` + +**Interfaces — Produces:** +```js +// utils/config.js endpoints +failoverChains: '/api/failover', +failoverChain: (name) => `/api/failover/${encodeURIComponent(name)}`, +failoverEvents: '/api/failover/events', +failoverPin: (name) => `/api/failover/${encodeURIComponent(name)}/pin`, +// utils/api.js +export const failoverApi = { + list: () => fetchJSON(API_CONFIG.endpoints.failoverChains), + get: (name) => fetchJSON(API_CONFIG.endpoints.failoverChain(name)), + pin: (name, target) => postJSON(API_CONFIG.endpoints.failoverPin(name), { target }), + unpin: (name) => fetchJSON(API_CONFIG.endpoints.failoverPin(name), { method: 'DELETE' }), + eventsUrl: () => API_CONFIG.endpoints.failoverEvents, +} +// hooks/useFailoverChains.js +export default function useFailoverChains() // → { chains: ChainStatus[], byName: {[name]: ChainStatus}, loading, error, refresh } +``` +Hook rules: initial `failoverApi.list()`; `new EventSource(apiUrl(failoverApi.eventsUrl()))`; `snapshot` → replace `chains` from `data.chains`; `chain.switched` → patch that chain's `active`, `state`, `active_since = data.at`; `target.state` → patch that target's `state`, `last_error = data.error` in every chain containing it; `onerror` no-op; poll `list()` every 15 s; close on unmount. + +`StatusPill` STATUS map additions: `primary: 'success'`, `fallback: 'warning'`, `recovering: 'warning'`, `degraded: 'error'`, `down: 'error'`, `missing: 'muted'` (`healthy` already maps to success). + +`FailoverChainStatus({ chain, onPin, onUnpin, canPin })` renders: chain `StatusPill` + active target + relative "since"; a compact table of targets (model, kind, warm, `StatusPill` state, last probe, last error truncated with title); per-target **Pin** button and a chain-level **Unpin** when `canPin`, each behind `ConfirmDialog`. Classes only. + +- [ ] **Step 1: Failing Playwright spec** `e2e/failover-health.spec.js` (import `test` from `./coverage-fixtures.js`; mock `**/api/auth/status`, config metadata, `**/api/failover` with one chain `chain` [a healthy active, b healthy], `**/api/failover/chain`, and `**/api/failover/events` fulfilled with `text/event-stream` body `event: snapshot\ndata: {"chains":[...]}\n\nevent: chain.switched\ndata: {"chain":"chain","from":"a","to":"b","state":"fallback","reason":"trip","at":"2026-09-26T10:00:00Z"}\n\n`). Assertions: `/app/model-editor/chain` shows the strip with the chain state; after the event the active target is `b` and the pill reads fallback; the Pin button is visible with auth disabled (admin) and hidden when auth status reports a non-admin user (mock `/api/auth/me` the way `users-tab-gating.spec.js` does); clicking Pin + confirm POSTs `{target}`. +- [ ] **Step 2** run `cd core/http/react-ui && npx playwright test e2e/failover-health.spec.js` (after `bun run build` if the harness serves the build; follow the repo's UI test instructions in `Makefile` `test-ui*` targets) → FAIL. +- [ ] **Step 3** implement; **Step 4** re-run → PASS; `npm run lint` and `npm run lint:inline-styles` (or the scripts in `package.json`) clean. +- [ ] **Step 5: Commit** — `feat(ui): show live failover chain health in the model editor`. + +--- + +### Task 11: Chain editor field and template + +**Files:** +- Create: `core/http/react-ui/src/components/FailoverTargetsEditor.jsx` +- Modify: `src/components/ConfigFieldRenderer.jsx` (branch `component === 'failover-targets'`, same `list-row` wrapper as `router-candidates`), `src/utils/modelTemplates.js` (template), `src/pages/ModelEditor.jsx` (`SECTION_ICONS.failover = 'fa-shuffle'`, `SECTION_COLORS.failover = 'var(--color-accent)'`), `core/config/meta/registry.go` (`failover.targets` `Component: "failover-targets"`), `core/config/meta/registry_test.go` (assert the component), `public/locales/*/modelEditor.json` +- Test: `e2e/failover-editor.spec.js` + +`FailoverTargetsEditor({ value, onChange })`: modelled on `RouterCandidatesEditor` — items `{model, warm}`; row = `SearchableModelSelect` (value/onChange), `Toggle` for warm (disabled with a title when the selected model's backend is a proxy — look it up from `useModels()` data if it carries the backend; otherwise leave enabled and rely on the load warning, and say so), move up/down, remove; "Add target" button; inline errors (fewer than 2 targets, duplicate model, the edited model's own name) from `useFormContext()` `formData.name`. + +Template entry: +```js +{ + id: 'failover', + label: 'Failover Chain', + icon: 'fa-shuffle', + description: 'Serve one model name from an ordered list of models. The first healthy one answers; the next takes over when it fails.', + fields: { + 'name': '', + 'failover.targets': [{ model: '' }, { model: '' }], + }, +}, +``` + +- [ ] Steps: failing spec (template card visible; `?template=failover` shows two target rows; adding/removing/moving rows; duplicate and too-few errors; saving sends `failover.targets` in the PATCH/import body — mock the save endpoint and assert the JSON) → implement → `npx playwright test e2e/failover-editor.spec.js` PASS, `go test ./core/config/meta/...` PASS → commit `feat(ui): edit failover chain targets with a dedicated field`. + +--- + +### Task 12: Chain badge and Failover overview page + +**Files:** +- Create: `core/http/react-ui/src/pages/Failover.jsx` +- Modify: `src/pages/InstalledModels.jsx` (chain badge next to the alias badge: `badge badge-info`, icon `fa-shuffle`, text `chain → `; data from `useFailoverChains()`), `src/router.jsx` (`const Failover = page('failover', () => import('./pages/Failover'))`; route `{ path: 'failover', element: }`), `src/components/console/consoleConfig.js` (`operate.runtime` item `{ path: '/app/failover', icon: 'fas fa-shuffle', labelKey: 'items.failover', adminOnly: true }`), `public/locales/*/nav.json`, `public/locales/*/models.json`, `src/App.css` +- Test: `e2e/failover-overview.spec.js` + +Overview: `useFailoverChains()`; a dense table (chain name link → `/app/model-editor/`, chain `StatusPill`, active target, one small pill per target, time since `active_since`); empty state with a link to `/app/model-editor?template=failover`. + +- [ ] Steps: failing spec (nav entry visible for admin; table rows from mocked `/api/failover`; SSE patch updates a row; empty state link; installed-models badge `chain → a`) → implement → specs PASS; run the full UI suite with coverage: `make test-ui-coverage-check` (UI coverage ≥ baseline) → commit `feat(ui): list failover chains and badge chain models`. + +--- + +## Part D — Contributor rule and final verification + +### Task 13: distributed-state rule, docs sweep, final verification + +**Files:** +- Create: `.agents/distributed-state.md` +- Modify: `AGENTS.md` (Topics table row; Quick Reference bullet), `.agents/api-endpoints-and-auth.md` (checklist line), `docs/content/features/model-failover.md` (UI section: editor, health strip, overview; confirm distributed section from Task 3) + +`.agents/distributed-state.md` content (write it in full, following the style of the other `.agents/*.md` guides): +- Title "Distributed-aware state". Why: frontends are stateless replicas; in-memory state diverges silently (the failover chains example). +- The rule: a feature that keeps runtime state (in-memory maps, caches, pins, schedulers, background loops, probes) chooses one mode and documents it on its docs page: + 1. **Shared** — `syncstate.SyncedMap` (`core/services/syncstate`); add a `Store` when the state must survive a restart. Example: finetune jobs (`core/services/finetune/service.go`), failover pins (`core/services/failover/distsync`). Gotcha: `Reconcile` without a `Store`/`Loader` re-hydrates the map empty — republish from a leader instead. + 2. **Single-runner** — `advisorylock.RunLeaderLoop` / `TryWithLockCtx` (`core/services/advisorylock`); new keys go in `keys.go`. Example: node health monitor (`core/services/nodes/health.go`), failover prober. + 3. **Stateless per request** — nothing to share. + 4. **Per-instance** — allowed only with the reason written in the feature's docs. +- Tests: shared and single-runner features include a two-instance test on `testutil.NewFakeBus()` (`core/services/testutil/fakebus.go`); note that the fake bus delivers synchronously, including the publisher's own message, so never publish while holding a lock the apply path takes. +- Checklist for PRs. + +AGENTS.md Quick Reference bullet: +`- **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).` + +AGENTS.md Topics row: +`| [.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) |` + +api-endpoints-and-auth.md checklist line (under Quality): +`- [ ] Stateful feature: distributed mode chosen and documented (see [distributed-state.md](distributed-state.md))` + +- [ ] Steps: write the files → final verification, in order, reading each output: +```bash +make protogen-go build-mock-backend build-cloud-proxy-backend build-localai-proxy-backend +go vet ./core/services/failover/... ./backend/go/localai-proxy/... ./pkg/grpc/... +go test -race ./core/services/failover/... ./core/application/... ./core/config/... ./core/http/middleware/... ./pkg/grpc/... ./pkg/mcp/localaitools/... ./backend/go/localai-proxy/... ./core/http/endpoints/openai/... +LOCALAI_TEST_HTTP_PORT=19391 go test ./core/http/... +go run github.com/onsi/ginkgo/v2/ginkgo --label-filter='!real-models' -v ./tests/e2e +make test-ui-coverage-check +LOCALAI_TEST_HTTP_PORT=19391 make test-coverage-check +``` +Expected: all green; both coverage checks at or above baseline. Record any pre-existing failure (e.g. `make swagger`) with its output tail. +- [ ] Commit — `docs: require distributed-aware state for stateful features`. From 65c22c53bc6eb6ab68b25ab1f5d1a6eac7bb7188 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 21:57:02 +0000 Subject: [PATCH 35/79] feat(failover): share state through a sync hook and gate probes on a leader Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/manager.go | 74 +++++++- core/services/failover/schedule.go | 47 ++++- core/services/failover/statesync.go | 218 +++++++++++++++++++++++ core/services/failover/statesync_test.go | 184 +++++++++++++++++++ 4 files changed, 512 insertions(+), 11 deletions(-) create mode 100644 core/services/failover/statesync.go create mode 100644 core/services/failover/statesync_test.go diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 847584d62..1412d7edf 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -60,6 +60,21 @@ type Manager struct { hasChains atomic.Bool // probes counts running probes; only tests wait on it. probes sync.WaitGroup + + // sync shares state with other frontends; nil when standalone. + sync StateSync + // gate grants probing and chain decisions to one frontend; nil means + // this manager is always the leader. + gate LeaderGate + leader bool + // applying is set while a peer's target state is applied, so the + // transition is not published back to the peers. + applying bool + // pending holds publishes queued under the lock; unlockAndFlush runs + // them after unlocking because the sync layer can call back into Apply*. + pending []func() + // ticks counts Run's ticks for the periodic republish; only Run uses it. + ticks int } type targetState struct { @@ -73,6 +88,9 @@ type targetState struct { lastProbe time.Time lastActivity time.Time lastError string + // since and reason describe the last state change, for snapshots. + since time.Time + reason Reason // probing is set while a probe runs, so the scheduler does not start a // second one for the same target. probing bool @@ -90,6 +108,9 @@ type chainState struct { activeSince time.Time pinned string state ChainState + // adopted is set once a follower received the leader's decision for this + // chain; from then on it stops choosing the active target itself. + adopted bool } func New(src ConfigSource, opts ...Option) *Manager { @@ -103,6 +124,7 @@ func New(src ConfigSource, opts ...Option) *Manager { for _, o := range opts { o(m) } + m.leader = m.gate == nil return m } @@ -112,7 +134,7 @@ func (m *Manager) Sync() { m.mu.Lock() m.syncLocked() warm, deliver := m.takeWarmLocked() - m.mu.Unlock() + m.unlockAndFlush() if deliver && m.onWarm != nil { m.onWarm(warm) } @@ -264,7 +286,7 @@ func (m *Manager) targetStates() map[string]TargetState { // the scheduler calls this on every tick. func (m *Manager) Reevaluate() { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() for _, ch := range m.chains { m.recomputeLocked(ch, "") } @@ -277,6 +299,7 @@ func (m *Manager) setTargetLocked(ts *targetState, to TargetState, reason Reason from := ts.state now := m.clock.Now() ts.state = to + ts.since, ts.reason = now, reason switch to { case StateDown: ts.downSince = now @@ -287,11 +310,19 @@ func (m *Manager) setTargetLocked(ts *targetState, to TargetState, reason Reason ts.failures = nil } m.emitLocked(Event{Type: EventTargetState, Target: ts.name, From: string(from), To: string(to), Reason: reason, Error: errMsg, At: now}) + m.queuePublishTargetLocked(ts, from) } // recomputeLocked picks the active target. override replaces the reason of a // resulting switch (pin and unpin are always "manual"). func (m *Manager) recomputeLocked(ch *chainState, override Reason) { + if m.sync != nil && !m.leader && ch.adopted && ch.pinned == "" { + // The leader decides; deciding here too would let frontends serve + // different targets. A pin is exempt: it fixes the active target + // the same way on every frontend, and applying it at once gives the + // caller read-your-writes. + return + } now := m.clock.Now() prev := ch.active next := prev @@ -333,6 +364,7 @@ func (m *Manager) recomputeLocked(ch *chainState, override Reason) { default: state = ChainFallback } + changed := next != prev || state != ch.state switch { case next != prev: ch.active = next @@ -348,6 +380,17 @@ func (m *Manager) recomputeLocked(ch *chainState, override Reason) { m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonRecovery, At: now}) } ch.state = state + if changed { + pub := reason + if next == prev { + // Same reasons as the events above for a state-only change. + pub = ReasonRecovery + if state == ChainDegraded { + pub = ReasonDegraded + } + } + m.queuePublishChainLocked(ch, pub) + } } func (m *Manager) recomputeForLocked(target string) { @@ -373,7 +416,7 @@ type Attempt struct { // order; a pinned chain only the pinned target. func (m *Manager) Plan(chain string) (*Attempt, error) { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() ch := m.chainLocked(chain) if ch == nil { return nil, fmt.Errorf("%w: %q", ErrChainNotFound, chain) @@ -447,7 +490,7 @@ func (a *Attempt) Succeed() { a.m.ReportSuccess(a.Target()) } func (m *Manager) ReportFailure(target string, err error) { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() ts := m.targetLocked(target) if ts == nil { return @@ -483,7 +526,7 @@ func (m *Manager) recordFailureLocked(ts *targetState, msg string) { func (m *Manager) ReportSuccess(target string) { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() ts := m.targetLocked(target) if ts == nil { return @@ -515,37 +558,50 @@ func (m *Manager) recordPassLocked(ts *targetState) { } } +// Pin takes effect here at once (read-your-writes), then is shared with the +// other frontends. func (m *Manager) Pin(chain, target string) error { m.mu.Lock() - defer m.mu.Unlock() ch := m.chainLocked(chain) if ch == nil { + m.unlockAndFlush() return fmt.Errorf("%w: %q", ErrChainNotFound, chain) } if !slices.Contains(ch.targets, target) { + m.unlockAndFlush() return fmt.Errorf("%w: %q", ErrTargetNotInChain, target) } ch.pinned = target m.recomputeLocked(ch, ReasonManual) + s := m.sync + m.unlockAndFlush() + if s != nil { + return s.SetPin(chain, target) + } return nil } func (m *Manager) Unpin(chain string) error { m.mu.Lock() - defer m.mu.Unlock() ch := m.chainLocked(chain) if ch == nil { + m.unlockAndFlush() return fmt.Errorf("%w: %q", ErrChainNotFound, chain) } ch.pinned = "" m.recomputeLocked(ch, ReasonManual) + s := m.sync + m.unlockAndFlush() + if s != nil { + return s.ClearPin(chain) + } return nil } // Status returns every chain, sorted by name. func (m *Manager) Status() []ChainStatus { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() m.syncLocked() names := make([]string, 0, len(m.chains)) for name := range m.chains { @@ -561,7 +617,7 @@ func (m *Manager) Status() []ChainStatus { func (m *Manager) ChainStatus(name string) (ChainStatus, bool) { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() ch := m.chainLocked(name) if ch == nil { return ChainStatus{}, false diff --git a/core/services/failover/schedule.go b/core/services/failover/schedule.go index 08f138d00..739c48b42 100644 --- a/core/services/failover/schedule.go +++ b/core/services/failover/schedule.go @@ -22,6 +22,12 @@ func (m *Manager) Run(ctx context.Context) { return case <-ticker.C: m.Tick(ctx) + // Publishes are fire-and-forget, so a frontend that missed one + // (restart, dropped message) converges within ten seconds. + m.ticks++ + if m.ticks%10 == 0 && m.IsLeader() { + m.Republish() + } } } } @@ -34,8 +40,45 @@ func (m *Manager) Run(ctx context.Context) { // hangs until its timeout) must not delay probing and fail-back of every // other chain. Each probe applies its own result, and a target whose probe is // still running is skipped until it ends. +// +// With a leader gate, only the leader probes and decides chains; followers +// still recompute, which only moves chains that have not yet adopted a +// leader decision. func (m *Manager) Tick(ctx context.Context) { m.Sync() + if m.gate == nil { + m.lead(ctx) + return + } + if m.gate(ctx, func() { m.lead(ctx) }) { + return + } + m.mu.Lock() + m.leader = false + m.mu.Unlock() + m.Reevaluate() +} + +// lead is the leader's share of a tick. +func (m *Manager) lead(ctx context.Context) { + m.mu.Lock() + became := !m.leader + m.leader = true + if became { + // The previous leader owned the warm-set callback's effects; + // deliver the set again so this frontend takes them over. + m.warmPending = true + } + warm, deliver := m.takeWarmLocked() + m.unlockAndFlush() + if deliver && m.onWarm != nil { + m.onWarm(warm) + } + if became { + // Followers hold the old leader's view; send ours at once instead of + // letting them wait for the periodic republish. + m.Republish() + } for _, j := range m.dueProbes() { m.probes.Add(1) go func(j probeJob) { @@ -61,7 +104,7 @@ type probeJob struct { func (m *Manager) dueProbes() []probeJob { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() now := m.clock.Now() var jobs []probeJob for _, ts := range m.targets { @@ -139,7 +182,7 @@ func (m *Manager) endProbe(target string) { func (m *Manager) applyProbe(j probeJob, err error) { m.mu.Lock() - defer m.mu.Unlock() + defer m.unlockAndFlush() ts := m.targets[j.target] if ts == nil { return diff --git a/core/services/failover/statesync.go b/core/services/failover/statesync.go new file mode 100644 index 000000000..d63a01b96 --- /dev/null +++ b/core/services/failover/statesync.go @@ -0,0 +1,218 @@ +package failover + +import ( + "context" + "slices" + "time" + + "github.com/mudler/xlog" +) + +// TargetSnapshot is one target's health as shared between frontends. +type TargetSnapshot struct { + Target string `json:"target"` + State TargetState `json:"state"` + Reason Reason `json:"reason"` + Error string `json:"error,omitempty"` + ConsecutiveOK int `json:"consecutive_ok"` + Since time.Time `json:"since"` +} + +// ChainSnapshot is one chain's active target as decided by the leader. +type ChainSnapshot struct { + Chain string `json:"chain"` + Active string `json:"active"` + ActiveSince time.Time `json:"active_since"` + State ChainState `json:"state"` + Reason Reason `json:"reason"` +} + +// StateSync shares failover state between frontends. Implementations may +// deliver a publish back to the publisher synchronously (NATS echoes), so the +// manager never calls it while holding its lock. +type StateSync interface { + PublishTarget(TargetSnapshot) + PublishChain(ChainSnapshot) + SetPin(chain, target string) error + ClearPin(chain string) error + Pins() map[string]string +} + +// LeaderGate runs fn only on the one frontend that holds leadership and +// reports whether it did. Probing and chain decisions happen on the leader +// only, so N frontends do not probe every target N times or disagree on the +// active target. +type LeaderGate func(ctx context.Context, fn func()) bool + +// WithLeaderGate makes the manager probe and decide chains only while the gate +// grants leadership. Without it the manager is always the leader. +func WithLeaderGate(g LeaderGate) Option { return func(m *Manager) { m.gate = g } } + +// SetStateSync attaches the sync layer and hydrates pins from it. Pins are +// read outside the lock because the store may need I/O. +func (m *Manager) SetStateSync(s StateSync) { + m.mu.Lock() + m.sync = s + m.mu.Unlock() + if s == nil { + return + } + for chain, target := range s.Pins() { + m.ApplyPin(chain, target) + } +} + +// IsLeader reports whether this manager probed and decided chains at its last +// tick. A standalone manager (no gate) is always the leader. +func (m *Manager) IsLeader() bool { + m.mu.Lock() + defer m.mu.Unlock() + return m.leader +} + +// unlockAndFlush releases the lock and then runs the publishes queued while +// it was held: the sync layer may call straight back into Apply*, which takes +// the lock again. +func (m *Manager) unlockAndFlush() { + pending := m.pending + m.pending = nil + m.mu.Unlock() + for _, f := range pending { + f() + } +} + +func (m *Manager) targetSnapshotLocked(ts *targetState) TargetSnapshot { + since := ts.since + if ts.state == StateDown { + since = ts.downSince + } + return TargetSnapshot{ + Target: ts.name, State: ts.state, Reason: ts.reason, Error: ts.lastError, + ConsecutiveOK: ts.consecutiveOK, Since: since, + } +} + +func (m *Manager) chainSnapshotLocked(ch *chainState, reason Reason) ChainSnapshot { + return ChainSnapshot{ + Chain: ch.name, Active: ch.targets[ch.active], ActiveSince: ch.activeSince, + State: ch.state, Reason: reason, + } +} + +// queuePublishTargetLocked shares a local target transition. Missing is a fact +// about this frontend's config, not about the target, so it is never shared. +func (m *Manager) queuePublishTargetLocked(ts *targetState, from TargetState) { + if m.sync == nil || m.applying || ts.state == StateMissing || from == StateMissing { + return + } + s, snap := m.sync, m.targetSnapshotLocked(ts) + m.pending = append(m.pending, func() { s.PublishTarget(snap) }) +} + +func (m *Manager) queuePublishChainLocked(ch *chainState, reason Reason) { + if m.sync == nil || !m.leader { + return + } + s, snap := m.sync, m.chainSnapshotLocked(ch, reason) + m.pending = append(m.pending, func() { s.PublishChain(snap) }) +} + +// ApplyTarget takes a target state published by any frontend, this one +// included. The echo of an own publish finds the same state and does nothing. +func (m *Manager) ApplyTarget(s TargetSnapshot) { + m.mu.Lock() + defer m.unlockAndFlush() + ts := m.targetLocked(s.Target) + if ts == nil || ts.state == StateMissing || s.State == StateMissing { + return + } + if ts.state == s.State { + return + } + m.applying = true + m.setTargetLocked(ts, s.State, s.Reason, s.Error) + m.applying = false + // Keep the publisher's clock, so dwell timers (cold recovery, fail-back) + // run from when the target actually changed, not from when we heard. + if !s.Since.IsZero() { + ts.since = s.Since + if s.State == StateDown { + ts.downSince = s.Since + } + } + ts.consecutiveOK = s.ConsecutiveOK + if s.Error != "" { + ts.lastError = s.Error + } + m.recomputeForLocked(ts.name) +} + +// ApplyChain adopts the leader's decision for a chain. The leader ignores it: +// it is the source of these decisions, and a late publish from a previous +// leader must not undo its own. +func (m *Manager) ApplyChain(s ChainSnapshot) { + m.mu.Lock() + defer m.unlockAndFlush() + if m.leader { + return + } + ch := m.chainLocked(s.Chain) + if ch == nil { + return + } + next := slices.Index(ch.targets, s.Active) + if next < 0 { + // The frontends disagree on the chain's targets while a config + // change propagates; keep deciding locally until they agree. + xlog.Debug("failover: ignoring chain state for an unknown target", "chain", s.Chain, "target", s.Active) + return + } + prev := ch.active + changed := next != prev || s.State != ch.state + ch.active, ch.activeSince, ch.state, ch.adopted = next, s.ActiveSince, s.State, true + if changed { + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(s.State), Reason: s.Reason, At: m.clock.Now()}) + } +} + +// ApplyPin sets (target != "") or clears a pin set on any frontend. +func (m *Manager) ApplyPin(chain, target string) { + m.mu.Lock() + defer m.unlockAndFlush() + ch := m.chainLocked(chain) + if ch == nil { + return + } + if target != "" && !slices.Contains(ch.targets, target) { + xlog.Debug("failover: ignoring pin to a target not in the chain", "chain", chain, "target", target) + return + } + if ch.pinned == target { + return + } + ch.pinned = target + m.recomputeLocked(ch, ReasonManual) +} + +// Republish sends every target and chain state, so frontends that missed a +// publish (joined late, dropped a message) converge. Only the leader's view +// is authoritative, so followers do nothing. +func (m *Manager) Republish() { + m.mu.Lock() + defer m.unlockAndFlush() + if m.sync == nil || !m.leader { + return + } + s := m.sync + for _, ts := range m.targets { + if ts.state == StateMissing { + continue + } + snap := m.targetSnapshotLocked(ts) + m.pending = append(m.pending, func() { s.PublishTarget(snap) }) + } + for _, ch := range m.chains { + m.queuePublishChainLocked(ch, ReasonInitial) + } +} diff --git a/core/services/failover/statesync_test.go b/core/services/failover/statesync_test.go new file mode 100644 index 000000000..7e6f7aa62 --- /dev/null +++ b/core/services/failover/statesync_test.go @@ -0,0 +1,184 @@ +package failover + +import ( + "context" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// loopSync is an in-process StateSync that delivers every publish to all +// managers synchronously, including the publisher (like NATS echo). +type loopSync struct { + mu sync.Mutex + peers []*Manager + pins map[string]string +} + +func (l *loopSync) add(m *Manager) { l.mu.Lock(); l.peers = append(l.peers, m); l.mu.Unlock() } +func (l *loopSync) each(f func(*Manager)) { + l.mu.Lock() + ps := append([]*Manager(nil), l.peers...) + l.mu.Unlock() + for _, p := range ps { + f(p) + } +} +func (l *loopSync) PublishTarget(s TargetSnapshot) { l.each(func(m *Manager) { m.ApplyTarget(s) }) } +func (l *loopSync) PublishChain(s ChainSnapshot) { l.each(func(m *Manager) { m.ApplyChain(s) }) } +func (l *loopSync) SetPin(c, t string) error { + l.mu.Lock() + l.pins[c] = t + l.mu.Unlock() + l.each(func(m *Manager) { m.ApplyPin(c, t) }) + return nil +} +func (l *loopSync) ClearPin(c string) error { + l.mu.Lock() + delete(l.pins, c) + l.mu.Unlock() + l.each(func(m *Manager) { m.ApplyPin(c, "") }) + return nil +} +func (l *loopSync) Pins() map[string]string { + l.mu.Lock() + defer l.mu.Unlock() + out := map[string]string{} + for k, v := range l.pins { + out[k] = v + } + return out +} + +var _ = Describe("Manager state sync", func() { + var ( + clock *fakeClock + src *fakeSource + bus *loopSync + a, b *Manager + leaderIsA bool + gateFor func(isA bool) LeaderGate + ctx = context.Background() + ) + + BeforeEach(func() { + clock = newFakeClock() + src = newFakeSource(remote("x"), local("y"), chainCfg("chain", nil, t("x"), t("y"))) + bus = &loopSync{pins: map[string]string{}} + leaderIsA = true + gateFor = func(isA bool) LeaderGate { + return func(_ context.Context, fn func()) bool { + if isA != leaderIsA { + return false + } + fn() + return true + } + } + a = New(src, WithClock(clock), WithLeaderGate(gateFor(true))) + b = New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + bus.add(a) + bus.add(b) + a.SetStateSync(bus) + b.SetStateSync(bus) + a.Tick(ctx) + b.Tick(ctx) + }) + + It("echo of own publish is a no-op and emits one event", func() { + events, cancel := a.Subscribe(16) + defer cancel() + a.ReportFailure("x", errBoom) + n := 0 + for _, e := range drain(events) { + if e.Type == EventTargetState && e.Target == "x" { + n++ + } + } + Expect(n).To(Equal(1)) + }) + + It("a trip on one frontend is skipped by the other's plan", func() { + b.ReportFailure("x", errBoom) + att, err := a.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("y")) + }) + + It("followers adopt the leader's chain state and emit the switch", func() { + events, cancel := b.Subscribe(16) + defer cancel() + a.ReportFailure("x", errBoom) // leader recomputes and publishes chain state + st, _ := b.ChainStatus("chain") + Expect(st.Active).To(Equal("y")) + var sw []Event + for _, e := range drain(events) { + if e.Type == EventChainSwitched { + sw = append(sw, e) + } + } + Expect(sw).ToNot(BeEmpty()) + }) + + It("a pin on one frontend applies on all", func() { + Expect(b.Pin("chain", "y")).To(Succeed()) + st, _ := a.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + Expect(a.Unpin("chain")).To(Succeed()) + st, _ = b.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + + It("hydrates pins when the sync is attached", func() { + bus.pins["chain"] = "y" + c := New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + c.SetStateSync(bus) + st, _ := c.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + }) + + It("only the leader probes", func() { + pa, pb := &fakeProber{fail: map[string]error{}}, &fakeProber{fail: map[string]error{}} + a = New(src, WithClock(clock), WithProber(pa), WithLeaderGate(gateFor(true))) + b = New(src, WithClock(clock), WithProber(pb), WithLeaderGate(gateFor(false))) + a.SetStateSync(bus) + b.SetStateSync(bus) + a.Tick(ctx) + b.Tick(ctx) + Eventually(func() int { return len(pa.take()) }).Should(BeNumerically(">", 0)) + Consistently(func() int { return len(pb.take()) }, 200*time.Millisecond).Should(Equal(0)) + Expect(a.IsLeader()).To(BeTrue()) + Expect(b.IsLeader()).To(BeFalse()) + }) + + It("new leader keeps activeSince across a leadership move", func() { + a.ReportFailure("x", errBoom) + before, _ := b.ChainStatus("chain") + leaderIsA = false + clock.Advance(5 * time.Second) + a.Tick(ctx) + b.Tick(ctx) + after, _ := b.ChainStatus("chain") + Expect(after.ActiveSince).To(Equal(before.ActiveSince)) + Expect(b.IsLeader()).To(BeTrue()) + }) + + It("republish sends every target and chain", func() { + c := New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + bus.add(c) + c.SetStateSync(bus) + a.ReportFailure("x", errBoom) + a.Republish() + st, _ := c.ChainStatus("chain") + Expect(st.Active).To(Equal("y")) + }) + + It("standalone manager (no sync, no gate) is always leader", func() { + m := New(src, WithClock(clock)) + m.Tick(ctx) + Expect(m.IsLeader()).To(BeTrue()) + }) +}) From 5f98dfcf95fe583c10c9e7f72e4b3d0612114111 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:01:13 +0000 Subject: [PATCH 36/79] fix(failover): keep shared pins for chains this frontend has not loaded yet A pin arrives from the sync layer once. Dropping it when the chain or target is unknown here left this frontend routing differently from the cluster whenever its config lagged or a chain was re-created. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/manager.go | 34 +++++++++---- core/services/failover/statesync.go | 11 ++++- core/services/failover/statesync_test.go | 61 ++++++++++++++++++++++++ 3 files changed, 96 insertions(+), 10 deletions(-) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 1412d7edf..da5ba6079 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -43,11 +43,16 @@ func WithOnWarmChanged(fn func(warm []string)) Option { return func(m *Manager) // Manager tracks health per target and the active target per chain. type Manager struct { - mu sync.Mutex - src ConfigSource - clock Clock - prober Prober - onWarm func([]string) + mu sync.Mutex + src ConfigSource + clock Clock + prober Prober + onWarm func([]string) + // pins holds every known pin by chain name, including pins for chains + // or targets this frontend's config does not have yet: with a sync + // layer the pin is delivered once, and a chain that appears (or is + // rebuilt) later must still pick it up. + pins map[string]string targets map[string]*targetState chains map[string]*chainState subs map[int]chan Event @@ -119,6 +124,7 @@ func New(src ConfigSource, opts ...Option) *Manager { clock: realClock{}, targets: map[string]*targetState{}, chains: map[string]*chainState{}, + pins: map[string]string{}, subs: map[int]chan Event{}, } for _, o := range opts { @@ -155,9 +161,14 @@ func (m *Manager) syncLocked() { } ch := m.chains[c.Name] if ch == nil || !slices.Equal(ch.targets, names) { - pinned := "" - if ch != nil && slices.Contains(names, ch.pinned) { - pinned = ch.pinned + pinned := m.pins[c.Name] + if !slices.Contains(names, pinned) { + if m.sync == nil { + // Standalone, the pin lives with the chain: a target + // dropped from it takes the pin along. + delete(m.pins, c.Name) + } + pinned = "" } ch = &chainState{name: c.Name, targets: names, activeSince: now, state: ChainPrimary, pinned: pinned} m.chains[c.Name] = ch @@ -191,6 +202,11 @@ func (m *Manager) syncLocked() { for name := range m.chains { if !seenChains[name] { delete(m.chains, name) + if m.sync == nil { + // With a sync layer the shared pin outlives a chain this + // frontend has not (re)loaded yet; standalone it does not. + delete(m.pins, name) + } } } for name := range m.targets { @@ -572,6 +588,7 @@ func (m *Manager) Pin(chain, target string) error { return fmt.Errorf("%w: %q", ErrTargetNotInChain, target) } ch.pinned = target + m.pins[chain] = target m.recomputeLocked(ch, ReasonManual) s := m.sync m.unlockAndFlush() @@ -589,6 +606,7 @@ func (m *Manager) Unpin(chain string) error { return fmt.Errorf("%w: %q", ErrChainNotFound, chain) } ch.pinned = "" + delete(m.pins, chain) m.recomputeLocked(ch, ReasonManual) s := m.sync m.unlockAndFlush() diff --git a/core/services/failover/statesync.go b/core/services/failover/statesync.go index d63a01b96..bcab901ef 100644 --- a/core/services/failover/statesync.go +++ b/core/services/failover/statesync.go @@ -176,16 +176,23 @@ func (m *Manager) ApplyChain(s ChainSnapshot) { } } -// ApplyPin sets (target != "") or clears a pin set on any frontend. +// ApplyPin sets (target != "") or clears a pin set on any frontend. A pin for +// a chain or target this frontend does not know yet is kept and applied by +// syncLocked once the config catches up. func (m *Manager) ApplyPin(chain, target string) { m.mu.Lock() defer m.unlockAndFlush() + if target == "" { + delete(m.pins, chain) + } else { + m.pins[chain] = target + } ch := m.chainLocked(chain) if ch == nil { return } if target != "" && !slices.Contains(ch.targets, target) { - xlog.Debug("failover: ignoring pin to a target not in the chain", "chain", chain, "target", target) + xlog.Debug("failover: deferring pin to a target not in the chain yet", "chain", chain, "target", target) return } if ch.pinned == target { diff --git a/core/services/failover/statesync_test.go b/core/services/failover/statesync_test.go index 7e6f7aa62..3901d3b95 100644 --- a/core/services/failover/statesync_test.go +++ b/core/services/failover/statesync_test.go @@ -176,6 +176,67 @@ var _ = Describe("Manager state sync", func() { Expect(st.Active).To(Equal("y")) }) + It("keeps a pin for a chain it does not know yet and applies it when the chain appears", func() { + b.ApplyPin("later", "y") + _, ok := b.ChainStatus("later") + Expect(ok).To(BeFalse()) + src.Put(chainCfg("later", nil, t("x"), t("y"))) + b.Sync() + st, ok := b.ChainStatus("later") + Expect(ok).To(BeTrue()) + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + Expect(st.Active).To(Equal("y")) + }) + + It("re-applies a shared pin when the chain is removed and re-added", func() { + Expect(a.Pin("chain", "y")).To(Succeed()) + src.Delete("chain") + b.Sync() + _, ok := b.ChainStatus("chain") + Expect(ok).To(BeFalse()) + src.Put(chainCfg("chain", nil, t("x"), t("y"))) + b.Sync() + st, _ := b.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + }) + + It("applies a deferred pin once its target joins the chain", func() { + src.Put(local("z")) + b.ApplyPin("chain", "z") + st, _ := b.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + src.Put(chainCfg("chain", nil, t("x"), t("y"), t("z"))) + b.Sync() + st, _ = b.ChainStatus("chain") + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("z")) + }) + + It("hydrates a pin for an unknown chain and applies it when the chain appears", func() { + bus.pins["later"] = "y" + c := New(src, WithClock(clock), WithLeaderGate(gateFor(false))) + c.SetStateSync(bus) + src.Put(chainCfg("later", nil, t("x"), t("y"))) + c.Sync() + st, ok := c.ChainStatus("later") + Expect(ok).To(BeTrue()) + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + }) + + It("standalone, a removed chain drops its pin", func() { + m := New(src, WithClock(clock)) + Expect(m.Pin("chain", "y")).To(Succeed()) + src.Delete("chain") + m.Sync() + src.Put(chainCfg("chain", nil, t("x"), t("y"))) + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + It("standalone manager (no sync, no gate) is always leader", func() { m := New(src, WithClock(clock)) m.Tick(ctx) From 7b88674fc58c67cb796729210af10062d539c6d3 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:11:54 +0000 Subject: [PATCH 37/79] feat(failover): sync pins, target health and chain state over NATS Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/distsync/distsync.go | 158 +++++++++++ .../failover/distsync/distsync_suite_test.go | 13 + .../failover/distsync/distsync_test.go | 258 ++++++++++++++++++ core/services/failover/distsync/pinstore.go | 63 +++++ 4 files changed, 492 insertions(+) create mode 100644 core/services/failover/distsync/distsync.go create mode 100644 core/services/failover/distsync/distsync_suite_test.go create mode 100644 core/services/failover/distsync/distsync_test.go create mode 100644 core/services/failover/distsync/pinstore.go diff --git a/core/services/failover/distsync/distsync.go b/core/services/failover/distsync/distsync.go new file mode 100644 index 000000000..349aef14f --- /dev/null +++ b/core/services/failover/distsync/distsync.go @@ -0,0 +1,158 @@ +package distsync + +import ( + "context" + "errors" + "fmt" + "reflect" + "time" + + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/syncstate" + "github.com/mudler/xlog" +) + +// compile-time assertions. +var ( + _ syncstate.Store[string, PinRecord] = (*PinStore)(nil) + _ failover.StateSync = (*Sync)(nil) +) + +// Sync is a failover.StateSync backed by three syncstate.SyncedMaps: +// +// - failover.pins keeps the durable source of truth: Store-backed (when a +// PinStore is given) so a late joiner hydrates every pin from the DB on +// Start, not just from whatever peers happen to broadcast afterwards. +// - failover.targets and failover.chains are ephemeral live-health state, +// NATS-only with no Store and no Reconcile: there is nothing durable to +// hydrate from, and a Reconcile tick would re-hydrate them empty and wipe +// live state clean off a running frontend. A late joiner instead catches +// up from the leader's periodic Republish (see failover.Manager.Republish). +type Sync struct { + pins *syncstate.SyncedMap[string, PinRecord] + targets *syncstate.SyncedMap[string, failover.TargetSnapshot] + chains *syncstate.SyncedMap[string, failover.ChainSnapshot] +} + +// New builds and starts the three maps, then attaches the result to m via +// SetStateSync so any already-durable pins hydrate onto m immediately. +func New(ctx context.Context, nats messaging.MessagingClient, pins syncstate.Store[string, PinRecord], m *failover.Manager) (*Sync, error) { + s := &Sync{} + + // pins is already typed as the Store interface (the brief fixes this + // signature), so a caller holding a nil *PinStore (e.g. standalone mode, + // no DB configured) and passing it straight through boxes it into a + // non-nil interface wrapping a nil pointer - "pins != nil" alone would + // not catch that, and the SyncedMap would then try to hydrate/write + // through a nil *gorm.DB. isNilStore catches both that and a literal nil + // argument, the same defense finetune/service.go gets for free by taking + // a concrete *distributed.FineTuneStore and nil-checking before boxing it. + var pinStore syncstate.Store[string, PinRecord] + if !isNilStore(pins) { + pinStore = pins + } + + s.pins = syncstate.New(syncstate.Config[string, PinRecord]{ + Name: "failover.pins", + Key: func(p PinRecord) string { return p.Chain }, + Nats: nats, + Store: pinStore, + OnApply: func(op string, chain string, v PinRecord) { + if op == "delete" { + m.ApplyPin(chain, "") + return + } + m.ApplyPin(chain, v.Target) + }, + }) + if err := s.pins.Start(ctx); err != nil { + return nil, fmt.Errorf("distsync: starting pins map: %w", err) + } + + s.targets = syncstate.New(syncstate.Config[string, failover.TargetSnapshot]{ + Name: "failover.targets", + Key: func(t failover.TargetSnapshot) string { return t.Target }, + Nats: nats, + OnApply: func(_ string, _ string, v failover.TargetSnapshot) { + m.ApplyTarget(v) + }, + }) + if err := s.targets.Start(ctx); err != nil { + _ = s.pins.Close() + return nil, fmt.Errorf("distsync: starting targets map: %w", err) + } + + s.chains = syncstate.New(syncstate.Config[string, failover.ChainSnapshot]{ + Name: "failover.chains", + Key: func(c failover.ChainSnapshot) string { return c.Chain }, + Nats: nats, + OnApply: func(_ string, _ string, v failover.ChainSnapshot) { + m.ApplyChain(v) + }, + }) + if err := s.chains.Start(ctx); err != nil { + _ = s.targets.Close() + _ = s.pins.Close() + return nil, fmt.Errorf("distsync: starting chains map: %w", err) + } + + m.SetStateSync(s) + return s, nil +} + +// isNilStore reports whether pins is nil - either a literal nil argument, or +// the classic Go footgun of a non-nil interface value wrapping a nil +// pointer (e.g. a nil *PinStore passed in directly). Both must disable the +// durable Store the same way, so syncstate hydrates from nothing rather than +// panicking on a nil *gorm.DB the first time it dereferences it. +func isNilStore(pins syncstate.Store[string, PinRecord]) bool { + if pins == nil { + return true + } + v := reflect.ValueOf(pins) + return v.Kind() == reflect.Ptr && v.IsNil() +} + +// Close releases all three maps' subscriptions and background workers. +func (s *Sync) Close() error { + return errors.Join(s.chains.Close(), s.targets.Close(), s.pins.Close()) +} + +// PublishTarget shares a target health transition. StateSync's methods +// return no error to the manager, and a publish must never block it (the +// manager calls this outside its lock precisely so a synchronous NATS echo +// is safe) - so a failure here is logged and dropped; a missed publish +// self-heals on the leader's next Republish. +func (s *Sync) PublishTarget(t failover.TargetSnapshot) { + if err := s.targets.Set(context.Background(), t); err != nil { + xlog.Warn("distsync: publishing target state failed", "target", t.Target, "error", err) + } +} + +// PublishChain shares the leader's decision for a chain. +func (s *Sync) PublishChain(c failover.ChainSnapshot) { + if err := s.chains.Set(context.Background(), c); err != nil { + xlog.Warn("distsync: publishing chain state failed", "chain", c.Chain, "error", err) + } +} + +// SetPin durably persists and broadcasts a pin. +func (s *Sync) SetPin(chain, target string) error { + return s.pins.Set(context.Background(), PinRecord{Chain: chain, Target: target, UpdatedAt: time.Now()}) +} + +// ClearPin durably removes and broadcasts a pin's removal. +func (s *Sync) ClearPin(chain string) error { + return s.pins.Delete(context.Background(), chain) +} + +// Pins returns every known pin (chain -> target), for Manager.SetStateSync's +// hydrate-on-attach. +func (s *Sync) Pins() map[string]string { + out := make(map[string]string) + for chain, rec := range s.pins.Snapshot() { + out[chain] = rec.Target + } + return out +} diff --git a/core/services/failover/distsync/distsync_suite_test.go b/core/services/failover/distsync/distsync_suite_test.go new file mode 100644 index 000000000..37a6d7092 --- /dev/null +++ b/core/services/failover/distsync/distsync_suite_test.go @@ -0,0 +1,13 @@ +package distsync_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestDistsync(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Distsync test suite") +} diff --git a/core/services/failover/distsync/distsync_test.go b/core/services/failover/distsync/distsync_test.go new file mode 100644 index 000000000..57a205f38 --- /dev/null +++ b/core/services/failover/distsync/distsync_test.go @@ -0,0 +1,258 @@ +package distsync_test + +import ( + "context" + "errors" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/core/services/failover/distsync" + "github.com/mudler/LocalAI/core/services/testutil" +) + +// fakeSource is a minimal failover.ConfigSource with one chain "chain" of two +// targets: "x" (remote, primary) and "y" (local warm, fallback). It is +// shared across the managers in a spec, mirroring how frontends in a real +// deployment read the same config loader. +type fakeSource struct { + mu sync.Mutex + cfgs map[string]config.ModelConfig +} + +func newChainSource() *fakeSource { + return &fakeSource{cfgs: map[string]config.ModelConfig{ + "x": {Name: "x", Backend: "cloud-proxy"}, + "y": {Name: "y", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{ + {Model: "x"}, + {Model: "y", Warm: true}, + }}}, + }} +} + +func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { + s.mu.Lock() + defer s.mu.Unlock() + c, ok := s.cfgs[n] + return c, ok +} + +func (s *fakeSource) GetAllModelsConfigs() []config.ModelConfig { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]config.ModelConfig, 0, len(s.cfgs)) + for _, c := range s.cfgs { + out = append(out, c) + } + return out +} + +// memPinStore is an in-memory syncstate.Store[string, distsync.PinRecord] +// shared by several distsync.Sync instances the way a real DB would be, so a +// spec can build a "late joiner" that hydrates from what earlier instances +// already wrote through. +type memPinStore struct { + mu sync.Mutex + data map[string]distsync.PinRecord +} + +func newMemPinStore() *memPinStore { return &memPinStore{data: map[string]distsync.PinRecord{}} } + +func (s *memPinStore) List(context.Context) ([]distsync.PinRecord, error) { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]distsync.PinRecord, 0, len(s.data)) + for _, v := range s.data { + out = append(out, v) + } + return out, nil +} + +func (s *memPinStore) Upsert(_ context.Context, v distsync.PinRecord) error { + s.mu.Lock() + defer s.mu.Unlock() + s.data[v.Chain] = v + return nil +} + +func (s *memPinStore) Delete(_ context.Context, k string) error { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.data, k) + return nil +} + +// leaderGateFor grants leadership to exactly one named manager at a time +// (whatever *leader currently holds), mirroring an advisory-lock leader loop +// where only one frontend probes and decides chains. +func leaderGateFor(name string, leader *string) failover.LeaderGate { + return func(_ context.Context, fn func()) bool { + if *leader != name { + return false + } + fn() + return true + } +} + +var errBoom = errors.New("boom") + +var _ = Describe("distsync", func() { + var ( + ctx context.Context + bus *testutil.FakeBus + src *fakeSource + pinStore *memPinStore + leader string + ) + + BeforeEach(func() { + ctx = context.Background() + bus = testutil.NewFakeBus() + src = newChainSource() + pinStore = newMemPinStore() + leader = "a" + }) + + // newManager wires a fresh failover.Manager to a fresh distsync.Sync on + // the shared bus and pin store, named so leaderGateFor can grant or deny + // it leadership. + newManager := func(name string) (*failover.Manager, *distsync.Sync) { + m := failover.New(src, failover.WithLeaderGate(leaderGateFor(name, &leader))) + s, err := distsync.New(ctx, bus, pinStore, m) + Expect(err).ToNot(HaveOccurred()) + return m, s + } + + It("a pin on A is visible on B and survives a new instance C built from the same store", func() { + a, _ := newManager("a") + b, _ := newManager("b") + a.Tick(ctx) + b.Tick(ctx) + + Expect(a.Pin("chain", "y")).To(Succeed()) + + stB, ok := b.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(stB.Pinned).ToNot(BeNil()) + Expect(*stB.Pinned).To(Equal("y")) + + // C is built after the pin was already written through to the shared + // store, and before ever ticking: SetStateSync (inside distsync.New) + // hydrates C's pins from s.Pins(), so the chain is created pinned the + // first time anything asks for it. + c, _ := newManager("c") + stC, ok := c.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(stC.Pinned).ToNot(BeNil()) + Expect(*stC.Pinned).To(Equal("y")) + }) + + It("a trip on B makes A's plan skip the target", func() { + a, _ := newManager("a") + b, _ := newManager("b") + a.Tick(ctx) + b.Tick(ctx) + + b.ReportFailure("x", errBoom) + + att, err := a.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("y")) + }) + + It("the leader's chain switch reaches the follower", func() { + a, _ := newManager("a") // leader + b, _ := newManager("b") // follower + a.Tick(ctx) + b.Tick(ctx) + + a.ReportFailure("x", errBoom) + + st, ok := b.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(st.Active).To(Equal("y")) + }) + + It("late joiner converges on heartbeat", func() { + a, _ := newManager("a") // leader + a.Tick(ctx) + + a.ReportFailure("x", errBoom) + + // C joins after the trip: failover.targets/chains are NATS-only with + // no Store, so C's hydrate on Start sees nothing and it starts out + // believing every target is healthy. + c, _ := newManager("c") // follower + c.Tick(ctx) + + before, ok := c.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(before.Active).To(Equal("x")) + + a.Republish() + + after, ok := c.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(after.Active).To(Equal("y")) + }) + + It("guards against a typed-nil PinStore passed as the Store interface", func() { + var nilStore *distsync.PinStore // deliberately typed, deliberately nil + m := failover.New(src, failover.WithLeaderGate(leaderGateFor("a", &leader))) + _, err := distsync.New(ctx, bus, nilStore, m) + Expect(err).ToNot(HaveOccurred()) + m.Tick(ctx) + + Expect(m.Pin("chain", "y")).To(Succeed()) + st, ok := m.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(st.Pinned).ToNot(BeNil()) + Expect(*st.Pinned).To(Equal("y")) + }) + + It("Close stops all three maps without erroring", func() { + _, s := newManager("a") + Expect(s.Close()).To(Succeed()) + }) + + It("PinStore round-trips through a real sqlite-backed gorm DB", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + + store, err := distsync.NewPinStore(db) + Expect(err).ToNot(HaveOccurred()) + + recs, err := store.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(recs).To(BeEmpty()) + + Expect(store.Upsert(ctx, distsync.PinRecord{Chain: "chain", Target: "y", UpdatedAt: time.Now()})).To(Succeed()) + + recs, err = store.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(recs).To(HaveLen(1)) + Expect(recs[0].Chain).To(Equal("chain")) + Expect(recs[0].Target).To(Equal("y")) + + // Upsert again on the same key updates rather than duplicating. + Expect(store.Upsert(ctx, distsync.PinRecord{Chain: "chain", Target: "x", UpdatedAt: time.Now()})).To(Succeed()) + recs, err = store.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(recs).To(HaveLen(1)) + Expect(recs[0].Target).To(Equal("x")) + + Expect(store.Delete(ctx, "chain")).To(Succeed()) + recs, err = store.List(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(recs).To(BeEmpty()) + }) +}) diff --git a/core/services/failover/distsync/pinstore.go b/core/services/failover/distsync/pinstore.go new file mode 100644 index 000000000..eaa682431 --- /dev/null +++ b/core/services/failover/distsync/pinstore.go @@ -0,0 +1,63 @@ +// Package distsync wires failover.StateSync to syncstate.SyncedMap so target +// health, chain decisions and pins are shared across frontends over NATS, +// with pins durable in a small gorm-backed table. +package distsync + +import ( + "context" + "fmt" + "time" + + "github.com/mudler/LocalAI/core/services/advisorylock" + "gorm.io/gorm" +) + +// PinRecord is the durable form of a chain pin. It doubles as the value type +// of the "failover.pins" SyncedMap, so a hydrate (Store.List) needs no +// conversion and the wire delta carries the exact row. +type PinRecord struct { + Chain string `gorm:"primaryKey" json:"chain"` + Target string `json:"target"` + UpdatedAt time.Time `json:"updated_at"` +} + +// TableName pins the table name independent of the Go type name. +func (PinRecord) TableName() string { return "failover_pins" } + +// PinStore is gorm-backed durable storage for chain pins, implementing +// syncstate.Store[string, PinRecord] (asserted in distsync.go). +type PinStore struct { + db *gorm.DB +} + +// NewPinStore migrates the failover_pins table under the schema-migrate +// advisory lock - the same guard jobs.NewJobStore uses - so several +// frontends starting at once do not race on the migration, and returns a +// ready-to-use store. +func NewPinStore(db *gorm.DB) (*PinStore, error) { + if err := advisorylock.WithLockCtx(context.Background(), db, advisorylock.KeySchemaMigrate, func() error { + return db.AutoMigrate(&PinRecord{}) + }); err != nil { + return nil, fmt.Errorf("distsync: migrating pin table: %w", err) + } + return &PinStore{db: db}, nil +} + +// List returns every durable pin, for hydrate on Start. +func (s *PinStore) List(ctx context.Context) ([]PinRecord, error) { + var out []PinRecord + if err := s.db.WithContext(ctx).Find(&out).Error; err != nil { + return nil, err + } + return out, nil +} + +// Upsert writes a pin through. Save inserts or updates by primary key. +func (s *PinStore) Upsert(ctx context.Context, v PinRecord) error { + return s.db.WithContext(ctx).Save(&v).Error +} + +// Delete removes a pin by chain name. +func (s *PinStore) Delete(ctx context.Context, k string) error { + return s.db.WithContext(ctx).Delete(&PinRecord{Chain: k}).Error +} From d30c33a074cc0b78cc6fb3c8c0daf374de3d0fbb Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:21:01 +0000 Subject: [PATCH 38/79] feat(failover): run one prober per cluster and pin warm targets on workers Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/application.go | 12 +++ core/application/distributed.go | 7 +- core/application/failover.go | 7 ++ core/application/failover_distributed.go | 69 ++++++++++++ core/application/failover_distributed_test.go | 102 ++++++++++++++++++ core/application/startup.go | 6 +- core/services/advisorylock/keys.go | 3 + core/services/failover/schedule.go | 7 +- core/services/failover/statesync.go | 12 +++ core/services/failover/statesync_test.go | 12 +++ docs/content/features/model-failover.md | 22 +++- 11 files changed, 252 insertions(+), 7 deletions(-) create mode 100644 core/application/failover_distributed.go create mode 100644 core/application/failover_distributed_test.go diff --git a/core/application/application.go b/core/application/application.go index f5454aba1..b15d848a0 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -17,6 +17,7 @@ import ( "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" @@ -94,6 +95,10 @@ 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 + // Upgrade checker (background service for detecting backend upgrades) upgradeChecker *UpgradeChecker @@ -515,6 +520,13 @@ 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. + if a.failoverSync != nil { + if closeErr := a.failoverSync.Close(); closeErr != nil { + xlog.Warn("failover: closing state sync", "error", closeErr) + } + } a.distributed.Shutdown() if a.modelLoader != nil { err = a.modelLoader.StopAllGRPC() diff --git a/core/application/distributed.go b/core/application/distributed.go index 4c6214e57..6ba8234cc 100644 --- a/core/application/distributed.go +++ b/core/application/distributed.go @@ -77,7 +77,9 @@ func (ds *DistributedServices) Shutdown() { // Returns nil if distributed mode is not enabled. // configLoader is used by the SmartRouter to compute concurrency-group // anti-affinity at placement time (#9659); it may be nil in tests. -func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader) (*DistributedServices, error) { +// pinned, when set, replaces configLoader as the source of models the router +// and reconciler must keep loaded (it adds warm failover targets). +func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader, pinned nodes.PinnedModelResolver) (*DistributedServices, error) { if !cfg.Distributed.Enabled { return nil, nil } @@ -383,6 +385,9 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade conflictResolver = configLoader pinnedResolver = configLoader } + if pinned != nil { + pinnedResolver = pinned + } modelCleanup := nodes.NewModelCleanupService(registry, remoteUnloader) router := nodes.NewSmartRouter(registry, nodes.SmartRouterOptions{ Unloader: remoteUnloader, diff --git a/core/application/failover.go b/core/application/failover.go index a513c1f10..ffe5a8e51 100644 --- a/core/application/failover.go +++ b/core/application/failover.go @@ -24,8 +24,15 @@ var preloadModelByName = backend.PreloadModelByName // 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 { diff --git a/core/application/failover_distributed.go b/core/application/failover_distributed.go new file mode 100644 index 000000000..7c1b63b9f --- /dev/null +++ b/core/application/failover_distributed.go @@ -0,0 +1,69 @@ +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" + "gorm.io/gorm" +) + +// 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. A lock error counts as "not leader": +// skipping one tick is safe, two leaders are not. +func failoverLeaderGate(db *gorm.DB) failover.LeaderGate { + return func(ctx context.Context, fn func()) bool { + ok, err := advisorylock.TryWithLockCtx(ctx, db, advisorylock.KeyFailoverProber, func() error { + fn() + return nil + }) + if err != nil { + xlog.Warn("failover: could not take the prober leader lock", "error", err) + return false + } + return ok + } +} + +// 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.failoverManager.SetLeaderGate(failoverLeaderGate(db)) +} diff --git a/core/application/failover_distributed_test.go b/core/application/failover_distributed_test.go new file mode 100644 index 000000000..b4cf84bf0 --- /dev/null +++ b/core/application/failover_distributed_test.go @@ -0,0 +1,102 @@ +package application + +import ( + "context" + "time" + + "github.com/mudler/LocalAI/core/config" + "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("lets only one frontend lead at a time on the same database", 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()) + first, second := failoverLeaderGate(db), failoverLeaderGate(db) + ctx := context.Background() + + var secondInside, secondRan bool + Expect(first(ctx, func() { + secondInside = second(ctx, func() { secondRan = true }) + })).To(BeTrue()) + Expect(secondInside).To(BeFalse(), "the second gate must not lead while the first holds the lock") + Expect(secondRan).To(BeFalse()) + + Expect(second(ctx, func() { secondRan = true })).To(BeTrue()) + Expect(secondRan).To(BeTrue()) + }) +}) + +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") + }) +}) diff --git a/core/application/startup.go b/core/application/startup.go index db9e1caaf..fe8ee1394 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -286,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) diff --git a/core/services/advisorylock/keys.go b/core/services/advisorylock/keys.go index 277817229..76a1834ed 100644 --- a/core/services/advisorylock/keys.go +++ b/core/services/advisorylock/keys.go @@ -12,4 +12,7 @@ const ( KeySchemaMigrate int64 = 105 KeyBackendUpgradeCheck int64 = 106 KeyStateReconciler int64 = 107 + // KeyFailoverProber elects the one frontend that probes failover + // targets, decides chains and preloads warm targets. + KeyFailoverProber int64 = 108 ) diff --git a/core/services/failover/schedule.go b/core/services/failover/schedule.go index 739c48b42..c910cb9b4 100644 --- a/core/services/failover/schedule.go +++ b/core/services/failover/schedule.go @@ -46,11 +46,14 @@ func (m *Manager) Run(ctx context.Context) { // leader decision. func (m *Manager) Tick(ctx context.Context) { m.Sync() - if m.gate == nil { + m.mu.Lock() + gate := m.gate // SetLeaderGate may replace it after construction + m.mu.Unlock() + if gate == nil { m.lead(ctx) return } - if m.gate(ctx, func() { m.lead(ctx) }) { + if gate(ctx, func() { m.lead(ctx) }) { return } m.mu.Lock() diff --git a/core/services/failover/statesync.go b/core/services/failover/statesync.go index bcab901ef..a14378a04 100644 --- a/core/services/failover/statesync.go +++ b/core/services/failover/statesync.go @@ -48,6 +48,18 @@ type LeaderGate func(ctx context.Context, fn func()) bool // grants leadership. Without it the manager is always the leader. func WithLeaderGate(g LeaderGate) Option { return func(m *Manager) { m.gate = g } } +// SetLeaderGate is WithLeaderGate for a manager that already exists: the +// application builds the manager before distributed init, where the gate's +// database becomes known. Leadership is reset to match New, so the first tick +// that wins the gate counts as becoming leader (warm set redelivered, state +// republished at once). +func (m *Manager) SetLeaderGate(g LeaderGate) { + m.mu.Lock() + defer m.mu.Unlock() + m.gate = g + m.leader = g == nil +} + // SetStateSync attaches the sync layer and hydrates pins from it. Pins are // read outside the lock because the store may need I/O. func (m *Manager) SetStateSync(s StateSync) { diff --git a/core/services/failover/statesync_test.go b/core/services/failover/statesync_test.go index 3901d3b95..d379bb674 100644 --- a/core/services/failover/statesync_test.go +++ b/core/services/failover/statesync_test.go @@ -140,6 +140,18 @@ var _ = Describe("Manager state sync", func() { Expect(st.Pinned).ToNot(BeNil()) }) + It("SetLeaderGate gates a manager built without one", func() { + // Production builds the manager before distributed init, so the gate + // arrives through the setter rather than the option. + p := &fakeProber{fail: map[string]error{}} + m := New(src, WithClock(clock), WithProber(p)) + Expect(m.IsLeader()).To(BeTrue()) + m.SetLeaderGate(gateFor(false)) + Expect(m.IsLeader()).To(BeFalse(), "a gated manager is not the leader until the gate grants it") + m.Tick(ctx) + Consistently(func() int { return len(p.take()) }, 200*time.Millisecond).Should(Equal(0), "a follower must not probe") + }) + It("only the leader probes", func() { pa, pb := &fakeProber{fail: map[string]error{}}, &fakeProber{fail: map[string]error{}} a = New(src, WithClock(clock), WithProber(pa), WithLeaderGate(gateFor(true))) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index 5915ff8b3..24b9ae8dd 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -167,7 +167,8 @@ curl -X DELETE http://localhost:8080/api/failover/assistant-llm/pin ``` While a chain is pinned, only the pinned target serves it. Health checks -continue. A restart removes the pin. +continue. On a single LocalAI instance, a restart removes the pin. In +[distributed mode](#distributed-mode), pins persist. ## Assistant and MCP @@ -175,10 +176,25 @@ The LocalAI Assistant and `local-ai mcp-server` offer `list_failover_chains`, `pin_failover_target` and `unpin_failover_target`. Create and edit chains with the model config tools, like any other model. +## Distributed mode + +In [distributed mode]({{%relref "features/distributed-mode" %}}), all frontends +share one failover state: + +- Pins apply to the whole cluster. LocalAI stores them in the database, so + they persist across restarts. A pin set on one frontend applies on all. +- Target health and the active target of each chain are shared over NATS. + All frontends send a chain's requests to the same target. +- One frontend, the probe leader, runs the health checks, decides fail-over + and fail-back, and loads warm targets. A PostgreSQL advisory lock selects + the leader. If the leader stops, another frontend takes over. +- Warm targets stay loaded on the workers. The router and the replica + reconciler treat them like pinned models and do not evict them. +- A frontend that starts late gets the current state within 10 seconds, + because the leader sends its full state again every 10 seconds. + ## Limits -- Failover state is kept in memory by each LocalAI instance. Several frontends - in distributed mode each keep their own view. - Chains do not nest. - See also [model aliases]({{%relref "features/model-aliases" %}}) and the [realtime API]({{%relref "features/openai-realtime" %}}). From b5792e4d175652c1bc650f06a18d2d2253b50989 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:27:04 +0000 Subject: [PATCH 39/79] fix(failover): keep the probe leader until its session ends A lock taken per tick passed between frontends on almost every tick, so several frontends probed at once and each change of leader re-sent the warm set and all state. The leader now holds a dedicated PostgreSQL session with the advisory lock and keeps it until it shuts down or the session dies. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/application.go | 10 +- core/application/failover_distributed.go | 49 ++++-- core/application/failover_distributed_test.go | 52 +++++-- core/services/advisorylock/held_lock.go | 147 ++++++++++++++++++ core/services/advisorylock/held_lock_test.go | 105 +++++++++++++ docs/content/features/model-failover.md | 14 +- 6 files changed, 345 insertions(+), 32 deletions(-) create mode 100644 core/services/advisorylock/held_lock.go create mode 100644 core/services/advisorylock/held_lock_test.go diff --git a/core/application/application.go b/core/application/application.go index b15d848a0..56e22eeed 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -13,6 +13,7 @@ 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" @@ -98,6 +99,9 @@ type Application struct { // 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 @@ -522,11 +526,7 @@ func (a *Application) Shutdown() error { a.shutdownOnce.Do(func() { // Before distributed shutdown: the sync's subscriptions live on the // NATS connection that closes there. - if a.failoverSync != nil { - if closeErr := a.failoverSync.Close(); closeErr != nil { - xlog.Warn("failover: closing state sync", "error", closeErr) - } - } + a.stopFailoverDistributed() a.distributed.Shutdown() if a.modelLoader != nil { err = a.modelLoader.StopAllGRPC() diff --git a/core/application/failover_distributed.go b/core/application/failover_distributed.go index 7c1b63b9f..37eb94edc 100644 --- a/core/application/failover_distributed.go +++ b/core/application/failover_distributed.go @@ -8,24 +8,31 @@ import ( "github.com/mudler/LocalAI/core/services/failover/distsync" "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/xlog" - "gorm.io/gorm" ) // 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. A lock error counts as "not leader": -// skipping one tick is safe, two leaders are not. -func failoverLeaderGate(db *gorm.DB) failover.LeaderGate { +// 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 { - ok, err := advisorylock.TryWithLockCtx(ctx, db, advisorylock.KeyFailoverProber, func() error { - fn() - return nil - }) - if err != nil { - xlog.Warn("failover: could not take the prober leader lock", "error", err) - return false + 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 + } } - return ok + fn() + return true } } @@ -65,5 +72,21 @@ func (a *Application) startFailoverDistributed(ctx context.Context) { // 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.failoverManager.SetLeaderGate(failoverLeaderGate(db)) + 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 { + a.failoverLock.Release() + } } diff --git a/core/application/failover_distributed_test.go b/core/application/failover_distributed_test.go index b4cf84bf0..1021d0105 100644 --- a/core/application/failover_distributed_test.go +++ b/core/application/failover_distributed_test.go @@ -5,6 +5,7 @@ import ( "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" @@ -57,23 +58,30 @@ var _ = Describe("failoverPinnedResolver", func() { }) var _ = Describe("failoverLeaderGate", func() { - It("lets only one frontend lead at a time on the same database", 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()) - first, second := failoverLeaderGate(db), failoverLeaderGate(db) + 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 secondInside, secondRan bool - Expect(first(ctx, func() { - secondInside = second(ctx, func() { secondRan = true }) - })).To(BeTrue()) - Expect(secondInside).To(BeFalse(), "the second gate must not lead while the first holds the lock") - Expect(secondRan).To(BeFalse()) + 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()) - Expect(second(ctx, func() { secondRan = true })).To(BeTrue()) - Expect(secondRan).To(BeTrue()) + 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()) }) }) @@ -99,4 +107,28 @@ var _ = Describe("applyFailoverWarmTargets in distributed mode", func() { app.applyFailoverWarmTargets([]string{"b"}) Consistently(preloaded, 200*time.Millisecond).ShouldNot(Receive(), "only the probe leader preloads warm targets") }) + + It("preloads warm targets on the probe leader", func() { + preloaded := make(chan string, 4) + orig := preloadModelByName + preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) { + preloaded <- name + return nil, nil + } + DeferCleanup(func() { preloadModelByName = orig }) + + app := &Application{ + applicationConfig: &config.ApplicationConfig{Context: context.Background()}, + distributed: &DistributedServices{}, + } + // A gate that always grants: this frontend is the leader. The warm + // set is delivered on the first tick that wins it. + app.failoverManager = failover.New(warmChainSource(), + failover.WithLeaderGate(func(_ context.Context, fn func()) bool { fn(); return true }), + failover.WithOnWarmChanged(app.applyFailoverWarmTargets)) + + app.failoverManager.Tick(context.Background()) + Eventually(preloaded, time.Second).Should(Receive(Equal("b"))) + Eventually(preloaded, time.Second).Should(Receive(Equal("c"))) + }) }) diff --git a/core/services/advisorylock/held_lock.go b/core/services/advisorylock/held_lock.go new file mode 100644 index 000000000..92e8acbe8 --- /dev/null +++ b/core/services/advisorylock/held_lock.go @@ -0,0 +1,147 @@ +package advisorylock + +import ( + "context" + "database/sql" + "database/sql/driver" + "fmt" + "sync" + "time" + + "github.com/mudler/xlog" + "gorm.io/gorm" +) + +// heldLockCheckTimeout bounds the liveness check and the unlock, so a hung +// database cannot stall the caller's loop. +const heldLockCheckTimeout = 5 * time.Second + +// HeldLock is an advisory lock that stays taken across calls, for leader +// election where leadership must be sticky. TryWithLockCtx releases the lock +// when fn returns; with a short fn, every contender wins it in turn and +// leadership flips on each tick. +// +// On PostgreSQL the lock belongs to one session, so HeldLock keeps a +// dedicated connection out of the pool while it holds the lock. If that +// session dies, the server drops the lock and another instance can take it. +// On other dialects it holds the package's in-process lock for the key. +type HeldLock struct { + db *gorm.DB + key int64 + + mu sync.Mutex + conn *sql.Conn // PostgreSQL: the session that holds the lock + local bool // other dialects: this lock holds the in-process slot +} + +// NewHeldLock returns a lock for key on db. It takes nothing until TryAcquire. +func NewHeldLock(db *gorm.DB, key int64) *HeldLock { + return &HeldLock{db: db, key: key} +} + +// TryAcquire takes the lock without blocking. It returns true if this +// HeldLock holds the lock afterwards, including when it already held it. +func (l *HeldLock) TryAcquire(ctx context.Context) (bool, error) { + l.mu.Lock() + defer l.mu.Unlock() + if l.conn != nil || l.local { + return true, nil + } + + if !isPostgres(l.db) { + select { + case localLockChan(l.key) <- struct{}{}: + l.local = true + return true, nil + default: + return false, nil + } + } + + sqlDB, err := l.db.DB() + if err != nil { + return false, fmt.Errorf("get sql.DB: %w", err) + } + conn, err := sqlDB.Conn(ctx) + if err != nil { + return false, fmt.Errorf("advisory lock conn: %w", err) + } + var acquired bool + if err := conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", l.key).Scan(&acquired); err != nil { + // The lock may have been granted before the error (a cancelled + // context, say); discarding the session is the only way to be sure + // it is not left held. + discardConn(conn) + return false, fmt.Errorf("pg_try_advisory_lock: %w", err) + } + if !acquired { + _ = conn.Close() // holds nothing, so it may go back to the pool + return false, nil + } + l.conn = conn + return true, nil +} + +// Held reports whether this HeldLock believes it holds the lock. It does not +// touch the database; use Verify for that. +func (l *HeldLock) Held() bool { + l.mu.Lock() + defer l.mu.Unlock() + return l.conn != nil || l.local +} + +// Verify reports whether the lock is still held. On PostgreSQL it checks +// that the holding session is alive; if it is not, the lock is gone +// server-side, so Verify drops it here too and returns false. +func (l *HeldLock) Verify(ctx context.Context) bool { + l.mu.Lock() + defer l.mu.Unlock() + if l.local { + return true + } + if l.conn == nil { + return false + } + pctx, cancel := context.WithTimeout(ctx, heldLockCheckTimeout) + defer cancel() + if err := l.conn.PingContext(pctx); err != nil { + xlog.Warn("advisory lock session lost, giving up the lock", "key", l.key, "error", err) + // A ping that only timed out may leave the session (and the lock) + // alive; discarding closes it, so the server releases the lock. + discardConn(l.conn) + l.conn = nil + return false + } + return true +} + +// Release gives up the lock. It is safe to call when the lock is not held. +func (l *HeldLock) Release() { + l.mu.Lock() + defer l.mu.Unlock() + if l.local { + <-localLockChan(l.key) + l.local = false + return + } + if l.conn == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), heldLockCheckTimeout) + defer cancel() + if _, err := l.conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", l.key); err != nil { + xlog.Warn("advisory lock unlock failed, closing its session instead", "key", l.key, "error", err) + } + // Never hand the session back to the pool: if the unlock failed it still + // holds the lock, and a later pool user would silently inherit it. + discardConn(l.conn) + l.conn = nil +} + +// discardConn closes conn's underlying session instead of returning it to +// the pool. database/sql drops a connection when Raw's callback returns +// driver.ErrBadConn. +func discardConn(conn *sql.Conn) { + _ = conn.Raw(func(any) error { return driver.ErrBadConn }) + _ = conn.Close() +} diff --git a/core/services/advisorylock/held_lock_test.go b/core/services/advisorylock/held_lock_test.go new file mode 100644 index 000000000..fa8322d10 --- /dev/null +++ b/core/services/advisorylock/held_lock_test.go @@ -0,0 +1,105 @@ +package advisorylock + +import ( + "context" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/testutil" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +// expectStickyHandover checks the contract both backends share: the first +// holder keeps the lock across repeated checks while a rival is denied, and +// the rival gets it once the holder releases. +func expectStickyHandover(db *gorm.DB, key int64) { + ctx := context.Background() + first, second := NewHeldLock(db, key), NewHeldLock(db, key) + DeferCleanup(first.Release) + DeferCleanup(second.Release) + + ok, err := first.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + for range 5 { + Expect(first.Held()).To(BeTrue()) + Expect(first.Verify(ctx)).To(BeTrue()) + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeFalse(), "a rival must not take a held lock") + Expect(second.Held()).To(BeFalse()) + } + + first.Release() + Expect(first.Held()).To(BeFalse()) + Expect(first.Verify(ctx)).To(BeFalse()) + + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "the lock is free once the holder releases it") + Expect(second.Verify(ctx)).To(BeTrue()) +} + +var _ = Describe("HeldLock (SQLite fallback)", Label("sqlite"), func() { + It("stays with its holder until Release", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + expectStickyHandover(db, 12101) + }) + + It("is idempotent: acquiring twice and releasing twice is safe", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + l := NewHeldLock(db, 12102) + ok, err := l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + ok, err = l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "a holder asking again keeps the lock") + l.Release() + l.Release() + ok, err = l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + l.Release() + }) +}) + +var _ = Describe("HeldLock (PostgreSQL)", func() { + var db *gorm.DB + + BeforeEach(func() { + db = testutil.SetupTestDB() + }) + + It("stays with its holder until Release, then hands over", func() { + expectStickyHandover(db, 12201) + }) + + It("drops leadership when the holding session dies, freeing the lock", func() { + ctx := context.Background() + first, second := NewHeldLock(db, 12202), NewHeldLock(db, 12202) + DeferCleanup(first.Release) + DeferCleanup(second.Release) + + ok, err := first.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + // Kill the holder's backend, as a network cut or DB restart would. + var pid int + Expect(first.conn.QueryRowContext(ctx, "SELECT pg_backend_pid()").Scan(&pid)).To(Succeed()) + Expect(db.Exec("SELECT pg_terminate_backend(?)", pid).Error).ToNot(HaveOccurred()) + + Expect(first.Verify(ctx)).To(BeFalse(), "a dead session no longer holds the lock") + Expect(first.Held()).To(BeFalse()) + + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "the server released the dead session's lock") + }) +}) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index 24b9ae8dd..f4f42027b 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -183,15 +183,21 @@ share one failover state: - Pins apply to the whole cluster. LocalAI stores them in the database, so they persist across restarts. A pin set on one frontend applies on all. -- Target health and the active target of each chain are shared over NATS. - All frontends send a chain's requests to the same target. +- Target health and the active target of each chain are shared over NATS, so + all frontends converge on the same target for a chain. - One frontend, the probe leader, runs the health checks, decides fail-over - and fail-back, and loads warm targets. A PostgreSQL advisory lock selects - the leader. If the leader stops, another frontend takes over. + and fail-back, and loads warm targets. The leader holds a PostgreSQL + advisory lock and keeps it until it shuts down or its database connection + fails. Then another frontend takes the lock and becomes the leader. - Warm targets stay loaded on the workers. The router and the replica reconciler treat them like pinned models and do not evict them. - A frontend that starts late gets the current state within 10 seconds, because the leader sends its full state again every 10 seconds. +- If PostgreSQL is not available, no frontend holds the lock. Health checks + and fail-back stop until the database is back. Requests still fail over to + the next target when a target fails during the request. +- If a frontend cannot start the shared state, it logs an error and manages + failover alone, as a single LocalAI instance does. ## Limits From e59fb854ee2fb769819976540e4fc79b60c3e17f Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:31:39 +0000 Subject: [PATCH 40/79] fix(failover): free the leader lock soon after the leader's host dies A leader whose host died without closing its connection kept the advisory lock for about two hours of OS keepalive defaults, and no other frontend could probe. The lock session now sets short TCP keepalives and tcp_user_timeout, so the server drops it within about 30 seconds. Shutdown now closes the lock for good, so a tick that runs after it cannot take the lock back. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/application/failover_distributed.go | 4 +- core/services/advisorylock/held_lock.go | 118 ++++++++++++++----- core/services/advisorylock/held_lock_test.go | 78 ++++++++++++ docs/content/features/model-failover.md | 6 +- 4 files changed, 175 insertions(+), 31 deletions(-) diff --git a/core/application/failover_distributed.go b/core/application/failover_distributed.go index 37eb94edc..87adbdeb5 100644 --- a/core/application/failover_distributed.go +++ b/core/application/failover_distributed.go @@ -87,6 +87,8 @@ func (a *Application) stopFailoverDistributed() { } } if a.failoverLock != nil { - a.failoverLock.Release() + // Close, not Release: Run may still be ticking and would take the + // lock straight back. + a.failoverLock.Close() } } diff --git a/core/services/advisorylock/held_lock.go b/core/services/advisorylock/held_lock.go index 92e8acbe8..46f2db6f5 100644 --- a/core/services/advisorylock/held_lock.go +++ b/core/services/advisorylock/held_lock.go @@ -16,22 +16,41 @@ import ( // database cannot stall the caller's loop. const heldLockCheckTimeout = 5 * time.Second +// heldLockSessionSettings make the server notice a lock holder whose host +// died without closing the connection (crash, power loss, partition) within +// about 30 seconds, instead of after the OS keepalive default of over two +// hours, during which nobody else could take the lock. +var heldLockSessionSettings = []string{ + "SET tcp_keepalives_idle = 10", + "SET tcp_keepalives_interval = 5", + "SET tcp_keepalives_count = 3", +} + +// heldLockUserTimeout bounds how long unacknowledged data (for example a +// keepalive reply to a dead host) may stay in flight. PostgreSQL 12 added it, +// so it is applied separately and an older server's error is ignored. +const heldLockUserTimeout = "SET tcp_user_timeout = 30000" + // HeldLock is an advisory lock that stays taken across calls, for leader // election where leadership must be sticky. TryWithLockCtx releases the lock // when fn returns; with a short fn, every contender wins it in turn and // leadership flips on each tick. // // On PostgreSQL the lock belongs to one session, so HeldLock keeps a -// dedicated connection out of the pool while it holds the lock. If that -// session dies, the server drops the lock and another instance can take it. -// On other dialects it holds the package's in-process lock for the key. +// dedicated connection out of the pool. The same session is reused for +// later attempts while the lock is taken elsewhere, so a follower does not +// open a connection per try. If the holding session dies, the server drops +// the lock and another instance can take it. On other dialects it holds the +// package's in-process lock for the key. type HeldLock struct { db *gorm.DB key int64 - mu sync.Mutex - conn *sql.Conn // PostgreSQL: the session that holds the lock - local bool // other dialects: this lock holds the in-process slot + mu sync.Mutex + conn *sql.Conn // PostgreSQL: the dedicated session, holding the lock or not + held bool // PostgreSQL: conn holds the lock + local bool // other dialects: this lock holds the in-process slot + closed bool // Close was called; the lock is never taken again } // NewHeldLock returns a lock for key on db. It takes nothing until TryAcquire. @@ -44,7 +63,10 @@ func NewHeldLock(db *gorm.DB, key int64) *HeldLock { func (l *HeldLock) TryAcquire(ctx context.Context) (bool, error) { l.mu.Lock() defer l.mu.Unlock() - if l.conn != nil || l.local { + if l.closed { + return false, nil + } + if l.held || l.local { return true, nil } @@ -58,28 +80,48 @@ func (l *HeldLock) TryAcquire(ctx context.Context) (bool, error) { } } - sqlDB, err := l.db.DB() - if err != nil { - return false, fmt.Errorf("get sql.DB: %w", err) - } - conn, err := sqlDB.Conn(ctx) - if err != nil { - return false, fmt.Errorf("advisory lock conn: %w", err) + if l.conn == nil { + conn, err := l.openSession(ctx) + if err != nil { + return false, err + } + l.conn = conn } var acquired bool - if err := conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", l.key).Scan(&acquired); err != nil { + if err := l.conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", l.key).Scan(&acquired); err != nil { // The lock may have been granted before the error (a cancelled // context, say); discarding the session is the only way to be sure // it is not left held. - discardConn(conn) + discardConn(l.conn) + l.conn = nil return false, fmt.Errorf("pg_try_advisory_lock: %w", err) } - if !acquired { - _ = conn.Close() // holds nothing, so it may go back to the pool - return false, nil + l.held = acquired + return acquired, nil +} + +// openSession takes a connection out of the pool for this lock and applies +// the keepalive settings to it. The session never goes back to the pool, so +// other pool users do not inherit those settings. +func (l *HeldLock) openSession(ctx context.Context) (*sql.Conn, error) { + sqlDB, err := l.db.DB() + if err != nil { + return nil, fmt.Errorf("get sql.DB: %w", err) } - l.conn = conn - return true, nil + conn, err := sqlDB.Conn(ctx) + if err != nil { + return nil, fmt.Errorf("advisory lock conn: %w", err) + } + for _, stmt := range heldLockSessionSettings { + if _, err := conn.ExecContext(ctx, stmt); err != nil { + discardConn(conn) + return nil, fmt.Errorf("advisory lock session %q: %w", stmt, err) + } + } + if _, err := conn.ExecContext(ctx, heldLockUserTimeout); err != nil { + xlog.Debug("advisory lock session: tcp_user_timeout not supported, relying on keepalives", "error", err) + } + return conn, nil } // Held reports whether this HeldLock believes it holds the lock. It does not @@ -87,7 +129,7 @@ func (l *HeldLock) TryAcquire(ctx context.Context) (bool, error) { func (l *HeldLock) Held() bool { l.mu.Lock() defer l.mu.Unlock() - return l.conn != nil || l.local + return l.held || l.local } // Verify reports whether the lock is still held. On PostgreSQL it checks @@ -99,7 +141,7 @@ func (l *HeldLock) Verify(ctx context.Context) bool { if l.local { return true } - if l.conn == nil { + if !l.held { return false } pctx, cancel := context.WithTimeout(ctx, heldLockCheckTimeout) @@ -110,15 +152,32 @@ func (l *HeldLock) Verify(ctx context.Context) bool { // alive; discarding closes it, so the server releases the lock. discardConn(l.conn) l.conn = nil + l.held = false return false } return true } -// Release gives up the lock. It is safe to call when the lock is not held. +// Release gives up the lock; a later TryAcquire may take it again. It is safe +// to call when the lock is not held. func (l *HeldLock) Release() { l.mu.Lock() defer l.mu.Unlock() + l.releaseLocked() +} + +// Close gives up the lock for good: TryAcquire returns false afterwards. Use +// it on shutdown, where a loop still running could otherwise take the lock +// back right after Release (and, on the in-process fallback, keep it for the +// rest of the process). +func (l *HeldLock) Close() { + l.mu.Lock() + defer l.mu.Unlock() + l.closed = true + l.releaseLocked() +} + +func (l *HeldLock) releaseLocked() { if l.local { <-localLockChan(l.key) l.local = false @@ -127,15 +186,18 @@ func (l *HeldLock) Release() { if l.conn == nil { return } - ctx, cancel := context.WithTimeout(context.Background(), heldLockCheckTimeout) - defer cancel() - if _, err := l.conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", l.key); err != nil { - xlog.Warn("advisory lock unlock failed, closing its session instead", "key", l.key, "error", err) + if l.held { + ctx, cancel := context.WithTimeout(context.Background(), heldLockCheckTimeout) + defer cancel() + if _, err := l.conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", l.key); err != nil { + xlog.Warn("advisory lock unlock failed, closing its session instead", "key", l.key, "error", err) + } } // Never hand the session back to the pool: if the unlock failed it still // holds the lock, and a later pool user would silently inherit it. discardConn(l.conn) l.conn = nil + l.held = false } // discardConn closes conn's underlying session instead of returning it to diff --git a/core/services/advisorylock/held_lock_test.go b/core/services/advisorylock/held_lock_test.go index fa8322d10..17f2c17b9 100644 --- a/core/services/advisorylock/held_lock_test.go +++ b/core/services/advisorylock/held_lock_test.go @@ -11,6 +11,29 @@ import ( "gorm.io/gorm" ) +// expectCloseIsFinal checks that Close frees the lock for a rival and that the +// closed HeldLock never takes it back, as a still-running loop would try to. +func expectCloseIsFinal(db *gorm.DB, key int64) { + ctx := context.Background() + l, rival := NewHeldLock(db, key), NewHeldLock(db, key) + DeferCleanup(rival.Release) + + ok, err := l.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + l.Close() + Expect(l.Held()).To(BeFalse()) + + ok, err = l.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeFalse(), "a closed lock is never taken again") + + ok, err = rival.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "Close frees the lock for others") + l.Close() // idempotent +} + // expectStickyHandover checks the contract both backends share: the first // holder keeps the lock across repeated checks while a rival is denied, and // the rival gets it once the holder releases. @@ -50,6 +73,12 @@ var _ = Describe("HeldLock (SQLite fallback)", Label("sqlite"), func() { expectStickyHandover(db, 12101) }) + It("never takes the lock again after Close", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + expectCloseIsFinal(db, 12103) + }) + It("is idempotent: acquiring twice and releasing twice is safe", func() { db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) Expect(err).ToNot(HaveOccurred()) @@ -80,6 +109,55 @@ var _ = Describe("HeldLock (PostgreSQL)", func() { expectStickyHandover(db, 12201) }) + It("never takes the lock again after Close", func() { + expectCloseIsFinal(db, 12203) + }) + + It("sets short TCP keepalives on its session so a dead host's lock expires", func() { + // Killing a host without closing its socket is not practical in a + // test; check the settings that make the server notice one. + ctx := context.Background() + l := NewHeldLock(db, 12204) + DeferCleanup(l.Release) + ok, err := l.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + show := func(name string) string { + var v string + Expect(l.conn.QueryRowContext(ctx, "SHOW "+name).Scan(&v)).To(Succeed()) + return v + } + Expect(show("tcp_keepalives_idle")).To(Equal("10")) + Expect(show("tcp_keepalives_interval")).To(Equal("5")) + Expect(show("tcp_keepalives_count")).To(Equal("3")) + Expect(show("tcp_user_timeout")).To(Equal("30000")) // milliseconds + }) + + It("reuses one session while the lock is taken elsewhere", func() { + ctx := context.Background() + holder, follower := NewHeldLock(db, 12205), NewHeldLock(db, 12205) + DeferCleanup(holder.Release) + DeferCleanup(follower.Release) + ok, err := holder.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + pid := func() int { + var p int + Expect(follower.conn.QueryRowContext(ctx, "SELECT pg_backend_pid()").Scan(&p)).To(Succeed()) + return p + } + ok, err = follower.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeFalse()) + before := pid() + ok, err = follower.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeFalse()) + Expect(pid()).To(Equal(before), "a follower must not open a connection per attempt") + }) + It("drops leadership when the holding session dies, freeing the lock", func() { ctx := context.Background() first, second := NewHeldLock(db, 12202), NewHeldLock(db, 12202) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index f4f42027b..32e91af96 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -187,8 +187,10 @@ share one failover state: all frontends converge on the same target for a chain. - One frontend, the probe leader, runs the health checks, decides fail-over and fail-back, and loads warm targets. The leader holds a PostgreSQL - advisory lock and keeps it until it shuts down or its database connection - fails. Then another frontend takes the lock and becomes the leader. + advisory lock and keeps it until it stops or its database connection fails. + Then another frontend takes the lock and becomes the leader: immediately + when the leader shuts down or its process exits, and within about 30 seconds + when the leader's host or network fails. - Warm targets stay loaded on the workers. The router and the replica reconciler treat them like pinned models and do not evict them. - A frontend that starts late gets the current state within 10 seconds, From 509cca35e0fdfcf2128dbbd972e1c89e832c0842 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 22:41:37 +0000 Subject: [PATCH 41/79] feat(grpc): serve Rerank from Go backends and skip Unimplemented targets Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/backend/options.go | 2 +- core/backend/options_internal_test.go | 25 +++++++ core/config/model_config_loader.go | 16 +++++ core/config/model_config_loader_test.go | 77 ++++++++++++++++++++++ core/http/middleware/failover.go | 9 ++- core/http/middleware/failover_test.go | 14 ++++ core/services/failover/classify.go | 12 ++++ core/services/failover/classify_test.go | 11 ++++ core/services/failover/manager.go | 7 ++ core/services/failover/manager_test.go | 17 +++++ pkg/grpc/interface.go | 10 +++ pkg/grpc/model_identity_modalities_test.go | 9 ++- pkg/grpc/server.go | 15 +++++ pkg/grpc/server_rerank_test.go | 49 ++++++++++++++ 14 files changed, 266 insertions(+), 7 deletions(-) create mode 100644 pkg/grpc/server_rerank_test.go diff --git a/core/backend/options.go b/core/backend/options.go index 93700f93d..12f2e29a1 100644 --- a/core/backend/options.go +++ b/core/backend/options.go @@ -542,7 +542,7 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions { Tokenizer: c.Tokenizer, } - if c.Backend == "cloud-proxy" { + if c.Backend == "cloud-proxy" || c.Backend == "localai-proxy" { opts.Proxy = &pb.ProxyOptions{ UpstreamUrl: c.Proxy.UpstreamURL, Mode: c.Proxy.Mode, diff --git a/core/backend/options_internal_test.go b/core/backend/options_internal_test.go index 18e081fb1..5f8370211 100644 --- a/core/backend/options_internal_test.go +++ b/core/backend/options_internal_test.go @@ -44,6 +44,31 @@ 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("leaves Proxy nil for a backend that is not a proxy", func() { + threads := 1 + opts := grpcModelOpts(config.ModelConfig{Threads: &threads, Backend: "llama-cpp"}, "/tmp/models") + Expect(opts.Proxy).To(BeNil()) + }) +}) + var _ = Describe("grpcModelOpts Diffusers options", func() { It("forwards original_config_file without rewriting it", func() { threads := 1 diff --git a/core/config/model_config_loader.go b/core/config/model_config_loader.go index a29c95c52..5e979d774 100644 --- a/core/config/model_config_loader.go +++ b/core/config/model_config_loader.go @@ -958,6 +958,22 @@ func (bcl *ModelConfigLoader) loadModelConfigsFromPath(path string, strict bool, } } + // 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 != "localai-proxy" { + continue + } + if cfg.Proxy.Mode == ProxyModeTranslate || cfg.Proxy.Provider != "" { + xlog.Warn("localai-proxy backend proxies to another LocalAI instance and ignores proxy.mode/proxy.provider", "model", name) + } + if len(cfg.KnownUsecaseStrings) == 0 { + xlog.Warn("localai-proxy config has no known_usecases; failover usecase matching will skip it", "model", name) + } + } + return nil } diff --git a/core/config/model_config_loader_test.go b/core/config/model_config_loader_test.go index aed79dbb8..e99a9136d 100644 --- a/core/config/model_config_loader_test.go +++ b/core/config/model_config_loader_test.go @@ -1,7 +1,9 @@ package config import ( + "bytes" "context" + "log/slog" "os" "path/filepath" @@ -9,6 +11,7 @@ import ( . "github.com/onsi/gomega" "github.com/mudler/LocalAI/pkg/modelartifacts" + "github.com/mudler/xlog" ) type preloadArtifactMaterializer struct { @@ -452,3 +455,77 @@ var _ = Describe("ModelConfigLoader failover validation", func() { Expect(failoverWarmRemoteTargets(chain("remote", "a"), loader.GetModelConfig)).To(BeEmpty()) }) }) + +var _ = Describe("ModelConfigLoader localai-proxy load-time warnings", func() { + var captured *bytes.Buffer + + BeforeEach(func() { + captured = &bytes.Buffer{} + handler := slog.NewTextHandler(captured, &slog.HandlerOptions{Level: slog.LevelWarn}) + xlog.SetLogger(xlog.NewLoggerWithHandler(handler, xlog.LogLevelWarn)) + }) + + AfterEach(func() { + // xlog exposes no getter for the package logger, so restore the same + // default the suite entrypoint installs rather than the prior value. + xlog.SetLogger(xlog.NewLogger(xlog.LogLevel("info"), "text")) + }) + + It("warns and still loads when proxy.mode/proxy.provider are set (ignored by localai-proxy)", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied +backend: localai-proxy +known_usecases: [chat] +proxy: + mode: translate + provider: openai + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + _, ok := loader.GetModelConfig("proxied") + Expect(ok).To(BeTrue()) + Expect(captured.String()).To(ContainSubstring("proxy.mode/proxy.provider")) + Expect(captured.String()).To(ContainSubstring("proxied")) + }) + + It("warns and still loads when known_usecases is empty", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied-no-usecase +backend: localai-proxy +proxy: + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + _, ok := loader.GetModelConfig("proxied-no-usecase") + Expect(ok).To(BeTrue()) + Expect(captured.String()).To(ContainSubstring("known_usecases")) + Expect(captured.String()).To(ContainSubstring("proxied-no-usecase")) + }) + + It("does not warn when localai-proxy uses only passthrough and known_usecases", func() { + modelsPath := GinkgoT().TempDir() + cfgYAML := ` +name: proxied-clean +backend: localai-proxy +known_usecases: [chat] +proxy: + upstream_url: http://127.0.0.1:8081 +` + Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed()) + + loader := NewModelConfigLoader(modelsPath) + Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + + Expect(captured.String()).To(BeEmpty()) + }) +}) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go index 853a3f22a..e0f84d1ed 100644 --- a/core/http/middleware/failover.go +++ b/core/http/middleware/failover.go @@ -143,9 +143,12 @@ func (re *RequestExtractor) failoverRetry(h echo.HandlerFunc) echo.HandlerFunc { } return nil } - if rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool); rejected && !w.committed { - // The target is at capacity, not broken: spill this request to - // the next target without counting a failure. + rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool) + gap := failover.IsCapabilityGap(err) && !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 diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index b5e8ba5b7..94588bdf9 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -23,6 +23,8 @@ import ( "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() { @@ -279,6 +281,18 @@ var _ = Describe("failover chains in the request pipeline", func() { 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("skips a disabled target without tripping it", func() { rec := chat("chain-off") Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) diff --git a/core/services/failover/classify.go b/core/services/failover/classify.go index 58868c089..a71d89d3f 100644 --- a/core/services/failover/classify.go +++ b/core/services/failover/classify.go @@ -61,6 +61,18 @@ func IsRetryable(err error, status int) bool { return !isRequestError(msg) } +// IsCapabilityGap reports a target that cannot serve this kind of request at +// all (gRPC Unimplemented, anywhere in the error chain). The next target may +// serve it, and this target is not broken: the failure carries no signal +// about its health, so callers must skip it without tripping. +func IsCapabilityGap(err error) bool { + if err == nil { + return false + } + st, ok := grpcstatus.FromError(err) + return ok && st.Code() == codes.Unimplemented +} + func retryableStatus(code int) bool { return code >= 500 && code != http.StatusNotImplemented } diff --git a/core/services/failover/classify_test.go b/core/services/failover/classify_test.go index 455a9f993..cc9994e70 100644 --- a/core/services/failover/classify_test.go +++ b/core/services/failover/classify_test.go @@ -37,3 +37,14 @@ var _ = DescribeTable("IsRetryable", Entry("context overflow", errors.New("the request exceeds the available context size"), 0, false), Entry("dial error", errors.New("dial tcp 10.0.0.1:8080: connect: connection refused"), 0, true), ) + +var _ = DescribeTable("IsCapabilityGap", + func(err error, want bool) { + Expect(IsCapabilityGap(err)).To(Equal(want)) + }, + Entry("nil error", nil, false), + Entry("grpc unimplemented", grpcstatus.Error(codes.Unimplemented, "x"), true), + Entry("wrapped grpc unimplemented", fmt.Errorf("call: %w", grpcstatus.Error(codes.Unimplemented, "x")), true), + Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), false), + Entry("plain error", errors.New("boom"), false), +) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index da5ba6079..071e9ab0c 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -727,6 +727,13 @@ func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Cont case err == nil: att.Succeed() return nil + case IsCapabilityGap(err) && !committed.Load(): + // This target cannot serve this kind of request at all; it is + // not broken, so move on without counting a failure. + if !att.Skip() { + return err + } + continue case ctx.Err() != nil || !IsRetryable(err, 0): return err case committed.Load(): diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 0d2746331..04f29d4db 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -9,6 +9,8 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" . "github.com/onsi/gomega/gstruct" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" ) var errBoom = errors.New("dial tcp: connection refused") @@ -288,5 +290,20 @@ var _ = Describe("Manager", func() { st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateHealthy)) }) + + It("skips an Unimplemented target without tripping it", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, _ func()) error { + tried = append(tried, target) + if target == "a" { + return grpcstatus.Error(codes.Unimplemented, "localai-proxy: Rerank has no upstream counterpart") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) }) }) diff --git a/pkg/grpc/interface.go b/pkg/grpc/interface.go index 3da403552..5287e18a4 100644 --- a/pkg/grpc/interface.go +++ b/pkg/grpc/interface.go @@ -123,3 +123,13 @@ type ClassifyModel interface { type ScoreModel interface { Score(context.Context, *pb.ScoreRequest) (*pb.ScoreResponse, error) } + +// RerankModel is an optional extension to AIModel for backends that +// implement the Rerank RPC (candidate document reranking against a query). +// The gRPC server type-asserts to this interface; backends that do not +// implement it fall through to the UnimplementedBackendServer default. This +// mirrors the ScoreModel pattern: adding a method to AIModel itself would +// break every backend, so the capability is opt-in. +type RerankModel interface { + Rerank(context.Context, *pb.RerankRequest) (*pb.RerankResult, error) +} diff --git a/pkg/grpc/model_identity_modalities_test.go b/pkg/grpc/model_identity_modalities_test.go index b96db9ed5..a166c77bf 100644 --- a/pkg/grpc/model_identity_modalities_test.go +++ b/pkg/grpc/model_identity_modalities_test.go @@ -113,9 +113,12 @@ var _ AIModel = (*modalityBackend)(nil) // server implements, sending `identity` in each request's ModelIdentity field. // One call per RPC, so the returned error count is also the RPC count. // -// Rerank, Score and TokenClassify are absent on purpose: the generic Go server -// does not implement them (they fall through to UnimplementedBackendServer), so -// only the C++ and Python backends can enforce them. +// Score and TokenClassify are absent on purpose: the generic Go server does +// not implement them (they fall through to UnimplementedBackendServer), so +// only the C++ and Python backends can enforce them. Rerank is different: the +// Go server does serve it (see server_rerank_test.go), but only for backends +// that opt in via RerankModel — modalityBackend does not, so it is left out +// here too rather than adding a no-op implementation with nothing to guard. func callAllModalities(c Backend, identity string) map[string]error { ctx := context.Background() errs := map[string]error{} diff --git a/pkg/grpc/server.go b/pkg/grpc/server.go index 40eec06fe..fa0560ed7 100644 --- a/pkg/grpc/server.go +++ b/pkg/grpc/server.go @@ -132,6 +132,21 @@ func (s *server) Score(ctx context.Context, in *pb.ScoreRequest) (*pb.ScoreRespo return sm.Score(ctx, in) } +func (s *server) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) { + if err := s.checkModelIdentity(in); err != nil { + return nil, err + } + rm, ok := s.llm.(RerankModel) + if !ok { + return nil, status.Errorf(codes.Unimplemented, "method Rerank not implemented") + } + if s.llm.Locking() { + s.llm.Lock() + defer s.llm.Unlock() + } + return rm.Rerank(ctx, in) +} + func (s *server) LoadModel(ctx context.Context, in *pb.ModelOptions) (*pb.Result, error) { if s.llm.Locking() { s.llm.Lock() diff --git a/pkg/grpc/server_rerank_test.go b/pkg/grpc/server_rerank_test.go new file mode 100644 index 000000000..d995b6b12 --- /dev/null +++ b/pkg/grpc/server_rerank_test.go @@ -0,0 +1,49 @@ +package grpc + +import ( + "context" + + "github.com/mudler/LocalAI/pkg/grpc/base" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" +) + +// rerankBackend implements RerankModel on top of the minimal AIModel surface, +// mirroring how a Go backend would opt into Score today. +type rerankBackend struct { + base.SingleThread +} + +func (b *rerankBackend) Rerank(context.Context, *pb.RerankRequest) (*pb.RerankResult, error) { + return &pb.RerankResult{ + Results: []*pb.DocumentResult{{Index: 0, RelevanceScore: 1}}, + }, nil +} + +var _ AIModel = (*rerankBackend)(nil) +var _ RerankModel = (*rerankBackend)(nil) + +var _ = Describe("Rerank", func() { + It("is served when the backend implements RerankModel", func() { + Provide("test://rerank-served", &rerankBackend{}) + c := NewClient("test://rerank-served", true, nil, false) + + res, err := c.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}}) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Results).To(HaveLen(1)) + }) + + It("reports Unimplemented when the backend does not implement RerankModel", func() { + Provide("test://rerank-unimplemented", &base.SingleThread{}) + c := NewClient("test://rerank-unimplemented", true, nil, false) + + _, err := c.Rerank(context.Background(), &pb.RerankRequest{Query: "q"}) + Expect(err).To(HaveOccurred()) + st, ok := grpcstatus.FromError(err) + Expect(ok).To(BeTrue()) + Expect(st.Code()).To(Equal(codes.Unimplemented)) + }) +}) From 6d5c9600d75fa0fe678086dda7eb50e13b94dc82 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 23:05:43 +0000 Subject: [PATCH 42/79] feat(localai-proxy): add a backend that serves text APIs from a remote LocalAI Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .github/backend-matrix.yml | 33 ++ .gitignore | 3 + Makefile | 17 +- backend/go/localai-proxy/Makefile | 13 + backend/go/localai-proxy/client.go | 220 ++++++++ .../go/localai-proxy/fake_upstream_test.go | 160 ++++++ .../localai-proxy/localai_proxy_suite_test.go | 17 + backend/go/localai-proxy/main.go | 32 ++ backend/go/localai-proxy/package.sh | 13 + backend/go/localai-proxy/proxy.go | 226 ++++++++ backend/go/localai-proxy/run.sh | 6 + backend/go/localai-proxy/text.go | 362 +++++++++++++ backend/go/localai-proxy/text_test.go | 483 ++++++++++++++++++ backend/index.yaml | 40 ++ 14 files changed, 1621 insertions(+), 4 deletions(-) create mode 100644 backend/go/localai-proxy/Makefile create mode 100644 backend/go/localai-proxy/client.go create mode 100644 backend/go/localai-proxy/fake_upstream_test.go create mode 100644 backend/go/localai-proxy/localai_proxy_suite_test.go create mode 100644 backend/go/localai-proxy/main.go create mode 100755 backend/go/localai-proxy/package.sh create mode 100644 backend/go/localai-proxy/proxy.go create mode 100755 backend/go/localai-proxy/run.sh create mode 100644 backend/go/localai-proxy/text.go create mode 100644 backend/go/localai-proxy/text_test.go diff --git a/.github/backend-matrix.yml b/.github/backend-matrix.yml index fa61fa514..8649209cf 100644 --- a/.github/backend-matrix.yml +++ b/.github/backend-matrix.yml @@ -5948,6 +5948,35 @@ include: dockerfile: "./backend/Dockerfile.golang" context: "./" ubuntu-version: '2404' + # localai-proxy + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/amd64' + platform-tag: 'amd64' + tag-latest: 'auto' + tag-suffix: '-cpu-localai-proxy' + runs-on: 'ubuntu-latest' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "localai-proxy" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' + - build-type: '' + cuda-major-version: "" + cuda-minor-version: "" + platforms: 'linux/arm64' + platform-tag: 'arm64' + tag-latest: 'auto' + tag-suffix: '-cpu-localai-proxy' + runs-on: 'ubuntu-24.04-arm' + base-image: "ubuntu:24.04" + skip-drivers: 'false' + backend: "localai-proxy" + dockerfile: "./backend/Dockerfile.golang" + context: "./" + ubuntu-version: '2404' # valkey-store - build-type: '' cuda-major-version: "" @@ -6753,6 +6782,10 @@ includeDarwin: tag-suffix: "-metal-darwin-arm64-cloud-proxy" build-type: "metal" lang: "go" + - backend: "localai-proxy" + tag-suffix: "-metal-darwin-arm64-localai-proxy" + build-type: "metal" + lang: "go" - backend: "valkey-store" tag-suffix: "-metal-darwin-arm64-valkey-store" build-type: "metal" diff --git a/.gitignore b/.gitignore index a6d890185..f59bbe453 100644 --- a/.gitignore +++ b/.gitignore @@ -29,6 +29,7 @@ LocalAI # Root-level build artifacts when running `go build ./...` against # Go backend packages whose main lives under backend/go/. /cloud-proxy +/localai-proxy /local-store /valkey-store # prevent above rules from omitting the helm chart @@ -50,6 +51,8 @@ tests/e2e-aio/backends /tests/e2e/mock-backend/mock-backend # The cloud-proxy backend binary the e2e suite runs next to the mock backend. /tests/e2e/mock-backend/cloud-proxy +# The localai-proxy backend binary, built next to it the same way. +/tests/e2e/mock-backend/localai-proxy release/ diff --git a/Makefile b/Makefile index e39c5dfa0..49812a8d7 100644 --- a/Makefile +++ b/Makefile @@ -1,5 +1,5 @@ # Disable parallel execution for backend builds -.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin +.NOTPARALLEL: backends/diffusers backends/llama-cpp backends/turboquant backends/bonsai backends/outetts backends/piper backends/stablediffusion-ggml backends/trellis2cpp backends/trellis2cpp-darwin backends/whisper backends/crispasr backends/parakeet-cpp backends/moss-transcribe-cpp backends/nemo-speech-cpp backends/faster-whisper backends/silero-vad backends/local-store backends/valkey-store backends/cloud-proxy backends/localai-proxy backends/huggingface backends/rfdetr backends/rfdetr-cpp backends/insightface backends/speaker-recognition backends/kitten-tts backends/kokoro backends/chatterbox backends/llama-cpp-darwin backends/neutts build-darwin-python-backend build-darwin-go-backend backends/mlx backends/mlx-video backends/diffuser-darwin backends/mlx-vlm backends/mlx-audio backends/mlx-distributed backends/stablediffusion-ggml-darwin backends/vllm backends/vllm-omni backends/longcat-video backends/sglang backends/moonshine backends/pocket-tts backends/qwen-tts backends/faster-qwen3-tts backends/qwen-asr backends/nemo backends/voxcpm backends/whisperx backends/ace-step backends/acestep-cpp backends/fish-speech backends/voxtral backends/opus backends/trl backends/llama-cpp-quantization backends/kokoros backends/sam3-cpp backends/qwen3-tts-cpp backends/moss-tts-cpp backends/magpie-tts-cpp backends/vllm-cpp backends/omnivoice-cpp backends/vibevoice-cpp backends/localvqe backends/tinygrad backends/sherpa-onnx backends/ds4 backends/ds4-darwin backends/liquid-audio backends/supertonic backends/depth-anything-cpp backends/privacy-filter backends/privacy-filter-darwin backends/audio-cpp backends/audio-cpp-darwin .NOTPARALLEL: backends/whisper-medusa .NOTPARALLEL: backends/funasr @@ -76,7 +76,7 @@ else GORELEASER=$(shell which goreleaser) endif -TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/... +TEST_PATHS?=./api/... ./pkg/... ./core/... ./backend/go/cloud-proxy/... ./backend/go/localai-proxy/... ./backend/go/local-store/... ./backend/go/valkey-store/... ## Coverage output and the committed baseline that CI compares against. ## The gate is strict: total coverage must never decrease (no tolerance). @@ -385,13 +385,14 @@ prepare-e2e: run-e2e-image: docker run -p 5390:8080 -e MODELS_PATH=/models -e THREADS=1 -e DEBUG=true -d --rm -v $(TEST_DIR):/models --name e2e-tests-$(RANDOM) localai-tests -test-e2e: build-mock-backend build-cloud-proxy-backend prepare-e2e run-e2e-image +test-e2e: build-mock-backend build-cloud-proxy-backend build-localai-proxy-backend prepare-e2e run-e2e-image @echo 'Running e2e tests' BUILD_TYPE=$(BUILD_TYPE) \ LOCALAI_API=http://$(E2E_BRIDGE_IP):5390 \ $(GOCMD) run github.com/onsi/ginkgo/v2/ginkgo --flake-attempts $(TEST_FLAKES) -v -r ./tests/e2e $(MAKE) clean-mock-backend $(MAKE) clean-cloud-proxy-backend + $(MAKE) clean-localai-proxy-backend $(MAKE) teardown-e2e docker rmi localai-tests @@ -1330,6 +1331,7 @@ BACKEND_PIPER = piper|golang|.|false|true BACKEND_LOCAL_STORE = local-store|golang|.|false|true BACKEND_VALKEY_STORE = valkey-store|golang|.|false|true BACKEND_CLOUD_PROXY = cloud-proxy|golang|.|false|true +BACKEND_LOCALAI_PROXY = localai-proxy|golang|.|false|true BACKEND_HUGGINGFACE = huggingface|golang|.|false|true BACKEND_SILERO_VAD = silero-vad|golang|.|false|true BACKEND_STABLEDIFFUSION_GGML = stablediffusion-ggml|golang|.|--progress=plain|true @@ -1436,6 +1438,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_PIPER))) $(eval $(call generate-docker-build-target,$(BACKEND_LOCAL_STORE))) $(eval $(call generate-docker-build-target,$(BACKEND_VALKEY_STORE))) $(eval $(call generate-docker-build-target,$(BACKEND_CLOUD_PROXY))) +$(eval $(call generate-docker-build-target,$(BACKEND_LOCALAI_PROXY))) $(eval $(call generate-docker-build-target,$(BACKEND_HUGGINGFACE))) $(eval $(call generate-docker-build-target,$(BACKEND_SILERO_VAD))) $(eval $(call generate-docker-build-target,$(BACKEND_STABLEDIFFUSION_GGML))) @@ -1511,7 +1514,7 @@ $(eval $(call generate-docker-build-target,$(BACKEND_SUPERTONIC))) docker-save-%: backend-images docker save local-ai-backend:$* -o backend-images/$*.tar -docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp +docker-build-backends: docker-build-llama-cpp docker-build-ik-llama-cpp docker-build-turboquant docker-build-bonsai docker-build-ds4 docker-build-rerankers docker-build-vllm docker-build-vllm-omni docker-build-longcat-video docker-build-sglang docker-build-transformers docker-build-outetts docker-build-diffusers docker-build-kokoro docker-build-faster-whisper docker-build-crispasr docker-build-coqui docker-build-chatterbox docker-build-vibevoice docker-build-liquid-audio docker-build-moonshine docker-build-pocket-tts docker-build-qwen-tts docker-build-fish-speech docker-build-faster-qwen3-tts docker-build-qwen-asr docker-build-nemo docker-build-voxcpm docker-build-whisperx docker-build-ace-step docker-build-acestep-cpp docker-build-voxtral docker-build-mlx-distributed docker-build-trl docker-build-llama-cpp-quantization docker-build-tinygrad docker-build-kokoros docker-build-sam3-cpp docker-build-rfdetr-cpp docker-build-qwen3-tts-cpp docker-build-moss-tts-cpp docker-build-magpie-tts-cpp docker-build-vllm-cpp docker-build-omnivoice-cpp docker-build-vibevoice-cpp docker-build-localvqe docker-build-insightface docker-build-speaker-recognition docker-build-sherpa-onnx docker-build-cloud-proxy docker-build-localai-proxy docker-build-supertonic docker-build-depth-anything-cpp docker-build-moss-transcribe-cpp docker-build-nemo-speech-cpp docker-build-privacy-filter docker-build-trellis2cpp docker-build-valkey-store docker-build-audio-cpp docker-build-backends: docker-build-whisper-medusa docker-build-backends: docker-build-funasr @@ -1531,6 +1534,12 @@ build-cloud-proxy-backend: protogen-go clean-cloud-proxy-backend: rm -f tests/e2e/mock-backend/cloud-proxy +build-localai-proxy-backend: protogen-go + $(GOCMD) build -o tests/e2e/mock-backend/localai-proxy ./backend/go/localai-proxy + +clean-localai-proxy-backend: + rm -f tests/e2e/mock-backend/localai-proxy + ######################################################## ### UI E2E Test Server ######################################################## diff --git a/backend/go/localai-proxy/Makefile b/backend/go/localai-proxy/Makefile new file mode 100644 index 000000000..94fc5cf48 --- /dev/null +++ b/backend/go/localai-proxy/Makefile @@ -0,0 +1,13 @@ +GOCMD=go + +# Packaged as a standalone gallery backend by backend/Dockerfile.golang. +localai-proxy: + CGO_ENABLED=0 $(GOCMD) build -ldflags "$(LD_FLAGS)" -tags "$(GO_TAGS)" -o localai-proxy ./ + +package: + bash package.sh + +build: localai-proxy package + +clean: + rm -f localai-proxy diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go new file mode 100644 index 000000000..3c8b252d3 --- /dev/null +++ b/backend/go/localai-proxy/client.go @@ -0,0 +1,220 @@ +package main + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "mime/multipart" + "net/http" + "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, 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 { + cfg, err := p.config() + if err != nil { + return err + } + var file *os.File + if fileField != "" { + // Open before contacting the upstream so a bad path is reported as a + // request error, not as a failure of the remote host. + if file, err = os.Open(filePath); err != nil { + return status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", filePath, err) + } + defer func() { _ = file.Close() }() + } + ctx, cancel := withTimeout(ctx, cfg) + defer cancel() + + // Stream the form through a pipe so large audio files are not buffered + // in memory. The transport closes pr when the request ends, which + // unblocks the writer on every error path. + pr, pw := io.Pipe() + mw := multipart.NewWriter(pw) + go func() { + pw.CloseWithError(writeMultipart(mw, fields, fileField, file)) + }() + + req, err := p.newRequest(ctx, cfg, path, pr) + if err != nil { + _ = pr.CloseWithError(err) + return err + } + req.Header.Set("Content-Type", mw.FormDataContentType()) + return p.do(req, path, out) +} + +func writeMultipart(mw *multipart.Writer, fields map[string]string, fileField string, file *os.File) error { + for k, v := range fields { + if err := mw.WriteField(k, v); err != nil { + return err + } + } + if file != nil { + part, err := mw.CreateFormFile(fileField, 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, path, bytes.NewReader(payload)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + 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, path string, body io.Reader) (*http.Request, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, 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). 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 >= 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)) +} diff --git a/backend/go/localai-proxy/fake_upstream_test.go b/backend/go/localai-proxy/fake_upstream_test.go new file mode 100644 index 000000000..6ced223a8 --- /dev/null +++ b/backend/go/localai-proxy/fake_upstream_test.go @@ -0,0 +1,160 @@ +package main + +import ( + "encoding/json" + "io" + "mime" + "mime/multipart" + "net/http" + "net/http/httptest" + "strings" + "sync" + + . "github.com/onsi/gomega" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// recordedRequest is what the fake upstream saw for one call. JSON bodies land +// in JSON; multipart bodies land in Fields and Files (field name to content). +type recordedRequest struct { + Method string + Path string + Auth string + JSON map[string]any + Fields map[string]string + Files map[string]string +} + +// scriptedResponse is the reply for one path. SSE, when set, is written as +// "data: " events and wins over Body. +type scriptedResponse struct { + Status int + ContentType 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) + } + w.WriteHeader(resp.Status) + _, _ = io.WriteString(w, resp.Body) +} + +// loadProxy returns a proxy loaded against the fake upstream with the given +// proxy options merged over sane defaults. +func loadProxy(f *fakeUpstream, mutate func(*pb.ModelOptions)) *LocalAIProxy { + opts := &pb.ModelOptions{ + Model: "local-name", + Proxy: &pb.ProxyOptions{UpstreamUrl: f.URL + "/", UpstreamModel: "remote-model"}, + } + if mutate != nil { + mutate(opts) + } + p := NewLocalAIProxy() + ExpectWithOffset(1, p.Load(opts)).To(Succeed()) + return p +} + +// sseJSON marshals v for use as one SSE frame. +func sseJSON(v any) string { + raw, err := json.Marshal(v) + Expect(err).NotTo(HaveOccurred()) + return strings.TrimSpace(string(raw)) +} diff --git a/backend/go/localai-proxy/localai_proxy_suite_test.go b/backend/go/localai-proxy/localai_proxy_suite_test.go new file mode 100644 index 000000000..62edc82f3 --- /dev/null +++ b/backend/go/localai-proxy/localai_proxy_suite_test.go @@ -0,0 +1,17 @@ +package main + +import ( + "testing" + + "github.com/mudler/xlog" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestLocalAIProxy(t *testing.T) { + RegisterFailHandler(Fail) + // The specs drive upstream failures on purpose; their warnings are + // expected and would only bury real failures in the output. + xlog.SetLogger(xlog.NewLogger(xlog.LogLevelError, xlog.TextFormat)) + RunSpecs(t, "localai-proxy specs") +} diff --git a/backend/go/localai-proxy/main.go b/backend/go/localai-proxy/main.go new file mode 100644 index 000000000..7c3c3f0e8 --- /dev/null +++ b/backend/go/localai-proxy/main.go @@ -0,0 +1,32 @@ +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) + } +} diff --git a/backend/go/localai-proxy/package.sh b/backend/go/localai-proxy/package.sh new file mode 100755 index 000000000..7bc9f64aa --- /dev/null +++ b/backend/go/localai-proxy/package.sh @@ -0,0 +1,13 @@ +#!/bin/bash + +# Script to copy the localai-proxy binary into the package dir for the +# final Dockerfile stage. Mirrors backend/go/local-store/package.sh — +# no extra runtime libs needed since the backend is pure Go. + +set -e + +CURDIR=$(dirname "$(realpath $0)") + +mkdir -p $CURDIR/package +cp -avf $CURDIR/localai-proxy $CURDIR/package/ +cp -rfv $CURDIR/run.sh $CURDIR/package/ diff --git a/backend/go/localai-proxy/proxy.go b/backend/go/localai-proxy/proxy.go new file mode 100644 index 000000000..8fcd715f3 --- /dev/null +++ b/backend/go/localai-proxy/proxy.go @@ -0,0 +1,226 @@ +package main + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/url" + "os" + "strings" + "sync/atomic" + "time" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + "github.com/mudler/LocalAI/pkg/grpc/base" + "github.com/mudler/LocalAI/pkg/grpc/grpcerrors" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/httpclient" +) + +const ( + backendName = "localai-proxy" + + // realtimePipelineOption names the upstream realtime pipeline that serves + // live transcription sessions (options: ["realtime_pipeline:"]). + realtimePipelineOption = "realtime_pipeline:" +) + +// LocalAIProxy serves backend methods by calling a remote LocalAI's REST API. +// base.SingleThread is not embedded: every call is an independent HTTP +// request, so serialising them would only add latency. +type LocalAIProxy struct { + base.Base + + cfg atomic.Pointer[proxyConfig] + client *http.Client +} + +type proxyConfig struct { + base string // upstream base URL without a trailing slash + upstreamModel string // model name sent upstream + apiKey string + realtimePipeline string + timeout time.Duration // per-request limit for non-streaming calls; 0 = none +} + +func NewLocalAIProxy() *LocalAIProxy { + // httpclient.New refuses redirects: the upstream is one configured + // LocalAI, so a 3xx means misconfiguration or a hijacked host, and + // following it would replay the bearer key to an unvetted host. It also + // sets no body deadline, so long SSE streams are not cut short. + return &LocalAIProxy{client: httpclient.New()} +} + +// Load refuses a model without proxy options so greedy backend probing, +// which tries every installed backend on a model file, never selects it. +func (p *LocalAIProxy) Load(opts *pb.ModelOptions) error { + po := opts.GetProxy() + if po == nil { + return errors.New("localai-proxy: Load requires proxy options (proxy.upstream_url)") + } + raw := po.GetUpstreamUrl() + if raw == "" { + return errors.New("localai-proxy: proxy.upstream_url is required") + } + u, err := url.ParseRequestURI(raw) + if err != nil { + return fmt.Errorf("localai-proxy: proxy.upstream_url %q invalid: %w", raw, err) + } + if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" { + return fmt.Errorf("localai-proxy: proxy.upstream_url %q must be an http(s) URL with a host", raw) + } + + // 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: strings.TrimRight(raw, "/"), + upstreamModel: model, + apiKey: key, + realtimePipeline: pipeline, + timeout: timeout, + }) + xlog.Info("localai-proxy: ready", "upstream", raw, "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 != "" { + b, err := os.ReadFile(filePath) + if err != nil { + return "", fmt.Errorf("localai-proxy: read api_key_file %q: %w", filePath, err) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} + +// unimplemented is the error for methods LocalAI's REST API cannot serve. +// Failover reads gRPC Unimplemented as a capability gap and moves to the next +// target without marking this one unhealthy. +func unimplemented(method string) error { + return status.Errorf(codes.Unimplemented, "localai-proxy: %s has no upstream counterpart", method) +} + +func (p *LocalAIProxy) AudioEncode(*pb.AudioEncodeRequest) (*pb.AudioEncodeResult, error) { + return nil, unimplemented("AudioEncode") +} + +func (p *LocalAIProxy) AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error) { + return nil, unimplemented("AudioDecode") +} + +// AudioToAudioStream closes out because the gRPC server drains it until +// closed; leaving it open would hang the call. +func (p *LocalAIProxy) AudioToAudioStream(_ <-chan *pb.AudioToAudioRequest, out chan<- *pb.AudioToAudioResponse) error { + close(out) + return unimplemented("AudioToAudioStream") +} + +func (p *LocalAIProxy) TokenClassify(context.Context, *pb.TokenClassifyRequest) (*pb.TokenClassifyResponse, error) { + return nil, unimplemented("TokenClassify") +} + +func (p *LocalAIProxy) ModelMetadata(*pb.ModelOptions) (*pb.ModelMetadataResponse, error) { + return nil, unimplemented("ModelMetadata") +} + +func (p *LocalAIProxy) StartFineTune(*pb.FineTuneRequest) (*pb.FineTuneJobResult, error) { + return nil, unimplemented("StartFineTune") +} + +// FineTuneProgress closes the channel: the gRPC server waits for it to close +// before returning, and base.Base leaves it open. +func (p *LocalAIProxy) FineTuneProgress(_ *pb.FineTuneProgressRequest, updates chan *pb.FineTuneProgressUpdate) error { + close(updates) + return unimplemented("FineTuneProgress") +} + +func (p *LocalAIProxy) StopFineTune(*pb.FineTuneStopRequest) error { + return unimplemented("StopFineTune") +} + +func (p *LocalAIProxy) ListCheckpoints(*pb.ListCheckpointsRequest) (*pb.ListCheckpointsResponse, error) { + return nil, unimplemented("ListCheckpoints") +} + +func (p *LocalAIProxy) ExportModel(*pb.ExportModelRequest) error { + return unimplemented("ExportModel") +} + +func (p *LocalAIProxy) StartQuantization(*pb.QuantizationRequest) (*pb.QuantizationJobResult, error) { + return nil, unimplemented("StartQuantization") +} + +// QuantizationProgress closes the channel for the same reason as +// FineTuneProgress. +func (p *LocalAIProxy) QuantizationProgress(_ *pb.QuantizationProgressRequest, updates chan *pb.QuantizationProgressUpdate) error { + close(updates) + return unimplemented("QuantizationProgress") +} + +func (p *LocalAIProxy) StopQuantization(*pb.QuantizationStopRequest) error { + return unimplemented("StopQuantization") +} diff --git a/backend/go/localai-proxy/run.sh b/backend/go/localai-proxy/run.sh new file mode 100755 index 000000000..f2023e3ec --- /dev/null +++ b/backend/go/localai-proxy/run.sh @@ -0,0 +1,6 @@ +#!/bin/bash +set -ex + +CURDIR=$(dirname "$(realpath "$0")") + +exec "$CURDIR"/localai-proxy "$@" diff --git a/backend/go/localai-proxy/text.go b/backend/go/localai-proxy/text.go new file mode 100644 index 000000000..c388cc5f7 --- /dev/null +++ b/backend/go/localai-proxy/text.go @@ -0,0 +1,362 @@ +package main + +import ( + "bufio" + "context" + "encoding/json" + "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. +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,omitempty"` + 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"` +} + +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 { + 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) { + 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 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 +} + +// 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: int32(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) { + path, body := p.textRequest(opts, false) + var resp textResponse + if err := p.postJSON(context.Background(), 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 { + path, body := p.textRequest(opts, true) + resp, err := p.postStream(context.Background(), 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.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()} + 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(), + "top_n": in.GetTopN(), + } + 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: int32(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 +} diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go new file mode 100644 index 000000000..bd12770b8 --- /dev/null +++ b/backend/go/localai-proxy/text_test.go @@ -0,0 +1,483 @@ +package main + +import ( + "context" + "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)) + }) + + It("falls back to the model name when upstream_model is unset", func() { + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamModel = "" }) + Expect(p.model("")).To(Equal("local-name")) + }) + }) + + Describe("PredictRich", func() { + It("sends messages to /v1/chat/completions with the upstream model and key", func() { + GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-test") + p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" }) + up.replyJSON("/v1/chat/completions", map[string]any{ + "choices": []any{map[string]any{"message": map[string]any{"role": "assistant", "content": "hello back"}}}, + "usage": map[string]any{"prompt_tokens": 3, "completion_tokens": 2}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{ + Messages: []*pb.Message{{Role: "user", Content: "hello"}}, + Tokens: 32, + Temperature: 0.5, + TopK: 40, + StopPrompts: []string{""}, + Seed: 9, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(string(reply.GetMessage())).To(Equal("hello back")) + Expect(reply.GetPromptTokens()).To(Equal(int32(3))) + Expect(reply.GetTokens()).To(Equal(int32(2))) + + req := up.last() + Expect(req.Method).To(Equal(http.MethodPost)) + Expect(req.Path).To(Equal("/v1/chat/completions")) + Expect(req.Auth).To(Equal("Bearer sk-test")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + Expect(req.JSON).To(HaveKeyWithValue("max_tokens", BeNumerically("==", 32))) + Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0.5))) + Expect(req.JSON).To(HaveKeyWithValue("top_k", BeNumerically("==", 40))) + Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 9))) + Expect(req.JSON).To(HaveKeyWithValue("stop", ConsistOf(""))) + Expect(req.JSON).NotTo(HaveKey("stream")) + Expect(req.JSON["messages"]).To(ConsistOf(HaveKeyWithValue("content", "hello"))) + }) + + It("returns upstream tool calls as chat deltas", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/chat/completions", map[string]any{ + "choices": []any{map[string]any{"message": map[string]any{ + "role": "assistant", + "tool_calls": []any{map[string]any{ + "id": "call_1", "type": "function", + "function": map[string]any{"name": "get_weather", "arguments": `{"city":"Rome"}`}, + }}, + }}}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{ + Messages: []*pb.Message{{Role: "user", Content: "weather?"}}, + Tools: `[{"type":"function","function":{"name":"get_weather"}}]`, + }) + Expect(err).NotTo(HaveOccurred()) + Expect(reply.GetChatDeltas()).To(HaveLen(1)) + tc := reply.GetChatDeltas()[0].GetToolCalls() + Expect(tc).To(HaveLen(1)) + Expect(tc[0].GetName()).To(Equal("get_weather")) + Expect(tc[0].GetArguments()).To(Equal(`{"city":"Rome"}`)) + Expect(up.last().JSON).To(HaveKey("tools")) + }) + + It("sends a bare prompt to /v1/completions", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/completions", map[string]any{ + "choices": []any{map[string]any{"text": "completed"}}, + }) + + reply, err := p.PredictRich(&pb.PredictOptions{Prompt: "once upon"}) + Expect(err).NotTo(HaveOccurred()) + Expect(string(reply.GetMessage())).To(Equal("completed")) + req := up.last() + Expect(req.Path).To(Equal("/v1/completions")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "once upon")) + Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model")) + }) + + It("maps a 5xx upstream to Unavailable with the body in the message", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusServiceUnavailable, Body: "backend is down"}) + + _, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(err.Error()).To(ContainSubstring("backend is down")) + }) + + It("maps a 4xx upstream to InvalidArgument", func() { + p := loadProxy(up, nil) + up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad request"}) + + _, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}}) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + Expect(err.Error()).To(ContainSubstring("bad request")) + }) + + It("truncates a long upstream error body", func() { + p := loadProxy(up, nil) + long := make([]byte, 2000) + for i := range long { + long[i] = 'a' + } + up.script("/v1/completions", scriptedResponse{Status: http.StatusInternalServerError, Body: string(long)}) + + _, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700)) + }) + + It("maps 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("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("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"})) + }) + }) + + 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("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")) + }) + }) + + 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()) + }) + }) +}) diff --git a/backend/index.yaml b/backend/index.yaml index 8b2efd417..d399d6498 100644 --- a/backend/index.yaml +++ b/backend/index.yaml @@ -2038,6 +2038,21 @@ capabilities: default: "cpu-cloud-proxy" metal: "metal-cloud-proxy" +- &localai-proxy + name: "localai-proxy" + alias: "localai-proxy" + urls: + - https://github.com/mudler/LocalAI/tree/master/backend/go/localai-proxy + description: | + Serve a model from another LocalAI instance: text, embeddings, rerank, audio, image and video requests are forwarded to its REST API. + tags: + - text-to-text + - proxy + - CPU + license: MIT + capabilities: + default: "cpu-localai-proxy" + metal: "metal-localai-proxy" - &valkey-store name: "valkey-store" urls: @@ -2623,6 +2638,31 @@ uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-cloud-proxy" mirrors: - localai/localai-backends:master-metal-darwin-arm64-cloud-proxy +- !!merge <<: *localai-proxy + name: "cpu-localai-proxy" + uri: "quay.io/go-skynet/local-ai-backends:latest-cpu-localai-proxy" + mirrors: + - localai/localai-backends:latest-cpu-localai-proxy +- !!merge <<: *localai-proxy + name: "cpu-localai-proxy-development" + uri: "quay.io/go-skynet/local-ai-backends:master-cpu-localai-proxy" + mirrors: + - localai/localai-backends:master-cpu-localai-proxy +- !!merge <<: *localai-proxy + name: "localai-proxy-development" + capabilities: + default: "cpu-localai-proxy-development" + metal: "metal-localai-proxy-development" +- !!merge <<: *localai-proxy + name: "metal-localai-proxy" + uri: "quay.io/go-skynet/local-ai-backends:latest-metal-darwin-arm64-localai-proxy" + mirrors: + - localai/localai-backends:latest-metal-darwin-arm64-localai-proxy +- !!merge <<: *localai-proxy + name: "metal-localai-proxy-development" + uri: "quay.io/go-skynet/local-ai-backends:master-metal-darwin-arm64-localai-proxy" + mirrors: + - localai/localai-backends:master-metal-darwin-arm64-localai-proxy - !!merge <<: *valkey-store name: "cpu-valkey-store" alias: "valkey-store" From 6df1767133f1699ee45cdf498590d80c72fee03a Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 23:13:22 +0000 Subject: [PATCH 43/79] fix(localai-proxy): keep rerank, stream errors and chat intact through the proxy Rerank no longer sends top_n 0, which the upstream rejects. A mid-stream upstream error frame now fails the call instead of ending it as a short success. Temperature 0 is forwarded. An upstream 429 becomes ResourceExhausted, which failover skips like Unimplemented. A localai-proxy config sends its own name upstream when upstream_model is unset, and a chat proxy defaults to the tokenizer template so chat reaches the upstream as messages. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/client.go | 6 ++- backend/go/localai-proxy/text.go | 46 +++++++++++++++++++++- backend/go/localai-proxy/text_test.go | 52 +++++++++++++++++++++++++ core/backend/options.go | 8 ++++ core/backend/options_internal_test.go | 27 +++++++++++++ core/config/hooks_localai_proxy.go | 28 +++++++++++++ core/config/hooks_test.go | 32 +++++++++++++++ core/http/middleware/failover_test.go | 12 ++++++ core/services/failover/classify.go | 19 ++++++--- core/services/failover/classify_test.go | 2 + core/services/failover/manager.go | 5 ++- core/services/failover/manager_test.go | 15 +++++++ 12 files changed, 242 insertions(+), 10 deletions(-) create mode 100644 core/config/hooks_localai_proxy.go diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go index 3c8b252d3..5d07562ec 100644 --- a/backend/go/localai-proxy/client.go +++ b/backend/go/localai-proxy/client.go @@ -189,7 +189,7 @@ func transportError(path string, err error) error { // 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). 501 is the upstream saying it +// trip a healthy target over a client error), except 429. 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 { @@ -208,6 +208,10 @@ func statusError(path string, resp *http.Response) error { code = codes.Unimplemented case resp.StatusCode >= 500: code = codes.Unavailable + case resp.StatusCode == http.StatusTooManyRequests: + // Rate limited: the upstream is healthy but out of capacity, so + // failover skips to the next target without tripping this one. + code = codes.ResourceExhausted case resp.StatusCode >= 400: code = codes.InvalidArgument default: diff --git a/backend/go/localai-proxy/text.go b/backend/go/localai-proxy/text.go index c388cc5f7..3b73defa2 100644 --- a/backend/go/localai-proxy/text.go +++ b/backend/go/localai-proxy/text.go @@ -16,13 +16,15 @@ import ( // 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,omitempty"` + Temperature float32 `json:"temperature"` TopP float32 `json:"top_p,omitempty"` TopK int32 `json:"top_k,omitempty"` Seed int32 `json:"seed,omitempty"` @@ -67,6 +69,11 @@ type choiceDelta struct { } 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"` @@ -79,6 +86,9 @@ type textResponse struct { // 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, @@ -113,6 +123,27 @@ func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string 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 { @@ -207,6 +238,13 @@ func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<- 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 @@ -273,7 +311,11 @@ func (p *LocalAIProxy) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.Re "model": p.model(""), "query": in.GetQuery(), "documents": in.GetDocuments(), - "top_n": in.GetTopN(), + } + // 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 { diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go index bd12770b8..b47051b4c 100644 --- a/backend/go/localai-proxy/text_test.go +++ b/backend/go/localai-proxy/text_test.go @@ -176,6 +176,27 @@ var _ = Describe("localai-proxy", func() { Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700)) }) + It("maps a 429 upstream to ResourceExhausted so failover skips without tripping", 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() @@ -229,6 +250,21 @@ var _ = Describe("localai-proxy", func() { Expect(string((<-results).GetMessage())).To(Equal("b")) }) + 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"}) @@ -260,6 +296,13 @@ var _ = Describe("localai-proxy", func() { }) }) + 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) @@ -314,6 +357,15 @@ var _ = Describe("localai-proxy", func() { }) }) + 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) diff --git a/core/backend/options.go b/core/backend/options.go index 12f2e29a1..0f345e154 100644 --- a/core/backend/options.go +++ b/core/backend/options.go @@ -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 == "localai-proxy" && opts.Proxy.UpstreamModel == "" { + opts.Proxy.UpstreamModel = c.Name + } } if c.MMProj != "" { diff --git a/core/backend/options_internal_test.go b/core/backend/options_internal_test.go index 5f8370211..13ad13b57 100644 --- a/core/backend/options_internal_test.go +++ b/core/backend/options_internal_test.go @@ -62,6 +62,33 @@ var _ = Describe("grpcModelOpts Proxy options", func() { 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") diff --git a/core/config/hooks_localai_proxy.go b/core/config/hooks_localai_proxy.go new file mode 100644 index 000000000..b1dd3bb36 --- /dev/null +++ b/core/config/hooks_localai_proxy.go @@ -0,0 +1,28 @@ +package config + +func init() { + RegisterBackendHook("localai-proxy", localAIProxyDefaults) +} + +// localAIProxyDefaults makes chat requests reach the upstream as structured +// messages. Without the tokenizer template core renders the prompt itself +// with no template, and the proxy can only send that text to +// /v1/completions, bypassing the upstream model's chat template, tool +// handling and reasoning parsing. +// +// Only configs that declare the chat usecase get it: usecase guessing reads +// the tokenizer template as "this model chats", so setting it on a +// transcription or embedding proxy would offer that model to chat pickers +// and default-model selection. A config that brings its own templates keeps +// them: the operator chose local templating. +func localAIProxyDefaults(cfg *ModelConfig, _ string) { + t := cfg.TemplateConfig + if t.UseTokenizerTemplate || t.Chat != "" || t.ChatMessage != "" || t.Completion != "" || t.Edit != "" { + return + } + declared := GetUsecasesFromYAML(cfg.KnownUsecaseStrings) + if declared == nil || *declared&FLAG_CHAT != FLAG_CHAT { + return + } + cfg.TemplateConfig.UseTokenizerTemplate = true +} diff --git a/core/config/hooks_test.go b/core/config/hooks_test.go index b69bc6989..60cc8fd82 100644 --- a/core/config/hooks_test.go +++ b/core/config/hooks_test.go @@ -127,6 +127,38 @@ var _ = Describe("Backend hooks and parser defaults", func() { }) }) + Context("localai-proxy hook", func() { + It("defaults a chat proxy to the tokenizer template so chat sends messages upstream", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeTrue()) + }) + + It("leaves non-chat and undeclared proxies alone so they are not guessed as chat", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"transcript"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + Expect(cfg.HasUsecases(FLAG_CHAT)).To(BeFalse()) + + bare := &ModelConfig{Backend: "localai-proxy"} + bare.SetDefaults() + Expect(bare.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + + It("keeps a config that brings its own templates", func() { + cfg := &ModelConfig{Backend: "localai-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.TemplateConfig.Chat = "{{.Input}}" + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + + It("does not touch other backends", func() { + cfg := &ModelConfig{Backend: "cloud-proxy", KnownUsecaseStrings: []string{"chat"}} + cfg.SetDefaults() + Expect(cfg.TemplateConfig.UseTokenizerTemplate).To(BeFalse()) + }) + }) + Context("vllmDefaults hook", func() { It("auto-sets parsers for known model families on vllm backend", func() { cfg := &ModelConfig{ diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index 94588bdf9..00439abcc 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -293,6 +293,18 @@ var _ = Describe("failover chains in the request pipeline", func() { Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) }) + It("spills a rate-limited target (gRPC ResourceExhausted) to the next target without tripping 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.StateHealthy)) + }) + It("skips a disabled target without tripping it", func() { rec := chat("chain-off") Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) diff --git a/core/services/failover/classify.go b/core/services/failover/classify.go index a71d89d3f..3a50299b6 100644 --- a/core/services/failover/classify.go +++ b/core/services/failover/classify.go @@ -61,16 +61,25 @@ func IsRetryable(err error, status int) bool { return !isRequestError(msg) } -// IsCapabilityGap reports a target that cannot serve this kind of request at -// all (gRPC Unimplemented, anywhere in the error chain). The next target may -// serve it, and this target is not broken: the failure carries no signal -// about its health, so callers must skip it without tripping. +// IsCapabilityGap reports a target that cannot serve this request right now +// for a reason that says nothing about its health: it cannot serve this kind +// of request at all (gRPC Unimplemented), or it is out of capacity, such as a +// rate-limited upstream (gRPC ResourceExhausted, what localai-proxy returns +// for an upstream 429). Matched anywhere in the error chain. The next target +// may serve it, so callers must skip this one without tripping it. func IsCapabilityGap(err error) bool { if err == nil { return false } st, ok := grpcstatus.FromError(err) - return ok && st.Code() == codes.Unimplemented + if !ok { + return false + } + switch st.Code() { + case codes.Unimplemented, codes.ResourceExhausted: + return true + } + return false } func retryableStatus(code int) bool { diff --git a/core/services/failover/classify_test.go b/core/services/failover/classify_test.go index cc9994e70..baafff337 100644 --- a/core/services/failover/classify_test.go +++ b/core/services/failover/classify_test.go @@ -45,6 +45,8 @@ var _ = DescribeTable("IsCapabilityGap", Entry("nil error", nil, false), Entry("grpc unimplemented", grpcstatus.Error(codes.Unimplemented, "x"), true), Entry("wrapped grpc unimplemented", fmt.Errorf("call: %w", grpcstatus.Error(codes.Unimplemented, "x")), true), + Entry("grpc resource exhausted", grpcstatus.Error(codes.ResourceExhausted, "x"), true), + Entry("wrapped grpc resource exhausted", fmt.Errorf("call: %w", grpcstatus.Error(codes.ResourceExhausted, "x")), true), Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), false), Entry("plain error", errors.New("boom"), false), ) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 071e9ab0c..147b95b34 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -728,8 +728,9 @@ func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Cont att.Succeed() return nil case IsCapabilityGap(err) && !committed.Load(): - // This target cannot serve this kind of request at all; it is - // not broken, so move on without counting a failure. + // This target cannot serve this kind of request, or is out of + // capacity; it is not broken, so move on without counting a + // failure. if !att.Skip() { return err } diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 04f29d4db..948f5aef0 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -305,5 +305,20 @@ var _ = Describe("Manager", func() { st, _ := m.ChainStatus("chain") Expect(st.Targets[0].State).To(Equal(StateHealthy)) }) + + It("skips a rate-limited (ResourceExhausted) target without tripping it", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, _ func()) error { + tried = append(tried, target) + if target == "a" { + return grpcstatus.Error(codes.ResourceExhausted, "localai-proxy: upstream /v1/rerank returned 429: slow down") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) }) }) From 882f51fd62ce3c9ab429ad16b2dc55755ad71dbe Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 23:20:31 +0000 Subject: [PATCH 44/79] fix(failover): trip a rate-limited or exhausted target instead of skipping it Treating ResourceExhausted as a capability gap skipped the target without counting a failure, so a target that stays rate limited or out of memory kept its traffic. It is now an ordinary retryable failure: the request moves to the next target and the exhausted one trips. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/client.go | 12 +++++++----- backend/go/localai-proxy/text_test.go | 2 +- core/http/middleware/failover_test.go | 4 ++-- core/services/failover/classify.go | 25 ++++++++++--------------- core/services/failover/classify_test.go | 4 ++-- core/services/failover/manager.go | 5 ++--- core/services/failover/manager_test.go | 4 ++-- 7 files changed, 26 insertions(+), 30 deletions(-) diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go index 5d07562ec..100a931dd 100644 --- a/backend/go/localai-proxy/client.go +++ b/backend/go/localai-proxy/client.go @@ -189,9 +189,10 @@ func transportError(path string, err error) error { // 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), except 429. 501 is the upstream saying it -// cannot serve this kind of request, which failover treats as a capability -// gap, like our own Unimplemented methods. +// 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)) @@ -209,8 +210,9 @@ func statusError(path string, resp *http.Response) error { case resp.StatusCode >= 500: code = codes.Unavailable case resp.StatusCode == http.StatusTooManyRequests: - // Rate limited: the upstream is healthy but out of capacity, so - // failover skips to the next target without tripping this one. + // 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 diff --git a/backend/go/localai-proxy/text_test.go b/backend/go/localai-proxy/text_test.go index b47051b4c..53ad73c9b 100644 --- a/backend/go/localai-proxy/text_test.go +++ b/backend/go/localai-proxy/text_test.go @@ -176,7 +176,7 @@ var _ = Describe("localai-proxy", func() { Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700)) }) - It("maps a 429 upstream to ResourceExhausted so failover skips without tripping", func() { + 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"}) diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index 00439abcc..173468d73 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -293,7 +293,7 @@ var _ = Describe("failover chains in the request pipeline", func() { Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) }) - It("spills a rate-limited target (gRPC ResourceExhausted) to the next target without tripping it", func() { + 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") } @@ -302,7 +302,7 @@ var _ = Describe("failover chains in the request pipeline", func() { 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)) + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) }) It("skips a disabled target without tripping it", func() { diff --git a/core/services/failover/classify.go b/core/services/failover/classify.go index 3a50299b6..ac45c3c20 100644 --- a/core/services/failover/classify.go +++ b/core/services/failover/classify.go @@ -46,7 +46,11 @@ func IsRetryable(err error, status int) bool { } if st, ok := grpcstatus.FromError(err); ok { switch st.Code() { - case codes.Unavailable, codes.Internal, codes.DeadlineExceeded, codes.Unknown: + // ResourceExhausted (a rate-limited upstream, what localai-proxy + // returns for a 429, or a backend out of memory) is retried + // elsewhere and trips the target: a gap would skip a chronically + // exhausted target forever without moving traffic off it. + case codes.Unavailable, codes.Internal, codes.DeadlineExceeded, codes.Unknown, codes.ResourceExhausted: return !isRequestError(st.Message()) default: return false @@ -61,25 +65,16 @@ func IsRetryable(err error, status int) bool { return !isRequestError(msg) } -// IsCapabilityGap reports a target that cannot serve this request right now -// for a reason that says nothing about its health: it cannot serve this kind -// of request at all (gRPC Unimplemented), or it is out of capacity, such as a -// rate-limited upstream (gRPC ResourceExhausted, what localai-proxy returns -// for an upstream 429). Matched anywhere in the error chain. The next target -// may serve it, so callers must skip this one without tripping it. +// IsCapabilityGap reports a target that cannot serve this kind of request at +// all (gRPC Unimplemented, anywhere in the error chain). The next target may +// serve it, and this target is not broken: the failure carries no signal +// about its health, so callers must skip it without tripping. func IsCapabilityGap(err error) bool { if err == nil { return false } st, ok := grpcstatus.FromError(err) - if !ok { - return false - } - switch st.Code() { - case codes.Unimplemented, codes.ResourceExhausted: - return true - } - return false + return ok && st.Code() == codes.Unimplemented } func retryableStatus(code int) bool { diff --git a/core/services/failover/classify_test.go b/core/services/failover/classify_test.go index baafff337..c19fd0aba 100644 --- a/core/services/failover/classify_test.go +++ b/core/services/failover/classify_test.go @@ -32,6 +32,7 @@ var _ = DescribeTable("IsRetryable", Entry("grpc deadline", grpcstatus.Error(codes.DeadlineExceeded, "x"), 0, true), Entry("grpc unknown", grpcstatus.Error(codes.Unknown, "x"), 0, true), Entry("grpc invalid argument", grpcstatus.Error(codes.InvalidArgument, "x"), 0, false), + Entry("grpc resource exhausted (rate limit, OOM)", grpcstatus.Error(codes.ResourceExhausted, "x"), 0, true), Entry("cloud-proxy upstream 503", errors.New("cloud-proxy: upstream 503: no healthy nodes"), 0, true), Entry("cloud-proxy upstream 429 stays 4xx", errors.New("cloud-proxy: upstream 429: slow down"), 0, false), Entry("context overflow", errors.New("the request exceeds the available context size"), 0, false), @@ -45,8 +46,7 @@ var _ = DescribeTable("IsCapabilityGap", Entry("nil error", nil, false), Entry("grpc unimplemented", grpcstatus.Error(codes.Unimplemented, "x"), true), Entry("wrapped grpc unimplemented", fmt.Errorf("call: %w", grpcstatus.Error(codes.Unimplemented, "x")), true), - Entry("grpc resource exhausted", grpcstatus.Error(codes.ResourceExhausted, "x"), true), - Entry("wrapped grpc resource exhausted", fmt.Errorf("call: %w", grpcstatus.Error(codes.ResourceExhausted, "x")), true), + Entry("grpc resource exhausted is a failure, not a gap", grpcstatus.Error(codes.ResourceExhausted, "x"), false), Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), false), Entry("plain error", errors.New("boom"), false), ) diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 147b95b34..071e9ab0c 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -728,9 +728,8 @@ func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Cont att.Succeed() return nil case IsCapabilityGap(err) && !committed.Load(): - // This target cannot serve this kind of request, or is out of - // capacity; it is not broken, so move on without counting a - // failure. + // This target cannot serve this kind of request at all; it is + // not broken, so move on without counting a failure. if !att.Skip() { return err } diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 948f5aef0..75a773f04 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -306,7 +306,7 @@ var _ = Describe("Manager", func() { Expect(st.Targets[0].State).To(Equal(StateHealthy)) }) - It("skips a rate-limited (ResourceExhausted) target without tripping it", func() { + It("fails a rate-limited (ResourceExhausted) target over and trips it", func() { var tried []string err := m.Do(context.Background(), "chain", func(_ context.Context, target string, _ func()) error { tried = append(tried, target) @@ -318,7 +318,7 @@ var _ = Describe("Manager", func() { Expect(err).ToNot(HaveOccurred()) Expect(tried).To(Equal([]string{"a", "b"})) st, _ := m.ChainStatus("chain") - Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Targets[0].State).To(Equal(StateDown)) }) }) }) From 87fd7da5315cfeab4ab1bcd8300b12a5d3ce78da Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 23:28:48 +0000 Subject: [PATCH 45/79] feat(localai-proxy): serve speech, transcription and audio APIs remotely The proxy now forwards TTS, streaming TTS, sound generation, transcription (plain and streaming), diarization, VAD, sound detection and audio transforms to the upstream LocalAI. Streaming TTS passes the upstream WAV bytes through unchanged. A streaming transcription that stops before its final frame, or sends an error frame, fails with Unavailable instead of ending as a short success. Transcription always sends diarize, because the upstream treats a missing field as true. Audio transforms also download the separation stems the upstream names and write them beside Dst. Sound generation from a source clip returns Unimplemented, because the REST endpoint has no field for the clip. The multipart helper now takes repeated fields and several files. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/audio.go | 453 ++++++++++++++++ backend/go/localai-proxy/audio_test.go | 489 ++++++++++++++++++ backend/go/localai-proxy/client.go | 213 ++++++-- .../go/localai-proxy/fake_upstream_test.go | 6 +- 4 files changed, 1128 insertions(+), 33 deletions(-) create mode 100644 backend/go/localai-proxy/audio.go create mode 100644 backend/go/localai-proxy/audio_test.go diff --git a/backend/go/localai-proxy/audio.go b/backend/go/localai-proxy/audio.go new file mode 100644 index 000000000..07e011209 --- /dev/null +++ b/backend/go/localai-proxy/audio.go @@ -0,0 +1,453 @@ +package main + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "io" + "net/url" + "path/filepath" + "strconv" + "strings" + "time" + + "github.com/mudler/xlog" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +const transcriptionsPath = "/v1/audio/transcriptions" + +// ttsRequest is the body of LocalAI's /tts (schema.TTSRequest). +type ttsRequest struct { + Model string `json:"model"` + Input string `json:"input"` + Voice string `json:"voice,omitempty"` + Language string `json:"language,omitempty"` + Instructions string `json:"instructions,omitempty"` + Params map[string]string `json:"params,omitempty"` + Stream bool `json:"stream,omitempty"` +} + +func (p *LocalAIProxy) ttsRequest(req *pb.TTSRequest, stream bool) ttsRequest { + return ttsRequest{ + Model: p.model(""), + Input: req.GetText(), + Voice: req.GetVoice(), + Language: req.GetLanguage(), + Instructions: req.GetInstructions(), + Params: req.GetParams(), + Stream: stream, + } +} + +func (p *LocalAIProxy) TTS(req *pb.TTSRequest) error { + return p.postJSONToFile(context.Background(), "/tts", p.ttsRequest(req, false), req.GetDst()) +} + +// TTSStream forwards the upstream's chunked audio unchanged. That body is +// already what core expects from a streaming backend: a WAV header followed +// by PCM. out is closed on every path because the gRPC server drains it until +// closed and would otherwise hang. +func (p *LocalAIProxy) TTSStream(req *pb.TTSRequest, out chan []byte) error { + defer close(out) + resp, err := p.postStream(context.Background(), "/tts", p.ttsRequest(req, true)) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + buf := make([]byte, 32*1024) + for { + n, err := resp.Body.Read(buf) + if n > 0 { + // The reader reuses buf, so each chunk needs its own copy. + out <- append([]byte(nil), buf[:n]...) + } + if errors.Is(err, io.EOF) { + return nil + } + if err != nil { + // A cut-off stream is a failed synthesis, not a short one. + return transportError("/tts", err) + } + } +} + +// soundGenerationRequest is the body of /v1/sound-generation +// (schema.ElevenLabsSoundGenerationRequest). Pointers keep "unset" distinct +// from zero so the upstream model's defaults apply. +type soundGenerationRequest struct { + ModelID string `json:"model_id"` + Text string `json:"text"` + Duration *float32 `json:"duration_seconds,omitempty"` + Temperature *float32 `json:"prompt_influence,omitempty"` + DoSample *bool `json:"do_sample,omitempty"` + Think *bool `json:"think,omitempty"` + Caption string `json:"caption,omitempty"` + Lyrics string `json:"lyrics,omitempty"` + BPM *int32 `json:"bpm,omitempty"` + Keyscale string `json:"keyscale,omitempty"` + Language string `json:"language,omitempty"` + Timesignature string `json:"timesignature,omitempty"` + Instrumental *bool `json:"instrumental,omitempty"` +} + +// SoundGeneration refuses audio-conditioned requests: the REST endpoint has no +// field for a source clip, and dropping it would return unconditioned audio as +// if it were the answer. +func (p *LocalAIProxy) SoundGeneration(req *pb.SoundGenerationRequest) error { + if req.GetSrc() != "" { + return unimplemented("SoundGeneration with src") + } + body := soundGenerationRequest{ + ModelID: p.model(""), + Text: req.GetText(), + Duration: req.Duration, + Temperature: req.Temperature, + DoSample: req.Sample, + Think: req.Think, + Caption: req.GetCaption(), + Lyrics: req.GetLyrics(), + BPM: req.Bpm, + Keyscale: req.GetKeyscale(), + Language: req.GetLanguage(), + Timesignature: req.GetTimesignature(), + Instrumental: req.Instrumental, + } + return p.postJSONToFile(context.Background(), "/v1/sound-generation", body, req.GetDst()) +} + +// transcriptionForm builds the /v1/audio/transcriptions upload. Dst is the +// input audio (core names it that way). Language and translate are only sent +// when set so the upstream model config's defaults still apply; diarize is +// always sent because the upstream treats a missing field as true. +func (p *LocalAIProxy) transcriptionForm(req *pb.TranscriptRequest, stream bool) multipartForm { + f := url.Values{} + f.Set("model", p.model("")) + f.Set("diarize", strconv.FormatBool(req.GetDiarize())) + // verbose_json keeps segments and words; the default drops nothing today, + // but naming it keeps the reply shape pinned. + f.Set("response_format", "verbose_json") + if v := req.GetLanguage(); v != "" { + f.Set("language", v) + } + if req.GetTranslate() { + f.Set("translate", "true") + } + if v := req.GetPrompt(); v != "" { + f.Set("prompt", v) + } + if v := req.GetTemperature(); v != 0 { + f.Set("temperature", formatFloat(v)) + } + for _, g := range req.GetTimestampGranularities() { + f.Add("timestamp_granularities[]", g) + } + if stream { + f.Set("stream", "true") + } + return multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}} +} + +// transcriptionResult is TranscriptionResultSeconds: the REST API reports +// times in seconds, pb in nanoseconds (core reads them as time.Duration). +type transcriptionResult struct { + Text string `json:"text"` + Language string `json:"language"` + Duration float64 `json:"duration"` + Segments []transcriptionSegment `json:"segments"` +} + +type transcriptionSegment struct { + ID int32 `json:"id"` + Start float64 `json:"start"` + End float64 `json:"end"` + Text string `json:"text"` + Tokens []int32 `json:"tokens"` + Speaker string `json:"speaker"` + Words []struct { + Start float64 `json:"start"` + End float64 `json:"end"` + Text string `json:"text"` + } `json:"words"` +} + +// toProto drops the top-level words: core rebuilds them from segment words. +func (r transcriptionResult) toProto() *pb.TranscriptResult { + out := &pb.TranscriptResult{Text: r.Text, Language: r.Language, Duration: float32(r.Duration)} + for _, s := range r.Segments { + seg := &pb.TranscriptSegment{ + Id: s.ID, Start: nanos(s.Start), End: nanos(s.End), Text: s.Text, + Tokens: s.Tokens, Speaker: s.Speaker, + } + for _, w := range s.Words { + seg.Words = append(seg.Words, &pb.TranscriptWord{Start: nanos(w.Start), End: nanos(w.End), Text: w.Text}) + } + out.Segments = append(out.Segments, seg) + } + return out +} + +func nanos(seconds float64) int64 { + return int64(seconds * float64(time.Second)) +} + +func (p *LocalAIProxy) AudioTranscription(ctx context.Context, req *pb.TranscriptRequest) (pb.TranscriptResult, error) { + var resp transcriptionResult + if err := p.postForm(ctx, transcriptionsPath, p.transcriptionForm(req, false), &resp); err != nil { + return pb.TranscriptResult{}, err + } + return *resp.toProto(), nil +} + +// transcriptEvent covers every frame of the upstream transcription SSE stream. +type transcriptEvent struct { + Type string `json:"type"` + Delta string `json:"delta"` + Error *struct { + Message string `json:"message"` + } `json:"error"` + transcriptionResult +} + +// AudioTranscriptionStream maps upstream SSE frames to stream responses. out +// is closed on every path: the gRPC server drains it until closed. +func (p *LocalAIProxy) AudioTranscriptionStream(ctx context.Context, req *pb.TranscriptRequest, out chan *pb.TranscriptStreamResponse) error { + defer close(out) + resp, err := p.postMultipartStream(ctx, transcriptionsPath, p.transcriptionForm(req, true)) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + scanner := bufio.NewScanner(resp.Body) + // The done frame carries every segment of the recording. + scanner.Buffer(make([]byte, 0, 64*1024), 4<<20) + for scanner.Scan() { + payload, ok := strings.CutPrefix(scanner.Text(), "data:") + if !ok { + continue + } + payload = strings.TrimSpace(payload) + if payload == "" || payload == "[DONE]" { + continue + } + var ev transcriptEvent + if err := json.Unmarshal([]byte(payload), &ev); err != nil { + xlog.Debug("localai-proxy: skip malformed SSE frame", "path", transcriptionsPath, "error", err) + continue + } + switch ev.Type { + case "transcript.text.delta": + out <- &pb.TranscriptStreamResponse{Delta: ev.Delta} + case "transcript.text.done": + out <- &pb.TranscriptStreamResponse{FinalResult: ev.toProto()} + return nil + case "error": + msg := "unknown error" + if ev.Error != nil { + msg = ev.Error.Message + } + xlog.Warn("localai-proxy: upstream stream error", "path", transcriptionsPath, "error", msg) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", transcriptionsPath, msg) + } + } + if err := scanner.Err(); err != nil { + return transportError(transcriptionsPath, err) + } + // The upstream always ends with a done or error frame, so a stream that + // stops without one was cut off; returning nil would pass the partial + // deltas off as the whole transcript. + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream ended before the final transcript", transcriptionsPath) +} + +func (p *LocalAIProxy) Diarize(req *pb.DiarizeRequest) (pb.DiarizeResponse, error) { + f := url.Values{} + f.Set("model", p.model("")) + // verbose_json keeps per-segment text; core strips it again for callers + // that did not ask for it. + f.Set("response_format", "verbose_json") + if v := req.GetLanguage(); v != "" { + f.Set("language", v) + } + for name, v := range map[string]int32{ + "num_speakers": req.GetNumSpeakers(), "min_speakers": req.GetMinSpeakers(), "max_speakers": req.GetMaxSpeakers(), + } { + if v != 0 { + f.Set(name, strconv.Itoa(int(v))) + } + } + for name, v := range map[string]float32{ + "clustering_threshold": req.GetClusteringThreshold(), + "min_duration_on": req.GetMinDurationOn(), + "min_duration_off": req.GetMinDurationOff(), + } { + if v != 0 { + f.Set(name, formatFloat(v)) + } + } + if req.GetIncludeText() { + f.Set("include_text", "true") + } + + var resp struct { + Duration float64 `json:"duration"` + Language string `json:"language"` + NumSpeakers int32 `json:"num_speakers"` + Segments []struct { + ID int32 `json:"id"` + Speaker string `json:"speaker"` + Label string `json:"label"` + Start float32 `json:"start"` + End float32 `json:"end"` + Text string `json:"text"` + } `json:"segments"` + } + if err := p.postForm(context.Background(), "/v1/audio/diarization", + multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetDst()}}}, &resp); err != nil { + return pb.DiarizeResponse{}, err + } + var segments []*pb.DiarizeSegment + for _, s := range resp.Segments { + // The upstream renames speakers to SPEAKER_NN and keeps the backend's + // own id in label. Core renames again, so hand it the raw id: the + // caller then sees the same speakers and labels a direct call gives. + speaker := s.Label + if speaker == "" { + speaker = s.Speaker + } + segments = append(segments, &pb.DiarizeSegment{ + Id: s.ID, Start: s.Start, End: s.End, Speaker: speaker, Text: s.Text, + }) + } + return pb.DiarizeResponse{ + Segments: segments, NumSpeakers: resp.NumSpeakers, Duration: float32(resp.Duration), Language: resp.Language, + }, nil +} + +func (p *LocalAIProxy) VAD(req *pb.VADRequest) (pb.VADResponse, error) { + var resp struct { + Segments []struct { + Start float32 `json:"start"` + End float32 `json:"end"` + } `json:"segments"` + } + body := map[string]any{"model": p.model(""), "audio": req.GetAudio()} + if err := p.postJSON(context.Background(), "/v1/vad", body, &resp); err != nil { + return pb.VADResponse{}, err + } + var segments []*pb.VADSegment + for _, s := range resp.Segments { + segments = append(segments, &pb.VADSegment{Start: s.Start, End: s.End}) + } + return pb.VADResponse{Segments: segments}, nil +} + +func (p *LocalAIProxy) SoundDetection(ctx context.Context, req *pb.SoundDetectionRequest) (*pb.SoundDetectionResponse, error) { + f := url.Values{} + f.Set("model", p.model("")) + if v := req.GetTopK(); v != 0 { + f.Set("top_k", strconv.Itoa(int(v))) + } + if v := req.GetThreshold(); v != 0 { + f.Set("threshold", formatFloat(v)) + } + var resp struct { + Detections []struct { + Index int32 `json:"index"` + Label string `json:"label"` + Score float32 `json:"score"` + } `json:"detections"` + } + if err := p.postForm(ctx, "/v1/audio/classification", + multipartForm{fields: f, files: []formFile{{field: "file", path: req.GetSrc()}}}, &resp); err != nil { + return nil, err + } + out := &pb.SoundDetectionResponse{} + for _, d := range resp.Detections { + out.Detections = append(out.Detections, &pb.SoundClass{Index: d.Index, Label: d.Label, Score: d.Score}) + } + return out, nil +} + +// stemsHeader names the other outputs of a separation run: the body carries +// one file, and the rest are served under /generated-audio/. +const stemsHeader = "X-Audio-Stems" + +// AudioTransform uploads the input (and reference) to /audio/transformations +// and writes the returned audio to Dst. SampleRate and Samples stay 0: the +// REST reply does not report them, and core only uses them for tracing. +func (p *LocalAIProxy) AudioTransform(req *pb.AudioTransformRequest) (*pb.AudioTransformResult, error) { + f := url.Values{} + f.Set("model", p.model("")) + for k, v := range req.GetParams() { + f.Set("params["+k+"]", v) + } + form := multipartForm{fields: f, files: []formFile{{field: "audio", path: req.GetAudioPath()}}} + if ref := req.GetReferencePath(); ref != "" { + form.files = append(form.files, formFile{field: "reference", path: ref}) + } + + ctx := context.Background() + header, err := p.postMultipartToFile(ctx, "/audio/transformations", form, req.GetDst()) + if err != nil { + return nil, err + } + return &pb.AudioTransformResult{ + Dst: req.GetDst(), + ReferenceProvided: req.GetReferencePath() != "", + Stems: p.fetchStems(ctx, header.Get(stemsHeader), req.GetDst()), + }, nil +} + +// fetchStems downloads each stem the upstream names and writes it beside +// dst, where core looks for them. A stem that cannot be fetched is dropped +// with a warning rather than failing the call: the caller still gets the +// audio it asked for. +func (p *LocalAIProxy) fetchStems(ctx context.Context, header, dst string) []*pb.AudioTransformStem { + if header == "" { + return nil + } + var entries []struct { + Name string `json:"name"` + URL string `json:"url"` + } + if err := json.Unmarshal([]byte(header), &entries); err != nil { + xlog.Warn("localai-proxy: ignore malformed stems header", "error", err) + return nil + } + dir := filepath.Dir(dst) + prefix := strings.TrimSuffix(filepath.Base(dst), filepath.Ext(dst)) + var stems []*pb.AudioTransformStem + for _, e := range entries { + // Only files under /generated-audio/ are stems. Anything else is not + // a path this API hands out, so it is not fetched. + escaped, ok := strings.CutPrefix(e.URL, "/generated-audio/") + if !ok || e.Name == "" { + xlog.Warn("localai-proxy: skip stem with unexpected url", "name", e.Name, "url", e.URL) + continue + } + name, err := url.PathUnescape(escaped) + if err != nil || name == "" || name != filepath.Base(name) || name == "." || name == ".." { + xlog.Warn("localai-proxy: skip stem with unsafe file name", "name", e.Name, "url", e.URL) + continue + } + local := filepath.Join(dir, prefix+"-"+name) + if err := p.getToFile(ctx, e.URL, local); err != nil { + xlog.Warn("localai-proxy: stem download failed", "name", e.Name, "error", err) + continue + } + stems = append(stems, &pb.AudioTransformStem{Name: e.Name, Dst: local}) + } + return stems +} + +// formatFloat prints a float32 form value without float64 noise +// (0.2, not 0.20000000298023224). +func formatFloat(v float32) string { + return strconv.FormatFloat(float64(v), 'g', -1, 32) +} diff --git a/backend/go/localai-proxy/audio_test.go b/backend/go/localai-proxy/audio_test.go new file mode 100644 index 000000000..8f5013c7b --- /dev/null +++ b/backend/go/localai-proxy/audio_test.go @@ -0,0 +1,489 @@ +package main + +import ( + "context" + "fmt" + "net/http" + "os" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// writeInput writes content to a fresh temp file and returns its path, so a +// spec can check the upstream received exactly those bytes. +func writeInput(name, content string) string { + path := filepath.Join(GinkgoT().TempDir(), name) + Expect(os.WriteFile(path, []byte(content), 0o600)).To(Succeed()) + return path +} + +// drainBytes collects every chunk sent on ch until it is closed. +func drainBytes(ch chan []byte) <-chan [][]byte { + done := make(chan [][]byte, 1) + go func() { + var got [][]byte + for c := range ch { + got = append(got, c) + } + done <- got + }() + return done +} + +func drainTranscript(ch chan *pb.TranscriptStreamResponse) <-chan []*pb.TranscriptStreamResponse { + done := make(chan []*pb.TranscriptStreamResponse, 1) + go func() { + var got []*pb.TranscriptStreamResponse + for c := range ch { + got = append(got, c) + } + done <- got + }() + return done +} + +// cutMidStream answers 200, flushes body, then drops the connection without +// finishing the chunked encoding, as a crashed upstream would. +func cutMidStream(contentType, body string) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(body)) + w.(http.Flusher).Flush() + conn, _, err := w.(http.Hijacker).Hijack() + if err == nil { + _ = conn.Close() + } + } +} + +var _ = Describe("audio methods", func() { + var up *fakeUpstream + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.Close) + }) + + Describe("TTS", func() { + It("posts to /tts and writes the upstream audio to Dst", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-audio-bytes"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.wav") + lang := "it" + instr := "cheerful" + + Expect(p.TTS(&pb.TTSRequest{ + Text: "ciao", Model: "local-path", Dst: dst, Voice: "v1", Language: &lang, + Instructions: &instr, Params: map[string]string{"speed": "1.2"}, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-audio-bytes")) + req := up.last() + Expect(req.Path).To(Equal("/tts")) + Expect(req.JSON).To(Equal(map[string]any{ + "model": "remote-model", "input": "ciao", "voice": "v1", "language": "it", + "instructions": "cheerful", "params": map[string]any{"speed": "1.2"}, + })) + }) + + It("maps an upstream failure and leaves no partial file", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"}) + dst := filepath.Join(GinkgoT().TempDir(), "out.wav") + err := p.TTS(&pb.TTSRequest{Text: "x", Dst: dst}) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(dst).NotTo(BeAnExistingFile()) + }) + }) + + Describe("TTSStream", func() { + It("forwards the chunked WAV body in order and closes the channel", func() { + header := "RIFF" + string(make([]byte, 40)) + up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "audio/wav") + w.WriteHeader(http.StatusOK) + for _, c := range []string{header, "pcm-1", "pcm-2"} { + _, _ = w.Write([]byte(c)) + w.(http.Flusher).Flush() + time.Sleep(10 * time.Millisecond) + } + }) + DeferCleanup(up.Close) + p := loadProxy(up, nil) + + out := make(chan []byte) + done := drainBytes(out) + Expect(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out)).To(Succeed()) + + var chunks [][]byte + Eventually(done).Should(Receive(&chunks)) + Expect(chunks).NotTo(BeEmpty()) + Expect(string(chunks[0][:4])).To(Equal("RIFF")) + var all []byte + for _, c := range chunks { + all = append(all, c...) + } + Expect(string(all)).To(Equal(header + "pcm-1pcm-2")) + }) + + It("asks the upstream to stream", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF"}) + out := make(chan []byte) + done := drainBytes(out) + Expect(p.TTSStream(&pb.TTSRequest{Text: "hi", Voice: "v"}, out)).To(Succeed()) + Eventually(done).Should(Receive()) + Expect(up.last().JSON).To(HaveKeyWithValue("stream", true)) + Expect(up.last().JSON).To(HaveKeyWithValue("input", "hi")) + }) + + It("reports a mid-stream disconnect as Unavailable and still closes the channel", func() { + cut := newFakeUpstreamWithHandler(cutMidStream("audio/wav", "RIFFpartial")) + DeferCleanup(cut.Close) + p := loadProxy(cut, nil) + out := make(chan []byte) + done := drainBytes(out) + err := p.TTSStream(&pb.TTSRequest{Text: "hi"}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Eventually(done).Should(Receive()) + }) + + It("closes the channel when the upstream refuses the request", func() { + p := loadProxy(up, nil) + up.script("/tts", scriptedResponse{Status: http.StatusBadRequest, Body: "bad"}) + out := make(chan []byte) + done := drainBytes(out) + Expect(codeOf(p.TTSStream(&pb.TTSRequest{Text: "hi"}, out))).To(Equal(codes.InvalidArgument)) + Eventually(done).Should(Receive()) + }) + }) + + Describe("SoundGeneration", func() { + It("posts the ElevenLabs body to /v1/sound-generation and writes Dst", func() { + p := loadProxy(up, nil) + up.script("/v1/sound-generation", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-sfx"}) + dst := filepath.Join(GinkgoT().TempDir(), "sfx.wav") + dur, temp, bpm := float32(4.5), float32(0.3), int32(120) + sample, instrumental := true, false + + Expect(p.SoundGeneration(&pb.SoundGenerationRequest{ + Text: "rain", Dst: dst, Duration: &dur, Temperature: &temp, Sample: &sample, + Bpm: &bpm, Caption: ptr("cap"), Lyrics: ptr("la"), Keyscale: ptr("C major"), + Language: ptr("en"), Timesignature: ptr("4/4"), Instrumental: &instrumental, + })).To(Succeed()) + + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-sfx")) + Expect(up.last().JSON).To(Equal(map[string]any{ + "model_id": "remote-model", "text": "rain", "duration_seconds": 4.5, + "prompt_influence": 0.3, "do_sample": true, "bpm": float64(120), + "caption": "cap", "lyrics": "la", "keyscale": "C major", "language": "en", + "timesignature": "4/4", "instrumental": false, + })) + }) + + It("refuses audio-conditioned generation it cannot upload", func() { + p := loadProxy(up, nil) + err := p.SoundGeneration(&pb.SoundGenerationRequest{Text: "x", Src: ptr("/tmp/in.wav")}) + Expect(codeOf(err)).To(Equal(codes.Unimplemented)) + Expect(up.recorded()).To(BeEmpty()) + }) + }) + + Describe("AudioTranscription", func() { + It("uploads Dst with the request fields and maps the seconds-based result", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/transcriptions", map[string]any{ + "text": "hello world", "language": "en", "duration": 2.5, + "segments": []any{map[string]any{ + "id": 0, "start": 0.5, "end": 1.25, "text": "hello world", "tokens": []int{1, 2}, + "speaker": "SPEAKER_00", + "words": []any{map[string]any{"start": 0.5, "end": 0.75, "text": "hello"}}, + }}, + }) + in := writeInput("speech.wav", "RIFF-speech") + + res, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{ + Dst: in, Language: "en", Translate: true, Diarize: true, Prompt: "names: Ada", + Temperature: 0.2, TimestampGranularities: []string{"word"}, + }) + Expect(err).NotTo(HaveOccurred()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/audio/transcriptions")) + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-speech"})) + Expect(req.Fields).To(Equal(map[string]string{ + "model": "remote-model", "language": "en", "translate": "true", "diarize": "true", + "prompt": "names: Ada", "temperature": "0.2", "timestamp_granularities[]": "word", + "response_format": "verbose_json", + })) + + Expect(res.Text).To(Equal("hello world")) + Expect(res.Language).To(Equal("en")) + Expect(res.Duration).To(BeNumerically("==", 2.5)) + Expect(res.Segments).To(HaveLen(1)) + seg := res.Segments[0] + // pb carries nanoseconds (core reads them as time.Duration). + Expect(seg.Start).To(Equal(int64(500 * time.Millisecond))) + Expect(seg.End).To(Equal(int64(1250 * time.Millisecond))) + Expect(seg.Text).To(Equal("hello world")) + Expect(seg.Tokens).To(Equal([]int32{1, 2})) + Expect(seg.Speaker).To(Equal("SPEAKER_00")) + Expect(seg.Words).To(HaveLen(1)) + Expect(seg.Words[0].Start).To(Equal(int64(500 * time.Millisecond))) + Expect(seg.Words[0].Text).To(Equal("hello")) + }) + + It("sends diarize=false explicitly because the upstream defaults it on", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "x"}) + _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}) + Expect(err).NotTo(HaveOccurred()) + f := up.last().Fields + Expect(f).To(HaveKeyWithValue("diarize", "false")) + // Unset language and translate are left to the upstream model config. + Expect(f).NotTo(HaveKey("language")) + Expect(f).NotTo(HaveKey("translate")) + }) + + It("maps a 4xx to InvalidArgument", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad audio"}) + _, err := p.AudioTranscription(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + }) + }) + + Describe("AudioTranscriptionStream", func() { + It("emits deltas, then the final result, and closes the channel", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}), + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "lo"}), + sseJSON(map[string]any{ + "type": "transcript.text.done", "text": "hello", "language": "en", "duration": 1.5, + "segments": []any{map[string]any{"id": 0, "start": 0.0, "end": 1.5, "text": "hello"}}, + }), + "[DONE]", + }}) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + + Expect(p.AudioTranscriptionStream(context.Background(), + &pb.TranscriptRequest{Dst: writeInput("a.wav", "RIFF-a"), Stream: true}, out)).To(Succeed()) + + var got []*pb.TranscriptStreamResponse + Eventually(done).Should(Receive(&got)) + Expect(got).To(HaveLen(3)) + Expect(got[0].Delta).To(Equal("hel")) + Expect(got[1].Delta).To(Equal("lo")) + final := got[2].FinalResult + Expect(final).NotTo(BeNil()) + Expect(final.Text).To(Equal("hello")) + Expect(final.Language).To(Equal("en")) + Expect(final.Duration).To(BeNumerically("==", 1.5)) + Expect(final.Segments).To(HaveLen(1)) + Expect(final.Segments[0].End).To(Equal(int64(1500 * time.Millisecond))) + + req := up.last() + Expect(req.Fields).To(HaveKeyWithValue("stream", "true")) + Expect(req.Files).To(HaveKeyWithValue("file", "RIFF-a")) + }) + + It("returns an upstream error event as Unavailable", func() { + p := loadProxy(up, nil) + up.script("/v1/audio/transcriptions", scriptedResponse{SSE: []string{ + sseJSON(map[string]any{"type": "error", "error": map[string]any{"message": "decoder died"}}), + "[DONE]", + }}) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + Expect(err).To(MatchError(ContainSubstring("decoder died"))) + Eventually(done).Should(Receive()) + }) + + It("reports an upstream disconnect before the final result as Unavailable", func() { + frame := "data: " + sseJSON(map[string]any{"type": "transcript.text.delta", "delta": "hel"}) + "\n\n" + cut := newFakeUpstreamWithHandler(cutMidStream("text/event-stream", frame)) + DeferCleanup(cut.Close) + p := loadProxy(cut, nil) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: writeInput("a.wav", "A")}, out) + Expect(codeOf(err)).To(Equal(codes.Unavailable)) + var got []*pb.TranscriptStreamResponse + Eventually(done).Should(Receive(&got)) + Expect(got).To(HaveLen(1)) + }) + + It("closes the channel when the local file is missing", func() { + p := loadProxy(up, nil) + out := make(chan *pb.TranscriptStreamResponse) + done := drainTranscript(out) + err := p.AudioTranscriptionStream(context.Background(), &pb.TranscriptRequest{Dst: "/nonexistent.wav"}, out) + Expect(codeOf(err)).To(Equal(codes.InvalidArgument)) + Eventually(done).Should(Receive()) + Expect(up.recorded()).To(BeEmpty()) + }) + }) + + Describe("Diarize", func() { + It("uploads Dst with the tuning fields and maps the segments", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/diarization", map[string]any{ + "task": "diarize", "duration": 3.0, "language": "en", "num_speakers": 2, + "segments": []any{ + map[string]any{"id": 0, "speaker": "SPEAKER_00", "label": "spk_a", "start": 0.0, "end": 1.5, "text": "hi"}, + map[string]any{"id": 1, "speaker": "SPEAKER_01", "start": 1.5, "end": 3.0}, + }, + }) + res, err := p.Diarize(&pb.DiarizeRequest{ + Dst: writeInput("talk.wav", "RIFF-talk"), Language: "en", NumSpeakers: 2, MinSpeakers: 1, + MaxSpeakers: 3, ClusteringThreshold: 0.5, MinDurationOn: 0.1, MinDurationOff: 0.2, IncludeText: true, + }) + Expect(err).NotTo(HaveOccurred()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/audio/diarization")) + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-talk"})) + Expect(req.Fields).To(Equal(map[string]string{ + "model": "remote-model", "language": "en", "num_speakers": "2", "min_speakers": "1", + "max_speakers": "3", "clustering_threshold": "0.5", "min_duration_on": "0.1", + "min_duration_off": "0.2", "include_text": "true", "response_format": "verbose_json", + })) + + Expect(res.NumSpeakers).To(Equal(int32(2))) + Expect(res.Duration).To(BeNumerically("==", 3.0)) + Expect(res.Language).To(Equal("en")) + Expect(res.Segments).To(HaveLen(2)) + // The raw upstream label survives, so core's own normalisation + // yields the same speakers a direct call would. + Expect(res.Segments[0].Speaker).To(Equal("spk_a")) + Expect(res.Segments[0].Text).To(Equal("hi")) + Expect(res.Segments[1].Speaker).To(Equal("SPEAKER_01")) + Expect(res.Segments[1].Start).To(BeNumerically("==", 1.5)) + Expect(res.Segments[1].Id).To(Equal(int32(1))) + }) + }) + + Describe("VAD", func() { + It("posts the samples to /v1/vad and maps the segments", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/vad", map[string]any{"segments": []any{map[string]any{"start": 0.25, "end": 1.5}}}) + res, err := p.VAD(&pb.VADRequest{Audio: []float32{0.5, -0.25}}) + Expect(err).NotTo(HaveOccurred()) + Expect(up.last().JSON).To(Equal(map[string]any{"model": "remote-model", "audio": []any{0.5, -0.25}})) + Expect(res.Segments).To(HaveLen(1)) + Expect(res.Segments[0].Start).To(BeNumerically("==", 0.25)) + Expect(res.Segments[0].End).To(BeNumerically("==", 1.5)) + }) + }) + + Describe("SoundDetection", func() { + It("uploads Src to /v1/audio/classification and maps the detections", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/audio/classification", map[string]any{"model": "remote-model", "detections": []any{ + map[string]any{"index": 74, "label": "Dog", "score": 0.9}, + map[string]any{"index": 0, "label": "Speech", "score": 0.4}, + }}) + res, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{ + Src: writeInput("bark.wav", "RIFF-bark"), TopK: 5, Threshold: 0.25, + }) + Expect(err).NotTo(HaveOccurred()) + req := up.last() + Expect(req.Files).To(Equal(map[string]string{"file": "RIFF-bark"})) + Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "top_k": "5", "threshold": "0.25"})) + Expect(res.Detections).To(HaveLen(2)) + Expect(res.Detections[0].Label).To(Equal("Dog")) + Expect(res.Detections[0].Index).To(Equal(int32(74))) + Expect(res.Detections[0].Score).To(BeNumerically("~", 0.9, 1e-6)) + }) + }) + + Describe("AudioTransform", func() { + It("uploads audio and reference with params and writes the result to Dst", func() { + p := loadProxy(up, nil) + up.script("/audio/transformations", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-clean"}) + dst := filepath.Join(GinkgoT().TempDir(), "transform.wav") + + res, err := p.AudioTransform(&pb.AudioTransformRequest{ + AudioPath: writeInput("mic.wav", "RIFF-mic"), ReferencePath: writeInput("ref.wav", "RIFF-ref"), + Dst: dst, Params: map[string]string{"noise_gate": "true"}, + }) + Expect(err).NotTo(HaveOccurred()) + got, err := os.ReadFile(dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-clean")) + Expect(res.Dst).To(Equal(dst)) + Expect(res.ReferenceProvided).To(BeTrue()) + Expect(res.Stems).To(BeEmpty()) + + req := up.last() + Expect(req.Path).To(Equal("/audio/transformations")) + Expect(req.Files).To(Equal(map[string]string{"audio": "RIFF-mic", "reference": "RIFF-ref"})) + Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "params[noise_gate]": "true"})) + }) + + It("fetches the stems the upstream names and writes them beside Dst", func() { + p := loadProxy(up, nil) + stems := `[{"name":"vocals","url":"/generated-audio/sep%20vocals.wav"},{"name":"evil","url":"/etc/passwd"}]` + up.script("/audio/transformations", scriptedResponse{ + Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-drums", + Header: map[string]string{"X-Audio-Stems": stems}, + }) + up.script("/generated-audio/sep vocals.wav", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "RIFF-vocals"}) + dir := GinkgoT().TempDir() + dst := filepath.Join(dir, "transform.wav") + + res, err := p.AudioTransform(&pb.AudioTransformRequest{AudioPath: writeInput("song.wav", "RIFF-song"), Dst: dst}) + Expect(err).NotTo(HaveOccurred()) + Expect(res.ReferenceProvided).To(BeFalse()) + Expect(res.Stems).To(HaveLen(1)) + Expect(res.Stems[0].Name).To(Equal("vocals")) + // Core keeps only stems that are direct children of Dst's directory. + Expect(filepath.Dir(res.Stems[0].Dst)).To(Equal(dir)) + got, err := os.ReadFile(res.Stems[0].Dst) + Expect(err).NotTo(HaveOccurred()) + Expect(string(got)).To(Equal("RIFF-vocals")) + + var paths []string + for _, r := range up.recorded() { + paths = append(paths, r.Method+" "+r.Path) + } + Expect(paths).To(Equal([]string{"POST /audio/transformations", "GET /generated-audio/sep vocals.wav"})) + }) + }) + + It("names the upload file after the local file", func() { + // The upstream saves the upload under its base name and some handlers + // pick the decoder by extension, so the name must reach it intact. + var name string + named := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { + _, fh, err := r.FormFile("file") + if err == nil { + name = fh.Filename + } + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprint(w, `{"detections":[]}`) + }) + DeferCleanup(named.Close) + p := loadProxy(named, nil) + _, err := p.SoundDetection(context.Background(), &pb.SoundDetectionRequest{Src: writeInput("clip.mp3", "ID3")}) + Expect(err).NotTo(HaveOccurred()) + Expect(name).To(Equal("clip.mp3")) + }) +}) + +func ptr[T any](v T) *T { return &v } diff --git a/backend/go/localai-proxy/client.go b/backend/go/localai-proxy/client.go index 100a931dd..166fa1ab8 100644 --- a/backend/go/localai-proxy/client.go +++ b/backend/go/localai-proxy/client.go @@ -9,6 +9,7 @@ import ( "io" "mime/multipart" "net/http" + "net/url" "os" "path/filepath" "strings" @@ -37,7 +38,7 @@ func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any) ctx, cancel := withTimeout(ctx, cfg) defer cancel() - req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload)) + req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload)) if err != nil { return err } @@ -51,48 +52,92 @@ func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any) // 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 } - var file *os.File - if fileField != "" { - // Open before contacting the upstream so a bad path is reported as a - // request error, not as a failure of the remote host. - if file, err = os.Open(filePath); err != nil { - return status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", filePath, err) - } - defer func() { _ = file.Close() }() - } ctx, cancel := withTimeout(ctx, cfg) defer cancel() - - // Stream the form through a pipe so large audio files are not buffered - // in memory. The transport closes pr when the request ends, which - // unblocks the writer on every error path. - pr, pw := io.Pipe() - mw := multipart.NewWriter(pw) - go func() { - pw.CloseWithError(writeMultipart(mw, fields, fileField, file)) - }() - - req, err := p.newRequest(ctx, cfg, path, pr) + req, err := p.newMultipartRequest(ctx, cfg, path, form) if err != nil { - _ = pr.CloseWithError(err) return err } - req.Header.Set("Content-Type", mw.FormDataContentType()) return p.do(req, path, out) } -func writeMultipart(mw *multipart.Writer, fields map[string]string, fileField string, file *os.File) error { - for k, v := range fields { - if err := mw.WriteField(k, v); err != nil { - return err +// 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() } } - if file != nil { - part, err := mw.CreateFormFile(fileField, filepath.Base(file.Name())) + 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 } @@ -115,11 +160,30 @@ func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (* if err != nil { return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err) } - req, err := p.newRequest(ctx, cfg, path, bytes.NewReader(payload)) + 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) @@ -131,8 +195,8 @@ func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (* return resp, nil } -func (p *LocalAIProxy) newRequest(ctx context.Context, cfg *proxyConfig, path string, body io.Reader) (*http.Request, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.base+path, body) +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) } @@ -224,3 +288,88 @@ func statusError(path string, resp *http.Response) error { 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() }() + + f, err := os.Create(dst) + if err != nil { + return nil, status.Errorf(codes.Internal, "localai-proxy: create %s: %v", dst, err) + } + _, copyErr := io.Copy(f, resp.Body) + closeErr := f.Close() + if copyErr != nil || closeErr != nil { + _ = os.Remove(dst) + if copyErr != nil { + if ctxErr := req.Context().Err(); ctxErr != nil { + return nil, transportError(path, ctxErr) + } + return nil, transportError(path, copyErr) + } + return nil, status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, closeErr) + } + return resp.Header, nil +} diff --git a/backend/go/localai-proxy/fake_upstream_test.go b/backend/go/localai-proxy/fake_upstream_test.go index 6ced223a8..933e012be 100644 --- a/backend/go/localai-proxy/fake_upstream_test.go +++ b/backend/go/localai-proxy/fake_upstream_test.go @@ -27,10 +27,11 @@ type recordedRequest struct { } // scriptedResponse is the reply for one path. SSE, when set, is written as -// "data: " events and wins over Body. +// "data: " events and wins over Body. Header adds response headers. type scriptedResponse struct { Status int ContentType string + Header map[string]string Body string SSE []string } @@ -133,6 +134,9 @@ func (f *fakeUpstream) serve(w http.ResponseWriter, r *http.Request) { 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) } From be9c039e65727816961245d7688a5e83309f8474 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 23:48:49 +0000 Subject: [PATCH 46/79] feat(localai-proxy): serve image, video, 3D and vision APIs remotely Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/media.go | 779 +++++++++++++++++++++++++ backend/go/localai-proxy/media_test.go | 434 ++++++++++++++ 2 files changed, 1213 insertions(+) create mode 100644 backend/go/localai-proxy/media.go create mode 100644 backend/go/localai-proxy/media_test.go diff --git a/backend/go/localai-proxy/media.go b/backend/go/localai-proxy/media.go new file mode 100644 index 000000000..e2094a092 --- /dev/null +++ b/backend/go/localai-proxy/media.go @@ -0,0 +1,779 @@ +package main + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "net/url" + "os" + "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 + } + data, err := os.ReadFile(path) + if err != nil { + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: read %s: %v", path, err) + } + return base64.StdEncoding.EncodeToString(data), 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) + } + if err := os.WriteFile(dst, data, 0o644); 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) + } + return p.getToFile(ctx, p.relativePath(item.URL), dst) +} + +// relativePath strips the configured upstream base from a URL the upstream +// handed back (e.g. "/generated-images/x.png"), so the follow-up +// download goes through the normal request path instead of concatenating two +// absolute URLs. +func (p *LocalAIProxy) relativePath(raw string) string { + if cfg := p.cfg.Load(); cfg != nil { + if rel, ok := strings.CutPrefix(raw, cfg.base); ok { + return rel + } + } + return 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 base64 payload (the same +// convention DetectionEndpoint uses to call this method locally), and maps +// each detection, decoding its PNG mask. +func (p *LocalAIProxy) Detect(req *pb.DetectOptions) (pb.DetectResponse, error) { + body := detectRequestBody{ + Model: p.model(""), + Image: req.GetSrc(), + 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 +} + +type depthRequestBody struct { + Model string `json:"model"` + Image string `json:"image"` + Dst string `json:"dst,omitempty"` + 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"` + Exports []string `json:"exports,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 base64 payload, per the same convention as Detect) and +// maps the full response, decoding the point-cloud color bytes. +func (p *LocalAIProxy) Depth(req *pb.DepthRequest) (pb.DepthResponse, error) { + body := depthRequestBody{ + Model: p.model(""), + Image: req.GetSrc(), + Dst: req.GetDst(), + IncludeDepth: req.GetIncludeDepth(), + IncludeConfidence: req.GetIncludeConfidence(), + IncludePose: req.GetIncludePose(), + IncludeSky: req.GetIncludeSky(), + IncludePoints: req.GetIncludePoints(), + PointsConfThresh: req.GetPointsConfThresh(), + Exports: req.GetExports(), + } + 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 base64 (the same +// convention FaceVerifyEndpoint uses to call this method locally). +func (p *LocalAIProxy) FaceVerify(req *pb.FaceVerifyRequest) (pb.FaceVerifyResponse, error) { + body := faceVerifyRequestBody{ + Model: p.model(""), Img1: req.GetImg1(), Img2: req.GetImg2(), + 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 base64. +func (p *LocalAIProxy) FaceAnalyze(req *pb.FaceAnalyzeRequest) (pb.FaceAnalyzeResponse, error) { + body := faceAnalyzeRequestBody{ + Model: p.model(""), Img: req.GetImg(), Actions: req.GetActions(), AntiSpoofing: req.GetAntiSpoofing(), + } + var resp faceAnalyzeResponseBody + if err := p.postJSON(context.Background(), "/v1/face/analyze", body, &resp); err != nil { + return pb.FaceAnalyzeResponse{}, err + } + var faces []*pb.FaceAnalysis + for _, f := range resp.Faces { + faces = append(faces, &pb.FaceAnalysis{ + Region: f.Region.toProto(), FaceConfidence: f.FaceConfidence, Age: f.Age, + DominantGender: f.DominantGender, Gender: f.Gender, + DominantEmotion: f.DominantEmotion, Emotion: f.Emotion, + DominantRace: f.DominantRace, Race: f.Race, + IsReal: boolValue(f.IsReal), AntispoofScore: float32Value(f.AntispoofScore), + }) + } + // A composite literal here, not a variable of type pb.FaceAnalyzeResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.FaceAnalyzeResponse{Faces: faces}, nil +} + +// --- Voice (speaker) recognition ----------------------------------------- + +type voiceVerifyRequestBody struct { + Model string `json:"model"` + Audio1 string `json:"audio1"` + Audio2 string `json:"audio2"` + Threshold float32 `json:"threshold,omitempty"` + AntiSpoofing bool `json:"anti_spoofing,omitempty"` +} + +type voiceVerifyResponseBody struct { + Verified bool `json:"verified"` + Distance float32 `json:"distance"` + Threshold float32 `json:"threshold"` + Confidence float32 `json:"confidence"` + Model string `json:"model"` + ProcessingTimeMs float32 `json:"processing_time_ms,omitempty"` +} + +// VoiceVerify sends Audio1/Audio2 (staged local paths) as base64, matching +// VoiceVerifyRequest's URL/base64/data-URI contract. +func (p *LocalAIProxy) VoiceVerify(req *pb.VoiceVerifyRequest) (pb.VoiceVerifyResponse, error) { + audio1, err := fileToBase64(req.GetAudio1()) + if err != nil { + return pb.VoiceVerifyResponse{}, err + } + audio2, err := fileToBase64(req.GetAudio2()) + if err != nil { + return pb.VoiceVerifyResponse{}, err + } + body := voiceVerifyRequestBody{ + Model: p.model(""), Audio1: audio1, Audio2: audio2, + Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(), + } + var resp voiceVerifyResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/verify", body, &resp); err != nil { + return pb.VoiceVerifyResponse{}, err + } + return pb.VoiceVerifyResponse{ + Verified: resp.Verified, Distance: resp.Distance, Threshold: resp.Threshold, + Confidence: resp.Confidence, Model: resp.Model, ProcessingTimeMs: resp.ProcessingTimeMs, + }, nil +} + +type voiceAnalyzeRequestBody struct { + Model string `json:"model"` + Audio string `json:"audio"` + Actions []string `json:"actions,omitempty"` +} + +type voiceAnalysisBody struct { + Start float32 `json:"start"` + End float32 `json:"end"` + Age float32 `json:"age,omitempty"` + DominantGender string `json:"dominant_gender,omitempty"` + Gender map[string]float32 `json:"gender,omitempty"` + DominantEmotion string `json:"dominant_emotion,omitempty"` + Emotion map[string]float32 `json:"emotion,omitempty"` +} + +type voiceAnalyzeResponseBody struct { + Segments []voiceAnalysisBody `json:"segments"` +} + +// VoiceAnalyze sends Audio (a staged local path) as base64. +func (p *LocalAIProxy) VoiceAnalyze(req *pb.VoiceAnalyzeRequest) (pb.VoiceAnalyzeResponse, error) { + audio, err := fileToBase64(req.GetAudio()) + if err != nil { + return pb.VoiceAnalyzeResponse{}, err + } + body := voiceAnalyzeRequestBody{Model: p.model(""), Audio: audio, Actions: req.GetActions()} + var resp voiceAnalyzeResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/analyze", body, &resp); err != nil { + return pb.VoiceAnalyzeResponse{}, err + } + var segments []*pb.VoiceAnalysis + for _, s := range resp.Segments { + segments = append(segments, &pb.VoiceAnalysis{ + Start: s.Start, End: s.End, Age: s.Age, + DominantGender: s.DominantGender, Gender: s.Gender, + DominantEmotion: s.DominantEmotion, Emotion: s.Emotion, + }) + } + // A composite literal here, not a variable of type pb.VoiceAnalyzeResponse, + // avoids copying the protobuf message's embedded lock on return. + return pb.VoiceAnalyzeResponse{Segments: segments}, nil +} + +type voiceEmbedRequestBody struct { + Model string `json:"model"` + Audio string `json:"audio"` +} + +type voiceEmbedResponseBody struct { + Embedding []float32 `json:"embedding"` + Model string `json:"model,omitempty"` +} + +// VoiceEmbed sends Audio (a staged local path) as base64. +func (p *LocalAIProxy) VoiceEmbed(req *pb.VoiceEmbedRequest) (pb.VoiceEmbedResponse, error) { + audio, err := fileToBase64(req.GetAudio()) + if err != nil { + return pb.VoiceEmbedResponse{}, err + } + body := voiceEmbedRequestBody{Model: p.model(""), Audio: audio} + var resp voiceEmbedResponseBody + if err := p.postJSON(context.Background(), "/v1/voice/embed", body, &resp); err != nil { + return pb.VoiceEmbedResponse{}, err + } + return pb.VoiceEmbedResponse{Embedding: resp.Embedding, Model: resp.Model}, nil +} + +// --- Stores --------------------------------------------------------------- + +// storeKeysToFloats and storeFloatsToKeys convert between the proto's boxed +// StoresKey/StoresValue slices and the plain [][]float32 / []string the REST +// stores endpoints take, per schema.StoresSet and friends. + +func storeKeysToFloats(keys []*pb.StoresKey) [][]float32 { + out := make([][]float32, len(keys)) + for i, k := range keys { + out[i] = k.GetFloats() + } + return out +} + +func storeFloatsToKeys(keys [][]float32) []*pb.StoresKey { + out := make([]*pb.StoresKey, len(keys)) + for i, k := range keys { + out[i] = &pb.StoresKey{Floats: k} + } + return out +} + +func storeValuesToStrings(values []*pb.StoresValue) []string { + out := make([]string, len(values)) + for i, v := range values { + out[i] = string(v.GetBytes()) + } + return out +} + +func storeStringsToValues(values []string) []*pb.StoresValue { + out := make([]*pb.StoresValue, len(values)) + for i, v := range values { + out[i] = &pb.StoresValue{Bytes: []byte(v)} + } + return out +} + +type storesSetRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` +} + +// StoresSet uses the configured upstream model name as the store name: the +// proxy is loaded per store, the same way it is loaded per model for every +// other method. +func (p *LocalAIProxy) StoresSet(req *pb.StoresSetOptions) error { + body := storesSetRequestBody{ + Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys()), Values: storeValuesToStrings(req.GetValues()), + } + return p.postJSON(context.Background(), "/stores/set", body, nil) +} + +type storesDeleteRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` +} + +func (p *LocalAIProxy) StoresDelete(req *pb.StoresDeleteOptions) error { + body := storesDeleteRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())} + return p.postJSON(context.Background(), "/stores/delete", body, nil) +} + +type storesGetRequestBody struct { + Store string `json:"store,omitempty"` + Keys [][]float32 `json:"keys"` +} + +type storesGetResponseBody struct { + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` +} + +func (p *LocalAIProxy) StoresGet(req *pb.StoresGetOptions) (pb.StoresGetResult, error) { + body := storesGetRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())} + var resp storesGetResponseBody + if err := p.postJSON(context.Background(), "/stores/get", body, &resp); err != nil { + return pb.StoresGetResult{}, err + } + return pb.StoresGetResult{Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values)}, nil +} + +type storesFindRequestBody struct { + Store string `json:"store,omitempty"` + Key []float32 `json:"key"` + Topk int `json:"topk,omitempty"` +} + +type storesFindResponseBody struct { + Keys [][]float32 `json:"keys"` + Values []string `json:"values"` + Similarities []float32 `json:"similarities"` +} + +func (p *LocalAIProxy) StoresFind(req *pb.StoresFindOptions) (pb.StoresFindResult, error) { + body := storesFindRequestBody{Store: p.model(""), Key: req.GetKey().GetFloats(), Topk: int(req.GetTopK())} + var resp storesFindResponseBody + if err := p.postJSON(context.Background(), "/stores/find", body, &resp); err != nil { + return pb.StoresFindResult{}, err + } + return pb.StoresFindResult{ + Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values), Similarities: resp.Similarities, + }, nil +} diff --git a/backend/go/localai-proxy/media_test.go b/backend/go/localai-proxy/media_test.go new file mode 100644 index 000000000..2e2a970e9 --- /dev/null +++ b/backend/go/localai-proxy/media_test.go @@ -0,0 +1,434 @@ +package main + +import ( + "encoding/base64" + "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" +) + +// 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()) + }) + }) + + 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")) + }) + }) + + Describe("Detect", func() { + It("posts the image and maps detections including the mask", func() { + p := loadProxy(up, nil) + mask := base64.StdEncoding.EncodeToString([]byte("png-mask")) + up.replyJSON("/v1/detection", 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}, + }, + }) + + res, err := p.Detect(&pb.DetectOptions{Src: "base64-image-data", 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"))) + + req := up.last() + Expect(req.Path).To(Equal("/v1/detection")) + Expect(req.JSON).To(HaveKeyWithValue("image", "base64-image-data")) + Expect(req.JSON).To(HaveKeyWithValue("prompt", "cat")) + Expect(req.JSON).To(HaveKeyWithValue("threshold", BeNumerically("~", 0.5, 1e-6))) + }) + }) + + Describe("Depth", func() { + It("posts the image and maps the full depth response", func() { + p := loadProxy(up, nil) + colors := base64.StdEncoding.EncodeToString([]byte("rgb")) + up.replyJSON("/v1/depth", map[string]any{ + "width": 2, "height": 1, "depth": []float64{0.1, 0.2}, + "point_colors": colors, "is_metric": true, + }) + + res, err := p.Depth(&pb.DepthRequest{Src: "base64-image-data", 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()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/depth")) + Expect(req.JSON).To(HaveKeyWithValue("image", "base64-image-data")) + Expect(req.JSON).To(HaveKeyWithValue("include_depth", true)) + }) + }) + + Describe("FaceVerify", func() { + It("posts both images and maps the response, including liveness fields", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/face/verify", 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, + }) + + res, err := p.FaceVerify(&pb.FaceVerifyRequest{Img1: "img1-b64", Img2: "img2-b64", 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()) + + req := up.last() + Expect(req.Path).To(Equal("/v1/face/verify")) + Expect(req.JSON).To(HaveKeyWithValue("img1", "img1-b64")) + Expect(req.JSON).To(HaveKeyWithValue("img2", "img2-b64")) + Expect(req.JSON).To(HaveKeyWithValue("anti_spoofing", true)) + }) + }) + + Describe("FaceAnalyze", func() { + It("posts the image and maps per-face demographic attributes", func() { + p := loadProxy(up, nil) + up.replyJSON("/v1/face/analyze", 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}, + }, + }, + }) + + res, err := p.FaceAnalyze(&pb.FaceAnalyzeRequest{Img: "img-b64", 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)))) + + req := up.last() + Expect(req.Path).To(Equal("/v1/face/analyze")) + Expect(req.JSON).To(HaveKeyWithValue("img", "img-b64")) + Expect(req.JSON).To(HaveKeyWithValue("actions", ConsistOf("age", "gender"))) + }) + }) + + 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)))) + }) + }) +}) From 37023939ad1a97f3a8755dbc7a49eb17378cd1c8 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 00:04:31 +0000 Subject: [PATCH 47/79] fix(localai-proxy): wrap bare base64 image inputs, block unreachable Depth exports, and validate download URLs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Detect/Depth/FaceVerify/FaceAnalyze forwarded the bare base64 core hands the backend, but the upstream's own REST handlers only accept a URL or a data:...;base64, string, so every real call 400'd. Wrap the payload as a data URI (sniffing its MIME type) before sending it. Depth requests for exports/dst now return Unimplemented: those files are written to the upstream's own local disk and are unreachable from here, so failover should move to a local target instead. Generation replies that hand back a URL are now re-fetched by path only, checked against the upstream's known generated-content prefixes, instead of stripping the configured base as a literal string prefix — the old approach broke (or silently trusted an arbitrary host) the moment the upstream advertised a different base via LOCALAI_BASE_URL, a reverse proxy, or X-Forwarded-Host. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/media.go | 147 +++++++++++++---- backend/go/localai-proxy/media_test.go | 213 ++++++++++++++++++------- 2 files changed, 265 insertions(+), 95 deletions(-) diff --git a/backend/go/localai-proxy/media.go b/backend/go/localai-proxy/media.go index e2094a092..b4bff219d 100644 --- a/backend/go/localai-proxy/media.go +++ b/backend/go/localai-proxy/media.go @@ -5,6 +5,7 @@ import ( "encoding/base64" "encoding/json" "fmt" + "net/http" "net/url" "os" "strconv" @@ -32,6 +33,30 @@ func fileToBase64(path string) (string, error) { return base64.StdEncoding.EncodeToString(data), nil } +// toDataURI wraps a base64 payload as a data URI. Detect/Depth/FaceVerify/ +// FaceAnalyze's REST endpoints decode their image field with +// utils.GetContentURIAsBase64, which only accepts an http(s) URL or a +// `data:;base64,` string — never bare base64. That is exactly +// what core hands the backend for these methods (see +// core/http/endpoints/localai/images.go's decodeImageInput, which already +// stripped any data: prefix off before we ever see it), so forwarding it +// unwrapped 400s on every real call. The MIME type isn't carried alongside +// the payload, so it's sniffed from the decoded bytes. +func toDataURI(b64 string) (string, error) { + if b64 == "" { + return "", nil + } + data, err := base64.StdEncoding.DecodeString(b64) + if err != nil { + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: decode base64 image: %v", err) + } + mime := http.DetectContentType(data) + if mime == "application/octet-stream" { + mime = "image/png" + } + return "data:" + mime + ";base64," + b64, nil +} + // genItem is one schema.Item as the image/video/3D generation endpoints // return it: either inline base64 data, or a URL to download the asset from. type genItem struct { @@ -66,20 +91,42 @@ func (p *LocalAIProxy) writeGenItem(ctx context.Context, path string, items []ge if item.URL == "" { return status.Errorf(codes.Internal, "localai-proxy: upstream %s returned neither b64_json nor url", path) } - return p.getToFile(ctx, p.relativePath(item.URL), dst) + rel, err := generatedContentPath(path, item.URL) + if err != nil { + return err + } + return p.getToFile(ctx, rel, dst) } -// relativePath strips the configured upstream base from a URL the upstream -// handed back (e.g. "/generated-images/x.png"), so the follow-up -// download goes through the normal request path instead of concatenating two -// absolute URLs. -func (p *LocalAIProxy) relativePath(raw string) string { - if cfg := p.cfg.Load(); cfg != nil { - if rel, ok := strings.CutPrefix(raw, cfg.base); ok { - return rel +// 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 raw + return "", status.Errorf(codes.InvalidArgument, "localai-proxy: upstream %s returned an unexpected url %q", callPath, raw) } // --- Images --------------------------------------------------------- @@ -335,13 +382,19 @@ type detectResponseBody struct { Detections []detectionBody `json:"detections"` } -// Detect posts Src, which core already carries as a base64 payload (the same -// convention DetectionEndpoint uses to call this method locally), and maps -// each detection, decoding its PNG mask. +// 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: req.GetSrc(), + Image: image, Prompt: req.GetPrompt(), Points: req.GetPoints(), Boxes: req.GetBoxes(), @@ -371,17 +424,18 @@ func (p *LocalAIProxy) Detect(req *pb.DetectOptions) (pb.DetectResponse, error) 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"` - Dst string `json:"dst,omitempty"` - 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"` - Exports []string `json:"exports,omitempty"` + 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 { @@ -399,20 +453,29 @@ type depthResponseBody struct { IsMetric bool `json:"is_metric"` } -// Depth posts Src (a base64 payload, per the same convention as Detect) and -// maps the full response, decoding the point-cloud color bytes. +// 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: req.GetSrc(), - Dst: req.GetDst(), + Image: image, IncludeDepth: req.GetIncludeDepth(), IncludeConfidence: req.GetIncludeConfidence(), IncludePose: req.GetIncludePose(), IncludeSky: req.GetIncludeSky(), IncludePoints: req.GetIncludePoints(), PointsConfThresh: req.GetPointsConfThresh(), - Exports: req.GetExports(), } var resp depthResponseBody if err := p.postJSON(context.Background(), "/v1/depth", body, &resp); err != nil { @@ -472,11 +535,21 @@ type faceVerifyResponseBody struct { Img2AntispoofScore *float32 `json:"img2_antispoof_score,omitempty"` } -// FaceVerify posts Img1/Img2, which core already carries as base64 (the same -// convention FaceVerifyEndpoint uses to call this method locally). +// 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: req.GetImg1(), Img2: req.GetImg2(), + Model: p.model(""), Img1: img1, Img2: img2, Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(), } var resp faceVerifyResponseBody @@ -536,10 +609,16 @@ type faceAnalyzeResponseBody struct { Faces []faceAnalysisBody `json:"faces"` } -// FaceAnalyze posts Img, which core already carries as base64. +// 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: req.GetImg(), Actions: req.GetActions(), AntiSpoofing: req.GetAntiSpoofing(), + 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 { diff --git a/backend/go/localai-proxy/media_test.go b/backend/go/localai-proxy/media_test.go index 2e2a970e9..491845fe3 100644 --- a/backend/go/localai-proxy/media_test.go +++ b/backend/go/localai-proxy/media_test.go @@ -2,6 +2,7 @@ package main import ( "encoding/base64" + "encoding/json" "net/http" "os" "path/filepath" @@ -11,6 +12,7 @@ import ( "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 @@ -84,6 +86,38 @@ var _ = Describe("media methods", func() { 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() { @@ -211,105 +245,162 @@ var _ = Describe("media methods", func() { }) }) - Describe("Detect", func() { - It("posts the image and maps detections including the mask", func() { - p := loadProxy(up, nil) - mask := base64.StdEncoding.EncodeToString([]byte("png-mask")) - up.replyJSON("/v1/detection", 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}, - }, - }) + // 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")...) - res, err := p.Detect(&pb.DetectOptions{Src: "base64-image-data", Prompt: "cat", Threshold: 0.5}) + // 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"))) - - req := up.last() - Expect(req.Path).To(Equal("/v1/detection")) - Expect(req.JSON).To(HaveKeyWithValue("image", "base64-image-data")) - Expect(req.JSON).To(HaveKeyWithValue("prompt", "cat")) - Expect(req.JSON).To(HaveKeyWithValue("threshold", BeNumerically("~", 0.5, 1e-6))) }) }) Describe("Depth", func() { - It("posts the image and maps the full depth response", func() { - p := loadProxy(up, nil) - colors := base64.StdEncoding.EncodeToString([]byte("rgb")) - up.replyJSON("/v1/depth", map[string]any{ - "width": 2, "height": 1, "depth": []float64{0.1, 0.2}, - "point_colors": colors, "is_metric": true, - }) + 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)) - res, err := p.Depth(&pb.DepthRequest{Src: "base64-image-data", IncludeDepth: 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()) + }) - req := up.last() - Expect(req.Path).To(Equal("/v1/depth")) - Expect(req.JSON).To(HaveKeyWithValue("image", "base64-image-data")) - Expect(req.JSON).To(HaveKeyWithValue("include_depth", true)) + 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("posts both images and maps the response, including liveness fields", func() { - p := loadProxy(up, nil) - up.replyJSON("/v1/face/verify", 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, - }) + 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)) - res, err := p.FaceVerify(&pb.FaceVerifyRequest{Img1: "img1-b64", Img2: "img2-b64", AntiSpoofing: 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()) - - req := up.last() - Expect(req.Path).To(Equal("/v1/face/verify")) - Expect(req.JSON).To(HaveKeyWithValue("img1", "img1-b64")) - Expect(req.JSON).To(HaveKeyWithValue("img2", "img2-b64")) - Expect(req.JSON).To(HaveKeyWithValue("anti_spoofing", true)) }) }) Describe("FaceAnalyze", func() { - It("posts the image and maps per-face demographic attributes", func() { - p := loadProxy(up, nil) - up.replyJSON("/v1/face/analyze", 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}, - }, - }, - }) + 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")) - res, err := p.FaceAnalyze(&pb.FaceAnalyzeRequest{Img: "img-b64", Actions: []string{"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)))) - - req := up.last() - Expect(req.Path).To(Equal("/v1/face/analyze")) - Expect(req.JSON).To(HaveKeyWithValue("img", "img-b64")) - Expect(req.JSON).To(HaveKeyWithValue("actions", ConsistOf("age", "gender"))) }) }) From c2375a7db8f48426695e7fc02c63dfc5d15724f3 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 00:14:32 +0000 Subject: [PATCH 48/79] feat(localai-proxy): bridge live transcription to the upstream realtime API AudioTranscriptionLive opens /v1/realtime?model= as a transcription session with server VAD, forwards PCM as base64 PCM16 appends, and maps transcription deltas and completions to Delta/Eou. Closing the send side waits briefly for an in-flight utterance, then sends the final text. An upstream error, failed transcription or disconnect ends the stream with Unavailable so failover reopens on the next target. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/live.go | 425 ++++++++++++++++++++++++++ backend/go/localai-proxy/live_test.go | 380 +++++++++++++++++++++++ 2 files changed, 805 insertions(+) create mode 100644 backend/go/localai-proxy/live.go create mode 100644 backend/go/localai-proxy/live_test.go diff --git a/backend/go/localai-proxy/live.go b/backend/go/localai-proxy/live.go new file mode 100644 index 000000000..622358729 --- /dev/null +++ b/backend/go/localai-proxy/live.go @@ -0,0 +1,425 @@ +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 VAD already started. 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 + + // 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 is + // bounded by request_timeout_seconds instead. + liveHandshakeTimeout = 30 * time.Second +) + +// 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 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{}, + 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() + + return s.run(in, cfg.timeout) +} + +// 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 + pending int // utterances started upstream and not yet transcribed + + 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 + setupTimer <-chan time.Time + drainTimer <-chan time.Time + drain *time.Timer + ) + defer func() { + if drain != nil { + drain.Stop() + } + }() + if setupTimeout > 0 { + t := time.NewTimer(setupTimeout) + defer t.Stop() + setupTimer = t.C + } + + // finishOrDrain finalizes once nothing is in flight, else waits (bounded) + // for the upstream to finish the utterance it already started. + finishOrDrain := func() (bool, error) { + if s.pending <= 0 { + return true, s.finish() + } + if drain == nil { + drain = time.NewTimer(finalWait) + drainTimer = drain.C + } + return false, nil + } + + for { + select { + case req, ok := <-in: + if !ok { + in, inClosed = nil, true + if ready { + if done, err := finishOrDrain(); done { + return err + } + } + continue + } + if !ready { + // Callers wait for the ready ack before streaming, but hold + // anything sent early rather than drop it. + backlog = append(backlog, req) + continue + } + 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 + if inClosed { + if done, err := finishOrDrain(); done { + return err + } + } + 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.pending++ + case "conversation.item.input_audio_transcription.delta": + s.delta(ev) + case "conversation.item.input_audio_transcription.completed": + s.completed(ev) + if inClosed && s.pending <= 0 { + return s.finish() + } + } + + case <-setupTimer: + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s did not set up the transcription session in %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() + } + } +} + +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 := math.Max(-1, math.Min(1, float64(f))) + 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 = "" + } + if s.pending > 0 { + s.pending-- + } + if t := strings.TrimSpace(ev.Transcript); t != "" { + s.final = append(s.final, t) + } + s.out <- &pb.TranscriptLiveResponse{Delta: rest, Eou: true} +} + +// finish sends the final transcript and closes the upstream session cleanly. +func (s *liveSession) finish() error { + s.out <- &pb.TranscriptLiveResponse{FinalResult: &pb.TranscriptResult{Text: strings.Join(s.final, " ")}} + _ = s.conn.WriteControl(websocket.CloseMessage, + websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second)) + return nil +} + +func (s *liveSession) write(v any) error { + if err := s.conn.SetWriteDeadline(time.Now().Add(liveWriteTimeout)); err != nil { + return s.writeFailed(err) + } + if err := s.conn.WriteJSON(v); err != nil { + return s.writeFailed(err) + } + return nil +} + +func (s *liveSession) writeFailed(err error) error { + xlog.Warn("localai-proxy: realtime write failed", "pipeline", s.pipeline, "error", err) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s: %v", realtimePath, err) +} + +// upstreamGone maps a socket that stopped delivering events. It is +// Unavailable even for a clean close: the client did not end the session, so +// the upstream dropped it, and failover should reopen on the next target. +func (s *liveSession) upstreamGone() error { + err := s.readErr + if err == nil { + err = errors.New("connection closed") + } + xlog.Warn("localai-proxy: realtime upstream disconnected", "pipeline", s.pipeline, "error", err) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s disconnected: %v", realtimePath, err) +} + +func (s *liveSession) upstreamError(what string, ev realtimeEvent) error { + msg := "no details" + if ev.Error != nil && ev.Error.Message != "" { + msg = ev.Error.Message + if ev.Error.Code != "" { + msg = fmt.Sprintf("%s (%s)", msg, ev.Error.Code) + } + } + xlog.Warn("localai-proxy: realtime upstream reported an error", "pipeline", s.pipeline, "event", what, "error", msg) + return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s %s: %s", realtimePath, what, msg) +} diff --git a/backend/go/localai-proxy/live_test.go b/backend/go/localai-proxy/live_test.go new file mode 100644 index 000000000..22e30e6d1 --- /dev/null +++ b/backend/go/localai-proxy/live_test.go @@ -0,0 +1,380 @@ +package main + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "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 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) + close(lc.in) + Expect(lc.finish()).To(Succeed()) + + Eventually(frames, 2*time.Second).Should(Receive(Equal([][]int16{ + {0, 16383, -32767, 32767, 32767}, + {-8191}, + }))) + }) + + 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, map[string]any{"type": "input_audio_buffer.speech_started"}) + 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, map[string]any{"type": "input_audio_buffer.speech_stopped"}) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "a", "transcript": "hello world"}) + wsSend(c, map[string]any{"type": "input_audio_buffer.speech_started"}) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "b", "transcript": "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 an in-flight utterance before the final result", func() { + up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { + wsHandshake(c) + wsSend(c, map[string]any{"type": "input_audio_buffer.speech_started"}) + // The client closes its side now; the transcription lands later. + time.Sleep(300 * time.Millisecond) + wsSend(c, map[string]any{"type": "input_audio_buffer.speech_stopped"}) + wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "a", "transcript": "late words"}) + wsDrain(c) + }) + DeferCleanup(up.Close) + lc := startLive(loadProxy(up, withPipeline)) + lc.config("en", 16000) + Expect(lc.next().GetReady()).To(BeTrue()) + close(lc.in) + + r := lc.next() + Expect(r.GetDelta()).To(Equal("late words")) + Expect(r.GetEou()).To(BeTrue()) + Expect(lc.next().GetFinalResult().GetText()).To(Equal("late words")) + Expect(lc.finish()).To(Succeed()) + }) + + 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) + }) +}) From e66d87b88770539dd59ecc38e17e28eeeff9ad44 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 00:22:50 +0000 Subject: [PATCH 49/79] fix(localai-proxy): bound live transcription setup and track in-flight turns precisely A caller that gives up before the ready ack now ends the call with Canceled and closes the upstream socket, and setup has a 3 minute default bound when request_timeout_seconds is unset, so a hung upstream cannot hold the call or block failover. A closing session now waits only for a turn the upstream committed (or is still speaking), not for turns it discarded. Audio held before the ready ack is capped at 5 s, and NaN samples become silence. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- backend/go/localai-proxy/live.go | 187 +++++++++++++++++--------- backend/go/localai-proxy/live_test.go | 175 ++++++++++++++++++++++-- 2 files changed, 284 insertions(+), 78 deletions(-) diff --git a/backend/go/localai-proxy/live.go b/backend/go/localai-proxy/live.go index 622358729..0193b973a 100644 --- a/backend/go/localai-proxy/live.go +++ b/backend/go/localai-proxy/live.go @@ -28,27 +28,53 @@ const ( defaultLiveSampleRate = 16000 // finalWait bounds how long a closing session waits for an utterance the - // upstream VAD already started. The upstream transcribes only after its + // 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 is - // bounded by request_timeout_seconds instead. + // 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"` + 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"` @@ -99,21 +125,26 @@ func (p *LocalAIProxy) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveReques defer func() { _ = conn.Close() }() s := &liveSession{ - conn: conn, - out: out, - pipeline: cfg.realtimePipeline, - language: lc.GetLanguage(), - rate: rate, - sent: map[string]string{}, - events: make(chan realtimeEvent), - done: make(chan struct{}), + 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() - return s.run(in, cfg.timeout) + setup := cfg.timeout + if setup <= 0 { + setup = liveSetupTimeout + } + return s.run(in, setup) } // dialRealtime opens the upstream WebSocket. A refused upgrade is mapped like @@ -159,9 +190,16 @@ type liveSession struct { language string rate int - sent map[string]string // text already sent as deltas, per upstream item - final []string // completed transcripts, in order - pending int // utterances started upstream and not yet transcribed + 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 @@ -199,53 +237,47 @@ func (s *liveSession) run(in <-chan *pb.TranscriptLiveRequest, setupTimeout time created, ready bool inClosed bool backlog []*pb.TranscriptLiveRequest - setupTimer <-chan time.Time + backlogSamples int drainTimer <-chan time.Time - drain *time.Timer + graceTimer <-chan time.Time ) - defer func() { - if drain != nil { - drain.Stop() - } - }() - if setupTimeout > 0 { - t := time.NewTimer(setupTimeout) - defer t.Stop() - setupTimer = t.C - } - - // finishOrDrain finalizes once nothing is in flight, else waits (bounded) - // for the upstream to finish the utterance it already started. - finishOrDrain := func() (bool, error) { - if s.pending <= 0 { - return true, s.finish() - } - if drain == nil { - drain = time.NewTimer(finalWait) - drainTimer = drain.C - } - return false, nil - } + 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 { - if done, err := finishOrDrain(); done { - return err - } + 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) } - continue - } - if !ready { - // Callers wait for the ready ack before streaming, but hold - // anything sent early rather than drop it. backlog = append(backlog, req) - continue - } - if err := s.forward(req); err != nil { + } else if err := s.forward(req); err != nil { return err } @@ -275,37 +307,54 @@ func (s *liveSession) run(in <-chan *pb.TranscriptLiveRequest, setupTimeout time } } backlog = nil - if inClosed { - if done, err := finishOrDrain(); done { - return err - } - } 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.pending++ + 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) - if inClosed && s.pending <= 0 { - return s.finish() - } } + 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 in %s", realtimePath, setupTimeout) + 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", @@ -343,7 +392,13 @@ func (s *liveSession) forward(req *pb.TranscriptLiveRequest) error { func pcm16LE(pcm []float32) []byte { buf := make([]byte, len(pcm)*2) for i, f := range pcm { - v := math.Max(-1, math.Min(1, float64(f))) + 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)) binary.LittleEndian.PutUint16(buf[i*2:], uint16(int16(v*math.MaxInt16))) } return buf @@ -368,9 +423,7 @@ func (s *liveSession) completed(ev realtimeEvent) { if !ok { rest = "" } - if s.pending > 0 { - s.pending-- - } + delete(s.committed, ev.ItemID) if t := strings.TrimSpace(ev.Transcript); t != "" { s.final = append(s.final, t) } diff --git a/backend/go/localai-proxy/live_test.go b/backend/go/localai-proxy/live_test.go index 22e30e6d1..c35b4463b 100644 --- a/backend/go/localai-proxy/live_test.go +++ b/backend/go/localai-proxy/live_test.go @@ -4,6 +4,7 @@ import ( "encoding/base64" "encoding/binary" "encoding/json" + "math" "net/http" "runtime" "strings" @@ -147,6 +148,27 @@ func liveGoroutines() int { 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") } @@ -249,25 +271,30 @@ var _ = Describe("AudioTranscriptionLive", func() { 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, map[string]any{"type": "input_audio_buffer.speech_started"}) + 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, map[string]any{"type": "input_audio_buffer.speech_stopped"}) - wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "a", "transcript": "hello world"}) - wsSend(c, map[string]any{"type": "input_audio_buffer.speech_started"}) - wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "b", "transcript": "again"}) + 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) @@ -291,29 +318,58 @@ var _ = Describe("AudioTranscriptionLive", func() { Expect(lc.finish()).To(Succeed()) }) - It("waits for an in-flight utterance before the final result", func() { + It("waits for a committed utterance before the final result", func() { up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { wsHandshake(c) - wsSend(c, map[string]any{"type": "input_audio_buffer.speech_started"}) - // The client closes its side now; the transcription lands later. + 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, map[string]any{"type": "input_audio_buffer.speech_stopped"}) - wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": "a", "transcript": "late words"}) + 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("late words")) + 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{}}) @@ -377,4 +433,101 @@ var _ = Describe("AudioTranscriptionLive", func() { 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) + }) }) From 4a2a6180a2eec8ef48fdc501dfce79f1588a029b Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 00:32:13 +0000 Subject: [PATCH 50/79] test(localai-proxy): proxy APIs and realtime stages end to end The e2e suite now registers the localai-proxy binary and points proxy models back at the test server itself, so a request leaves LocalAI through the backend, returns over REST and is answered by a mock model. Chat, embeddings, TTS and transcription through the proxy return the upstream model's answer; a chain whose proxy target's upstream model fails to load serves from the local target; and a realtime pipeline whose LLM stage is a chain on a remote target completes a turn, then switches to the local target with a localai.model.failover trip event when a gate in front of the upstream starts answering 503. The docs describe the localai-proxy backend next to cloud-proxy and add a per-stage remote LocalAI example to the failover page. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- docs/content/features/model-failover.md | 75 ++++++++- docs/content/operations/cloud-proxy.md | 90 ++++++++++ tests/e2e/e2e_failover_test.go | 23 +-- tests/e2e/e2e_localai_proxy_test.go | 215 ++++++++++++++++++++++++ tests/e2e/e2e_suite_test.go | 21 +++ tests/e2e/realtime_ws_test.go | 63 +++++++ 6 files changed, 461 insertions(+), 26 deletions(-) create mode 100644 tests/e2e/e2e_localai_proxy_test.go diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index 32e91af96..8f62a0314 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -20,7 +20,7 @@ provider, and to fall back to a local model when the remote one is down. name: assistant-llm failover: targets: - - model: argus-llm # for example a cloud-proxy model + - model: argus-llm # for example a localai-proxy or cloud-proxy model - model: gemma-local warm: true # keep it loaded ``` @@ -50,9 +50,9 @@ Rules: - A chain cannot also set `alias` or `backend`. - Responses name the chain as the model. The `X-LocalAI-Served-Model` header names the target that served the request. -- A remote (`cloud-proxy`) target receives its own model name, never the chain - name: `proxy.upstream_model`, or the target name when `upstream_model` is - empty. The health check looks for the same name. +- A remote (`localai-proxy` or `cloud-proxy`) target receives its own model + name, never the chain name: `proxy.upstream_model`, or the target name when + `upstream_model` is empty. The health check looks for the same name. ## How the target is chosen @@ -89,7 +89,7 @@ targets were down. | Target | Regular check | Check before moving back | |---|---|---| -| Remote (`cloud-proxy`) | `GET /v1/models` on the upstream lists the model | one small real request, for example a 1-token completion | +| Remote (`localai-proxy`, `cloud-proxy`) | `GET /v1/models` on the upstream lists the model | one small real request, for example a 1-token completion | | Local, `warm: true` | the backend answers a health check. A check never loads the model: while it is not loaded, the check passes and real requests judge it | one small real request. While the model is not loaded, the target is used again after `min_dwell` | | Local, not warm | none: judged only by real requests; it is never loaded only to check it | none: the target is used again after `min_dwell` | @@ -107,8 +107,9 @@ count toward the active backend limit (`--max-active-backends`) like any pinned model: LocalAI never evicts them to make room, and if they fill the limit, a new model still loads rather than being blocked. -`warm` applies only to local targets. On a remote (`cloud-proxy`) target it -has no effect, and LocalAI logs a warning when it loads the chain. +`warm` applies only to local targets. On a remote (`localai-proxy` or +`cloud-proxy`) target it has no effect, and LocalAI logs a warning when it +loads the chain. ## Realtime pipelines @@ -136,6 +137,66 @@ it starts (`reason: initial`) and each time a chain switches: "from":"argus-llm","to":"gemma-local","state":"fallback","reason":"trip"} ``` +### Example: stages on a remote LocalAI + +This pipeline runs its transcription, LLM and TTS stages on a remote LocalAI +(`argus`) through [`localai-proxy`]({{% relref "operations/cloud-proxy" %}}) +models, and uses local models when the remote instance is down. Each stage has +its own chain, so one stage can fail over while the others stay remote. + +```yaml +# Remote targets: each one names the model on the upstream LocalAI. +name: argus-stt +backend: localai-proxy +known_usecases: [transcript] +options: + - realtime_pipeline:asr-pipeline # upstream pipeline for live transcription +proxy: + upstream_url: http://argus.lan:8080 + upstream_model: parakeet +--- +name: argus-llm +backend: localai-proxy +known_usecases: [chat] +proxy: + upstream_url: http://argus.lan:8080 + upstream_model: gemma-3-12b +--- +name: argus-tts +backend: localai-proxy +known_usecases: [tts] +proxy: + upstream_url: http://argus.lan:8080 + upstream_model: kokoro +--- +# One chain per stage, remote first, local second. +name: stt-chain +failover: + targets: [{model: argus-stt}, {model: whisper-local, warm: true}] +--- +name: llm-chain +failover: + targets: [{model: argus-llm}, {model: gemma-local, warm: true}] +--- +name: tts-chain +failover: + targets: [{model: argus-tts}, {model: piper-local}] +--- +name: assistant +pipeline: + vad: silero-vad + transcription: stt-chain + llm: llm-chain + tts: tts-chain +``` + +The example shows the configs as one YAML stream; put each config in its own +file in the models directory. When `argus` stops +answering, the next call of each stage fails over to the local model and the +session receives a `localai.model.failover` event for that stage. A remote +target that does not support a call (it returns `Unimplemented`) is skipped for +that call and is not marked down. + Limits: - After a `session.update` that changes the pipeline, `localai.model.failover` diff --git a/docs/content/operations/cloud-proxy.md b/docs/content/operations/cloud-proxy.md index bd1d806d4..ec178e2f8 100644 --- a/docs/content/operations/cloud-proxy.md +++ b/docs/content/operations/cloud-proxy.md @@ -242,6 +242,96 @@ ACLs, and the cloud-proxy fork all run against the resolved target. See [Middleware: PII filtering and intelligent routing]({{< relref "middleware.md" >}}) for the full router and PII-filter reference. +## Proxying to another LocalAI (`localai-proxy`) + +`cloud-proxy` forwards chat and Messages requests only. To serve a model from +another LocalAI instance for every API it has, use `backend: localai-proxy`. +The backend receives the request from the local pipeline like any other +backend and sends it to the REST API of the upstream LocalAI. Because it is a +normal backend, a `localai-proxy` model can be a stage of a realtime pipeline +or a target of a [failover chain]({{% relref "features/model-failover" %}}). + +```yaml +name: remote-llm +backend: localai-proxy +known_usecases: [chat] +proxy: + # Base URL of the upstream LocalAI. Do not add /v1 or an endpoint path: + # the backend adds the path for each API. + upstream_url: https://argus.lan:8080 + # The model name on the upstream. When empty, the name of this config. + upstream_model: gemma-3-12b + # Optional. The upstream API key, from an environment variable + # (or api_key_file). Sent as "Authorization: Bearer ". + api_key_env: ARGUS_API_KEY + # Optional. Time limit for each non-streaming request. Streams have no limit. + request_timeout_seconds: 120 +``` + +A model that does live transcription in a realtime pipeline also names a +realtime pipeline on the upstream. The backend opens a transcription session +on the upstream `/v1/realtime` endpoint with that pipeline: + +```yaml +name: remote-stt +backend: localai-proxy +known_usecases: [transcript] +options: + - realtime_pipeline:asr-pipeline +proxy: + upstream_url: https://argus.lan:8080 + upstream_model: parakeet +``` + +Set `known_usecases` on every `localai-proxy` model. Failover uses it to match +targets, and LocalAI cannot guess the usecases of a remote model. For a chat +model, `known_usecases: [chat]` has one more effect: LocalAI sends the chat +messages to the upstream `/v1/chat/completions` endpoint, and the upstream +applies its own chat template, tool parsing and reasoning parsing. Without +`chat`, or when the config has its own templates, LocalAI renders the prompt +locally and sends it to `/v1/completions`. `proxy.mode` and `proxy.provider` +have no effect on this backend. + +Supported APIs: + +- Text: chat and completions (also streamed), embeddings, rerank, tokenize, + detokenize, score. +- Audio: TTS (also streamed), sound generation, transcription (also streamed), + live transcription (with `realtime_pipeline`), diarization, VAD, sound + classification, audio transformations. +- Image, video and 3D: image generation, upscaling, video generation, 3D + generation and animation. The backend downloads the files that the upstream + generates. +- Vision: object detection, depth, face verification and analysis, voice + verification, analysis and embeddings. +- Stores: set, get, delete, find. + +Methods that have no REST API on the upstream return the gRPC error +`Unimplemented` ("localai-proxy: has no upstream counterpart"): +audio encoding and decoding, audio-to-audio streams, token classification +(PII NER), model metadata, fine-tuning, quantization and model export. A +failover chain skips a target that returns `Unimplemented` and tries the next +target, but does not mark the target down. + +Errors from the upstream: a 5xx response or a connection failure becomes +`Unavailable`, and a failover chain marks the target down. A 4xx response +becomes `InvalidArgument`, and LocalAI returns it to the client without a +retry. + +Known limits: + +- Voice-profile paths pass through unresolved. When LocalAI resolves a TTS + voice to a local file (for example a voice clone reference), the backend + sends that path to the upstream, where it does not exist. Use voices that + the upstream knows by name. +- Depth exports are not supported. The upstream writes them to its own disk, + so a depth request with exports or a destination file returns + `Unimplemented`. Depth maps and points without exports work. +- The REST transcription API has no end-of-utterance (`eou`) flag, so + transcriptions through the proxy never set it. Live transcription through + `realtime_pipeline` sets `eou` at the end of each utterance. +- Sound generation from a source audio file is not supported. + ## Limitations - **Passthrough does no wire-shape translation.** Use `mode: translate` (with diff --git a/tests/e2e/e2e_failover_test.go b/tests/e2e/e2e_failover_test.go index b8e1632d2..22a4226f6 100644 --- a/tests/e2e/e2e_failover_test.go +++ b/tests/e2e/e2e_failover_test.go @@ -6,23 +6,13 @@ import ( "io" "mime/multipart" "net/http" - "os" - "path/filepath" "time" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "gopkg.in/yaml.v3" ) var _ = Describe("Failover chains", Label("failover"), func() { - postJSON := func(path string, body map[string]any) *http.Response { - b, err := json.Marshal(body) - Expect(err).ToNot(HaveOccurred()) - resp, err := http.Post(apiURL+path, "application/json", bytes.NewReader(b)) - Expect(err).ToNot(HaveOccurred()) - return resp - } expectServedByMock := func(resp *http.Response) { defer func() { _ = resp.Body.Close() }() body, _ := io.ReadAll(resp.Body) @@ -34,7 +24,7 @@ var _ = Describe("Failover chains", Label("failover"), func() { // The entry name is the chain suffix: chain- is written by the suite. DescribeTable("retries every endpoint family on the next target", func(path string, body func(model string) map[string]any) { - expectServedByMock(postJSON(path, body("chain-"+CurrentSpecReport().LeafNodeText))) + expectServedByMock(postJSONTo(path, body("chain-"+CurrentSpecReport().LeafNodeText))) }, Entry("chat", "/chat/completions", func(m string) map[string]any { return map[string]any{"model": m, "messages": []map[string]string{{"role": "user", "content": "hi"}}} @@ -99,7 +89,7 @@ var _ = Describe("Failover chains", Label("failover"), func() { }) up2.SetScript(chatReply) - resp := postJSON("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) + resp := postJSONTo("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) defer func() { _ = resp.Body.Close() }() body, _ := io.ReadAll(resp.Body) Expect(resp.StatusCode).To(Equal(200), string(body)) @@ -116,7 +106,7 @@ var _ = Describe("Failover chains", Label("failover"), func() { Eventually(func() string { return chainActive("chain-remote") }, 30*time.Second, 500*time.Millisecond). Should(Equal("up-1")) - resp2 := postJSON("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) + resp2 := postJSONTo("/chat/completions", map[string]any{"model": "chain-remote", "messages": []map[string]string{{"role": "user", "content": "hi"}}}) defer func() { _ = resp2.Body.Close() }() Expect(resp2.StatusCode).To(Equal(200)) Expect(resp2.Header.Get("X-LocalAI-Served-Model")).To(Equal("up-1")) @@ -178,12 +168,7 @@ func registerFailoverRemoteModels(url1, url2 string) { "recovery": map[string]any{"probes": 2, "min_dwell": "2s"}, }, } - for _, cfg := range []map[string]any{proxyModel("up-1", url1), proxyModel("up-2", url2), chain} { - data, err := yaml.Marshal(cfg) - Expect(err).ToNot(HaveOccurred()) - Expect(os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) - } - Expect(localAIApp.ModelConfigLoader().LoadModelConfigsFromPath(modelsPath)).To(Succeed()) + registerModelConfigs(proxyModel("up-1", url1), proxyModel("up-2", url2), chain) Eventually(func() string { return chainActive("chain-remote") }, 10*time.Second, 200*time.Millisecond). Should(Equal("up-1")) } diff --git a/tests/e2e/e2e_localai_proxy_test.go b/tests/e2e/e2e_localai_proxy_test.go new file mode 100644 index 000000000..fc86f612e --- /dev/null +++ b/tests/e2e/e2e_localai_proxy_test.go @@ -0,0 +1,215 @@ +package e2e_test + +import ( + "bytes" + "encoding/json" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "net/http/httputil" + "net/url" + "os" + "path/filepath" + "sync/atomic" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gopkg.in/yaml.v3" +) + +// The localai-proxy specs point the backend at this same test server: a +// request to an lp-* model leaves LocalAI through the localai-proxy process, +// comes back in over REST and is answered by the mock model it names, so the +// whole round trip (core -> gRPC -> REST upstream -> gRPC reply) is real. +var _ = Describe("localai-proxy backend", Label("failover"), Ordered, func() { + BeforeAll(func() { + if localAIProxyPath == "" { + Skip("localai-proxy backend binary not built (make build-localai-proxy-backend)") + } + registerModelConfigs( + localAIProxyModel("lp-chat", anthropicBaseURL, "mock-model", "chat"), + localAIProxyModel("lp-embeddings", anthropicBaseURL, "mock-model", "embeddings"), + localAIProxyModel("lp-tts", anthropicBaseURL, "mock-model", "tts"), + localAIProxyModel("lp-transcription", anthropicBaseURL, "mock-model", "transcript"), + // The upstream serves this target from a model whose load always + // fails, so every request through it errors with a 5xx. + localAIProxyModel("lp-broken", anthropicBaseURL, "fail-chat", "chat"), + map[string]any{ + "name": "chain-lp", + "failover": map[string]any{ + "targets": []map[string]any{{"model": "lp-broken"}, {"model": "mock-model"}}, + }, + }, + ) + Eventually(func() string { return chainActive("chain-lp") }, 10*time.Second, 200*time.Millisecond). + Should(Equal("lp-broken")) + }) + + // sameBody posts the same request to the proxied model and to the model + // the upstream serves it from, and returns both bodies: the proxy answered + // with the upstream's answer when they match. + sameBody := func(path string, body func(model string) map[string]any, proxied, upstream string) (string, string) { + get := func(model string) string { + resp := postJSONTo(path, body(model)) + defer func() { _ = resp.Body.Close() }() + b, err := io.ReadAll(resp.Body) + Expect(err).ToNot(HaveOccurred()) + Expect(resp.StatusCode).To(Equal(http.StatusOK), "%s: %s", model, b) + return string(b) + } + return get(proxied), get(upstream) + } + + It("answers chat with the upstream model's reply", func() { + chat := func(m string) map[string]any { + return map[string]any{"model": m, "messages": []map[string]string{{"role": "user", "content": "hi"}}} + } + got, want := sameBody("/chat/completions", chat, "lp-chat", "mock-model") + Expect(chatContent(got)).ToNot(BeEmpty()) + Expect(chatContent(got)).To(Equal(chatContent(want))) + }) + + It("answers embeddings with the upstream model's vector", func() { + embed := func(m string) map[string]any { return map[string]any{"model": m, "input": "hello"} } + got, want := sameBody("/embeddings", embed, "lp-embeddings", "mock-model") + Expect(embeddingVector(got)).ToNot(BeEmpty()) + Expect(embeddingVector(got)).To(Equal(embeddingVector(want))) + }) + + It("answers TTS with the upstream model's audio", func() { + speech := func(m string) map[string]any { return map[string]any{"model": m, "input": "hello", "voice": "default"} } + got, want := sameBody("/audio/speech", speech, "lp-tts", "mock-model") + Expect(got).To(HavePrefix("RIFF")) + Expect(got).To(Equal(want)) + }) + + It("answers transcription with the upstream model's text", func() { + transcribe := func(model string) string { + var body bytes.Buffer + mw := multipart.NewWriter(&body) + Expect(mw.WriteField("model", model)).To(Succeed()) + fw, err := mw.CreateFormFile("file", "a.wav") + Expect(err).ToNot(HaveOccurred()) + _, err = fw.Write(wavFromPCM(make([]byte, 6400), 16000)) + Expect(err).ToNot(HaveOccurred()) + Expect(mw.Close()).To(Succeed()) + resp, err := http.Post(apiURL+"/audio/transcriptions", mw.FormDataContentType(), &body) + Expect(err).ToNot(HaveOccurred()) + defer func() { _ = resp.Body.Close() }() + b, _ := io.ReadAll(resp.Body) + Expect(resp.StatusCode).To(Equal(http.StatusOK), "%s: %s", model, b) + var out struct { + Text string `json:"text"` + } + Expect(json.Unmarshal(b, &out)).To(Succeed(), string(b)) + return out.Text + } + got := transcribe("lp-transcription") + Expect(got).To(HavePrefix("transcribed:")) + Expect(got).To(Equal(transcribe("mock-model"))) + }) + + It("fails over from a proxy target whose upstream model fails to the local target", func() { + resp := postJSONTo("/chat/completions", map[string]any{ + "model": "chain-lp", + "messages": []map[string]string{{"role": "user", "content": "hi"}}, + }) + defer func() { _ = resp.Body.Close() }() + b, _ := io.ReadAll(resp.Body) + Expect(resp.StatusCode).To(Equal(http.StatusOK), string(b)) + Expect(resp.Header.Get("X-LocalAI-Served-Model")).To(Equal("mock-model")) + Expect(resp.Header.Get("X-LocalAI-Failover")).To(Equal("fallback")) + Expect(chainActive("chain-lp")).To(Equal("mock-model")) + }) +}) + +// postJSONTo posts body as JSON to an /v1 path of the test server. +func postJSONTo(path string, body map[string]any) *http.Response { + b, err := json.Marshal(body) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + resp, err := http.Post(apiURL+path, "application/json", bytes.NewReader(b)) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + return resp +} + +func chatContent(body string) string { + var out struct { + Choices []struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + } `json:"choices"` + } + ExpectWithOffset(1, json.Unmarshal([]byte(body), &out)).To(Succeed(), body) + ExpectWithOffset(1, out.Choices).ToNot(BeEmpty(), body) + return out.Choices[0].Message.Content +} + +func embeddingVector(body string) []float32 { + var out struct { + Data []struct { + Embedding []float32 `json:"embedding"` + } `json:"data"` + } + ExpectWithOffset(1, json.Unmarshal([]byte(body), &out)).To(Succeed(), body) + ExpectWithOffset(1, out.Data).ToNot(BeEmpty(), body) + return out.Data[0].Embedding +} + +// localAIProxyModel is a localai-proxy config serving upstreamModel from the +// LocalAI at baseURL. known_usecases is what failover matches targets on, and +// "chat" also makes the proxy send structured messages upstream. +func localAIProxyModel(name, baseURL, upstreamModel string, usecases ...string) map[string]any { + return map[string]any{ + "name": name, + "backend": "localai-proxy", + "known_usecases": usecases, + "parameters": map[string]any{"model": name + ".bin"}, + "proxy": map[string]any{ + "upstream_url": baseURL, + "upstream_model": upstreamModel, + }, + } +} + +// registerModelConfigs writes model YAMLs after startup and has the loader +// re-read the models directory, for configs that embed runtime URLs. The +// failover manager picks new chains up on its next tick. +func registerModelConfigs(cfgs ...map[string]any) { + for _, cfg := range cfgs { + data, err := yaml.Marshal(cfg) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + ExpectWithOffset(1, os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) + } + ExpectWithOffset(1, localAIApp.ModelConfigLoader().LoadModelConfigsFromPath(modelsPath)).To(Succeed()) +} + +// upstreamGate is a reverse proxy in front of the test server that a spec +// can take down: while down it answers 503, as an upstream LocalAI with no +// healthy backend would, so a localai-proxy target behind it starts failing +// without restarting its backend process. +type upstreamGate struct { + srv *httptest.Server + down atomic.Bool +} + +func newUpstreamGate(target string) *upstreamGate { + u, err := url.Parse(target) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + g := &upstreamGate{} + rp := httputil.NewSingleHostReverseProxy(u) + g.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if g.down.Load() { + http.Error(w, `{"error":{"message":"no healthy backend"}}`, http.StatusServiceUnavailable) + return + } + rp.ServeHTTP(w, r) + })) + return g +} + +func (g *upstreamGate) URL() string { return g.srv.URL } +func (g *upstreamGate) Close() { g.srv.Close() } +func (g *upstreamGate) SetDown(down bool) { g.down.Store(down) } diff --git a/tests/e2e/e2e_suite_test.go b/tests/e2e/e2e_suite_test.go index e7bb844d5..0d1f51608 100644 --- a/tests/e2e/e2e_suite_test.go +++ b/tests/e2e/e2e_suite_test.go @@ -39,6 +39,7 @@ var ( apiURL string mockBackendPath string cloudProxyPath string + localAIProxyPath string mcpServerURL string mcpServerShutdown func() localAIApp *localaiapp.Application @@ -646,6 +647,23 @@ var _ = BeforeSuite(func() { } } + // localai-proxy backend: its models point back at this server, whose URL + // exists only once it listens, so the specs register them at runtime. + // Like cloud-proxy, a missing binary makes those specs Skip. + for _, p := range []string{ + filepath.Join("..", "e2e", "mock-backend", "localai-proxy"), + filepath.Join("tests", "e2e", "mock-backend", "localai-proxy"), + filepath.Join("..", "..", "tests", "e2e", "mock-backend", "localai-proxy"), + } { + if _, err := os.Stat(p); err == nil { + localAIProxyPath = p + break + } + } + if localAIProxyPath != "" { + Expect(os.Chmod(localAIProxyPath, 0755)).To(Succeed()) + } + // Live PII NER tier. When PII_NER_MODEL_GGUF points at a downloaded // privacy-filter GGUF, register two detector models that drive the real // gRPC TokenClassify path on the privacy-filter backend (discovered via @@ -703,6 +721,9 @@ var _ = BeforeSuite(func() { if cloudProxyPath != "" { localAIApp.ModelLoader().SetExternalBackend("cloud-proxy", cloudProxyPath) } + if localAIProxyPath != "" { + localAIApp.ModelLoader().SetExternalBackend("localai-proxy", localAIProxyPath) + } // Create HTTP app app, err = httpapi.API(localAIApp) diff --git a/tests/e2e/realtime_ws_test.go b/tests/e2e/realtime_ws_test.go index 68181ef35..5556a608f 100644 --- a/tests/e2e/realtime_ws_test.go +++ b/tests/e2e/realtime_ws_test.go @@ -273,6 +273,69 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { Expect(retrieveItem(conn, firstReplyID)).To(HaveKeyWithValue("id", firstReplyID)) }) + It("serves the LLM stage from a remote LocalAI and switches to the local target when it fails", func() { + if localAIProxyPath == "" { + Skip("localai-proxy backend binary not built (make build-localai-proxy-backend)") + } + // The remote target reaches this server through a gate the spec + // can take down, as if the remote LocalAI lost its backends. + gate := newUpstreamGate(anthropicBaseURL) + DeferCleanup(gate.Close) + registerModelConfigs( + localAIProxyModel("lp-rt-llm", gate.URL(), "mock-llm", "chat"), + map[string]any{ + "name": "chain-rt-lp", + "failover": map[string]any{ + "targets": []map[string]any{{"model": "lp-rt-llm"}, {"model": "mock-llm"}}, + }, + }, + map[string]any{ + "name": "rt-lp", + "pipeline": map[string]any{ + "vad": "mock-vad", + "transcription": "mock-stt", + "llm": "chain-rt-lp", + "tts": "mock-tts", + "disable_warmup": true, + }, + }, + ) + Eventually(func() string { return chainActive("chain-rt-lp") }, 10*time.Second, 200*time.Millisecond). + Should(Equal("lp-rt-llm")) + + conn := connectWS("rt-lp") + defer conn.Close() + + Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) + initial := drainUntil(conn, "localai.model.failover", 10*time.Second) + Expect(initial).To(HaveKeyWithValue("stage", "llm")) + Expect(initial).To(HaveKeyWithValue("chain", "chain-rt-lp")) + Expect(initial).To(HaveKeyWithValue("reason", "initial")) + Expect(initial).To(HaveKeyWithValue("to", "lp-rt-llm")) + + sendClientEvent(conn, disableVADEvent()) + drainUntil(conn, "session.updated", 10*time.Second) + + // The first turn goes through the proxy to the remote mock-llm. + _, done, failovers := userTurn(conn, "Hello, how are you?") + Expect(failovers).To(BeEmpty()) + resp, _ := done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + Expect(chainActive("chain-rt-lp")).To(Equal("lp-rt-llm")) + + gate.SetDown(true) + _, done, failovers = userTurn(conn, "And now?") + Expect(failovers).To(ContainElement(And( + HaveKeyWithValue("stage", "llm"), + HaveKeyWithValue("from", "lp-rt-llm"), + HaveKeyWithValue("to", "mock-llm"), + HaveKeyWithValue("reason", "trip"), + ))) + resp, _ = done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + Expect(chainActive("chain-rt-lp")).To(Equal("mock-llm")) + }) + It("starts the session on the next target when the active one fails to warm up", func() { conn := connectWS("rt-failover-warm") defer func() { _ = conn.Close() }() From d65b4d3e6a777b123f1a729cc515568f4f8764b6 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sun, 27 Sep 2026 00:50:12 +0000 Subject: [PATCH 51/79] feat(ui): show live failover chain health in the model editor Editing a failover chain now shows its state, the target serving it, and a per-target health table under the editor header. The strip reads GET /api/failover and follows /api/failover/events, with a 15 s re-list to cover SSE reconnect gaps. Admins can pin a target or unpin the chain after a confirmation. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- .../http/react-ui/e2e/failover-health.spec.js | 124 +++++++++++++ .../react-ui/public/locales/de/models.json | 46 +++++ .../react-ui/public/locales/en/models.json | 46 +++++ .../react-ui/public/locales/es/models.json | 46 +++++ .../react-ui/public/locales/id/models.json | 46 +++++ .../react-ui/public/locales/it/models.json | 46 +++++ .../react-ui/public/locales/ko/models.json | 46 +++++ .../react-ui/public/locales/pt-BR/models.json | 46 +++++ .../react-ui/public/locales/zh-CN/models.json | 46 +++++ core/http/react-ui/src/App.css | 36 ++++ .../src/components/FailoverChainStatus.jsx | 163 ++++++++++++++++++ .../react-ui/src/components/StatusPill.jsx | 7 + .../react-ui/src/hooks/useFailoverChains.js | 82 +++++++++ core/http/react-ui/src/pages/ModelEditor.jsx | 4 + core/http/react-ui/src/utils/api.js | 10 ++ core/http/react-ui/src/utils/config.js | 6 + 16 files changed, 800 insertions(+) create mode 100644 core/http/react-ui/e2e/failover-health.spec.js create mode 100644 core/http/react-ui/src/components/FailoverChainStatus.jsx create mode 100644 core/http/react-ui/src/hooks/useFailoverChains.js diff --git a/core/http/react-ui/e2e/failover-health.spec.js b/core/http/react-ui/e2e/failover-health.spec.js new file mode 100644 index 000000000..bc2b3070d --- /dev/null +++ b/core/http/react-ui/e2e/failover-health.spec.js @@ -0,0 +1,124 @@ +import { test, expect } from './coverage-fixtures.js' + +// Live failover chain health in the Model Editor. A model that is a failover +// chain shows a strip under the editor header with the chain state, the +// active target, and a per-target health table fed by GET /api/failover plus +// the /api/failover/events SSE stream. Pin controls are admin-only. + +const MOCK_METADATA = { + sections: [{ id: 'general', label: 'General', icon: 'settings', order: 0 }], + fields: [ + { path: 'name', yaml_key: 'name', go_type: 'string', ui_type: 'string', section: 'general', label: 'Model Name', description: 'id', component: 'input', order: 0 }, + ], +} +const MOCK_YAML = 'name: chain\nfailover:\n targets: [a, b]\n' + +const CHAIN = { + name: 'chain', + state: 'primary', + active: 'a', + active_since: '2026-09-26T09:00:00Z', + pinned: null, + targets: [ + { model: 'a', kind: 'local', warm: true, state: 'healthy', consecutive_ok: 5, last_probe: '2026-09-26T09:59:00Z' }, + { model: 'b', kind: 'remote', warm: false, state: 'healthy', consecutive_ok: 3, last_error: 'dial tcp: connection refused while probing upstream' }, + ], +} + +// The stream replays a snapshot (primary on a) then a switch to b. The body +// ends after the two frames; EventSource reconnects and replays them, so the +// settled state stays fallback on b. +const SSE_BODY = + `event: snapshot\ndata: ${JSON.stringify({ chains: [CHAIN] })}\n\n` + + 'event: chain.switched\ndata: {"type":"chain.switched","chain":"chain","from":"a","to":"b","state":"fallback","reason":"trip","at":"2026-09-26T10:00:00Z"}\n\n' + +async function mockEditor(page, authStatus) { + await page.route('**/api/auth/status', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(authStatus) })) + await page.route('**/api/models/config-metadata*', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK_METADATA) })) + await page.route('**/api/models/edit/**', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: MOCK_YAML, name: 'chain' }) })) + await page.route('**/api/models/config-json/**', (route) => + route.fulfill({ contentType: 'application/json', body: '{}' })) + await page.route('**/api/failover', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [CHAIN] }) })) + await page.route('**/api/failover/chain', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(CHAIN) })) + await page.route('**/api/failover/events', (route) => + route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: SSE_BODY })) +} + +const NO_AUTH = { authEnabled: false, staticApiKeyRequired: false, providers: [] } +const NON_ADMIN = { + authEnabled: true, + staticApiKeyRequired: false, + providers: ['local'], + user: { id: 'user-uuid', name: 'User', role: 'user', provider: 'local' }, +} + +test.describe('Model Editor — failover chain health', () => { + test('shows the chain strip and follows chain.switched to fallback', async ({ page }) => { + await mockEditor(page, NO_AUTH) + await page.goto('/app/model-editor/chain') + + const strip = page.locator('.failover-status') + await expect(strip).toBeVisible({ timeout: 10_000 }) + const chainPill = strip.locator('.failover-status__state .status-pill') + await expect(chainPill).toHaveText(/fallback/i) + await expect(chainPill).toHaveClass(/status-pill--warning/) + await expect(strip.locator('.failover-status__active')).toHaveText('b') + + // Both targets are listed with their health. + const rows = strip.locator('tbody tr') + await expect(rows).toHaveCount(2) + await expect(rows.nth(0)).toContainText('a') + await expect(rows.nth(1).locator('.status-pill')).toHaveClass(/status-pill--success/) + // Long errors are truncated but kept whole in the title. + await expect(rows.nth(1).locator('.failover-status__error')).toHaveAttribute('title', /connection refused/) + }) + + test('admin can pin a target after confirming', async ({ page }) => { + await mockEditor(page, NO_AUTH) + let pinBody = null + await page.route('**/api/failover/chain/pin', async (route) => { + if (route.request().method() === 'POST') pinBody = route.request().postDataJSON() + await route.fulfill({ contentType: 'application/json', body: JSON.stringify({ ...CHAIN, pinned: 'b' }) }) + }) + await page.goto('/app/model-editor/chain') + + const strip = page.locator('.failover-status') + await expect(strip).toBeVisible({ timeout: 10_000 }) + const pinB = strip.locator('tbody tr').nth(1).getByRole('button', { name: /pin/i }) + await expect(pinB).toBeVisible() + await pinB.click() + + const dialog = page.getByRole('alertdialog') + await expect(dialog).toBeVisible() + await dialog.getByRole('button', { name: /^pin$/i }).click() + await expect.poll(() => pinBody).toEqual({ target: 'b' }) + }) + + // The editor route itself is admin-only (RequireAdmin), so a non-admin never + // reaches the strip or its pin controls; the component additionally gates + // pinning on isAdmin for surfaces that are not admin-gated. + test('non-admin users get no pin controls', async ({ page }) => { + await mockEditor(page, NON_ADMIN) + await page.route('**/api/auth/me', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify(NON_ADMIN.user) })) + await page.goto('/app/model-editor/chain') + + await page.waitForURL(/\/app(?!\/model-editor)/, { timeout: 5000 }) + await expect(page.locator('.failover-status')).toHaveCount(0) + await expect(page.getByRole('button', { name: /^pin/i })).toHaveCount(0) + }) + + test('models that are not failover chains show no strip', async ({ page }) => { + await mockEditor(page, NO_AUTH) + await page.route('**/api/models/edit/**', (route) => + route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: 'name: plain\n', name: 'plain' }) })) + await page.goto('/app/model-editor/plain') + await expect(page.locator('h1.page-title')).toBeVisible({ timeout: 10_000 }) + await expect(page.locator('.failover-status')).toHaveCount(0) + }) +}) diff --git a/core/http/react-ui/public/locales/de/models.json b/core/http/react-ui/public/locales/de/models.json index 779e0e7f6..8ac715477 100644 --- a/core/http/react-ui/public/locales/de/models.json +++ b/core/http/react-ui/public/locales/de/models.json @@ -224,5 +224,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/en/models.json b/core/http/react-ui/public/locales/en/models.json index 04e32e4fa..72e841600 100644 --- a/core/http/react-ui/public/locales/en/models.json +++ b/core/http/react-ui/public/locales/en/models.json @@ -241,5 +241,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/es/models.json b/core/http/react-ui/public/locales/es/models.json index d833fba55..ed22edb98 100644 --- a/core/http/react-ui/public/locales/es/models.json +++ b/core/http/react-ui/public/locales/es/models.json @@ -224,5 +224,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/id/models.json b/core/http/react-ui/public/locales/id/models.json index 9088647b9..812fa755f 100644 --- a/core/http/react-ui/public/locales/id/models.json +++ b/core/http/react-ui/public/locales/id/models.json @@ -237,5 +237,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/it/models.json b/core/http/react-ui/public/locales/it/models.json index cbc06d22c..d02010234 100644 --- a/core/http/react-ui/public/locales/it/models.json +++ b/core/http/react-ui/public/locales/it/models.json @@ -224,5 +224,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/ko/models.json b/core/http/react-ui/public/locales/ko/models.json index 8ed38bf0f..20741f534 100644 --- a/core/http/react-ui/public/locales/ko/models.json +++ b/core/http/react-ui/public/locales/ko/models.json @@ -208,5 +208,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/pt-BR/models.json b/core/http/react-ui/public/locales/pt-BR/models.json index c354e89ae..1c4fa1067 100644 --- a/core/http/react-ui/public/locales/pt-BR/models.json +++ b/core/http/react-ui/public/locales/pt-BR/models.json @@ -240,5 +240,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/public/locales/zh-CN/models.json b/core/http/react-ui/public/locales/zh-CN/models.json index 667009901..3a2a0aa9c 100644 --- a/core/http/react-ui/public/locales/zh-CN/models.json +++ b/core/http/react-ui/public/locales/zh-CN/models.json @@ -224,5 +224,51 @@ "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}}" + } } } diff --git a/core/http/react-ui/src/App.css b/core/http/react-ui/src/App.css index c6ea6e71d..45668f96a 100644 --- a/core/http/react-ui/src/App.css +++ b/core/http/react-ui/src/App.css @@ -12354,6 +12354,42 @@ 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; } .me-tabs { display: flex; gap: 0; padding: 0 var(--spacing-lg); border-bottom: 1px solid var(--color-border); } .me-tab { padding: var(--spacing-sm) var(--spacing-md); diff --git a/core/http/react-ui/src/components/FailoverChainStatus.jsx b/core/http/react-ui/src/components/FailoverChainStatus.jsx new file mode 100644 index 000000000..c5afcd944 --- /dev/null +++ b/core/http/react-ui/src/components/FailoverChainStatus.jsx @@ -0,0 +1,163 @@ +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. +function relative(ts, lng) { + const ms = Date.parse(ts) + if (!ms) return null + const seconds = Math.round((ms - Date.now()) / 1000) + const rtf = new Intl.RelativeTimeFormat(lng, { numeric: 'auto' }) + for (const [unit, size] of UNITS) { + if (Math.abs(seconds) >= size) return rtf.format(Math.round(seconds / size), unit) + } + return rtf.format(seconds, 'second') +} + +// FailoverChainStatus renders the live health of one failover chain: its +// state, the target serving it, and a per-target health table. Pin controls +// appear only when canPin, and every pin change needs a confirmation because +// it overrides automatic failover for all callers. +export default function FailoverChainStatus({ chain, onPin, onUnpin, canPin = false }) { + const { t, i18n } = useTranslation('models') + const [confirm, setConfirm] = useState(null) + const [pending, setPending] = useState(false) + + const runConfirmed = async () => { + setPending(true) + try { + if (confirm.kind === 'pin') await onPin?.(confirm.target) + else await onUnpin?.() + } finally { + setPending(false) + setConfirm(null) + } + } + + const since = relative(chain.active_since, i18n.language) + + return ( +
+
+ {t('failover.title')} + + + + + {t('failover.servedBy')} {chain.active} + + {since && ( + + {t('failover.changed', { time: since })} + + )} + {chain.pinned && ( + +