Merge PR #12285: feat(failover): serve a model name from a chain of local and remote targets

Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
Ettore Di Giacinto committed 2026-09-28 15:32:30 +00:00
commit 9c156656bd
186 files changed
+17157 -1085

No files matched your search

+1
View File
@@ -394,6 +394,7 @@ When adding a new endpoint:
- [ ] Error responses use `schema.ErrorResponse` format (or `echo.NewHTTPError` with a mapped gRPC status — see the `mapBackendError` helper in `core/http/endpoints/localai/images.go`)
- [ ] Tests cover both authenticated and unauthenticated access
- [ ] Swagger regenerated (`make swagger`) if you changed any `@Router`/`@Tags`/`@Param` annotation
- [ ] Stateful feature: distributed mode chosen and documented (see [distributed-state.md](distributed-state.md))
## Companion: MCP admin tool surface
+142
View File
@@ -0,0 +1,142 @@
# Distributed-aware state
Every frontend process is a stateless replica: any of them can take a
request, and any of them can be the one that starts, stops, or dies first.
Runtime state that lives only in a Go map on one frontend — an in-memory
cache, a set of pins, a scheduler's next-run time, a background probe loop —
diverges silently the moment there is more than one frontend. Each replica
believes its own copy, clients see different answers depending on which
frontend they hit, and nothing errors: the failure mode is quiet disagreement,
not a crash.
The failover chains feature is the concrete example this rule generalizes
from: chain pins, target health, and the active target of each chain used to
live in a plain map per frontend. Two frontends could disagree about which
target a chain was using, and a pin set through one frontend would be
invisible through another.
## The rule
Before you add runtime state (an in-memory map, a cache, a set of pins, a
scheduler, a background loop, a health prober, anything that is not re-derived
fresh from the database or the request itself), pick one of the four modes
below and write down which one you picked on the feature's own docs page
under `docs/content/`. "I didn't think about it" is not one of the modes.
### 1. Shared
Use `syncstate.SyncedMap` (`core/services/syncstate`) when every frontend
needs to see the same state and any frontend may be the one that reads or
writes it next. Deltas replicate over the message bus to every subscriber,
including the publisher, so all frontends converge without an extra polling
loop.
Add a `Store` (`syncstate.Config.Store`) when the state must survive a
restart of the whole cluster — the store is the durable source of truth, and
new frontends rehydrate from it instead of starting empty. Without a `Store`,
give a `Loader` for local, disk-backed rehydration in standalone mode.
Examples:
- Finetune jobs — `core/services/finetune/service.go`. A `Store` persists jobs
across restarts; without one, a `Loader` reloads from disk.
- Failover pins, targets, and chains — `core/services/failover/distsync`.
`failover.pins` is `Store`-backed (pins must survive a restart);
`failover.targets` and `failover.chains` are ephemeral live-health state
with no `Store`, republished by the leader every 10 s.
**Gotcha:** `Reconcile` with neither a `Store` nor a `Loader` does **nothing**
— there is nothing for it to pull from, so it cannot help a late joiner catch
up. If your map has no durable backing, a leader (or another privileged
writer) must republish its live state after a reconnect instead of relying on
`Reconcile` to recover it.
**Gotcha:** a hydrate (on `Start`, after a NATS reconnect, and on every
`Reconcile` tick) replaces the map's contents **without firing `OnApply`**.
If `OnApply` feeds derived state (the failover manager's pins, a cache, a
running process), that state stays stale after a reconnect or a repaired
missed delta. Re-sync the derived state from the map's `Snapshot()`: on a
periodic tick, and/or from an `OnReconnect` callback registered after the
map's `Start` (callbacks run in registration order, so the map has already
re-hydrated). `failover.Manager.ReconcilePins` is the example.
### 2. Single-runner
Use `advisorylock.RunLeaderLoop` or `advisorylock.TryWithLockCtx`
(`core/services/advisorylock`) when exactly one frontend in the cluster should
run some work on a schedule, and it does not matter which one. Add your lock
key to `keys.go` — keys are a single global namespace across the database, so
pick a number that is not already taken.
`RunLeaderLoop`/`TryWithLockCtx` leadership is **not sticky**: the lock is
acquired and released around each tick, so a different frontend can win it
next time. That is fine for idempotent, tick-shaped work (a cleanup pass, a
periodic scan) where it does not matter who ran the last one.
When leadership must be sticky — the same frontend keeps a role across many
ticks, for example because it holds live connections or in-memory session
state tied to that role — use `advisorylock.HeldLock` instead. It keeps a
dedicated database session and PostgreSQL TCP keepalives so a dead host's
lock is freed in about 30 s rather than after the OS's multi-hour keepalive
default. The failover prober (`core/application/failover_distributed.go`)
uses `HeldLock` for exactly this reason: it needs to stay the same leader
across probe ticks, not re-elect on every one.
Examples:
- Tick-shaped, non-sticky: the node health monitor
(`core/services/nodes/health.go`), via `TryWithLockCtx`.
- Sticky: the failover prober (`core/application/failover_distributed.go`),
via `HeldLock`.
### 3. Stateless per request
Nothing to share: the feature re-derives everything it needs from the
request, the database, or another already-distributed source, so there is no
in-memory state to diverge. Most handlers are this by default — no action
needed beyond noting it if a reviewer might otherwise ask.
### 4. Per-instance (documented exception)
Some state is legitimately local to a frontend — a warm in-process cache
that only saves work when it is warm, a local rate limiter for the process's
own outbound connections. This is allowed, but only when the feature's docs
page says so explicitly and states the consequence (which behavior differs
between frontends, and why that is acceptable). An undocumented per-instance
map is a bug, not a design choice.
## Tests
A shared or single-runner feature includes a **two-instance test** that
proves the two frontends actually agree, not just that each one works alone.
Build it on `testutil.NewFakeBus()` (`core/services/testutil/fakebus.go`):
construct two instances of the feature against the same fake bus, mutate
state through one, and assert the other observes it.
`FakeBus` delivers every publish **synchronously**, including back to the
publisher itself — that is what makes the two-instance test deterministic
without polling, but it is also a trap: **never publish while holding a lock
that the apply path also takes.** The publish call re-enters your own apply
handler on the same goroutine before `Publish` returns, so a lock held across
the publish deadlocks against itself. Release the lock, then publish.
## Checklist
When your PR adds or changes state that lives longer than a single request:
- [ ] Mode chosen: shared (`syncstate.SyncedMap`) / single-runner
(`advisorylock`) / stateless / documented per-instance
- [ ] If shared: `Store` added if the state must survive a cluster restart,
or a `Loader` for standalone rehydration; if neither, a leader
republishes after reconnect instead of relying on bare `Reconcile`;
state derived through `OnApply` re-syncs from `Snapshot()` after a
hydrate (hydrate fires no `OnApply`)
- [ ] If single-runner: new lock key added to `keys.go`; `HeldLock` chosen
over `RunLeaderLoop`/`TryWithLockCtx` if leadership must be sticky
across ticks
- [ ] Publishing to the shared state never happens while holding a lock the
apply path also takes (see the `FakeBus` gotcha above)
- [ ] Two-instance test on `testutil.NewFakeBus()` for shared/single-runner
state
- [ ] The chosen mode (and, for per-instance, the reason) is written on the
feature's docs page under `docs/content/`
- [ ] Standalone mode (no NATS/DB) still behaves correctly — most
shared/single-runner code paths must degrade gracefully alone
+33
View File
@@ -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"
+5
View File
@@ -29,6 +29,7 @@ LocalAI
# Root-level build artifacts when running `go build ./...` against
# Go backend packages whose main lives under backend/go/.
/cloud-proxy
/localai-proxy
/local-store
/valkey-store
# prevent above rules from omitting the helm chart
@@ -48,6 +49,10 @@ tests/e2e-aio/backends
# tests/e2e/mock-backend/.gitignore covers the same binary; kept here too so
# the artifact stays ignored if that scoped file is ever removed.
/tests/e2e/mock-backend/mock-backend
# The cloud-proxy backend binary the e2e suite runs next to the mock backend.
/tests/e2e/mock-backend/cloud-proxy
# The localai-proxy backend binary, built next to it the same way.
/tests/e2e/mock-backend/localai-proxy
release/
+2
View File
@@ -33,6 +33,7 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants]
| [.agents/localai-assistant-mcp.md](.agents/localai-assistant-mcp.md) | LocalAI Assistant chat modality — adding admin tools to the in-process MCP server, editing skill prompts, keeping REST + MCP + skills in sync |
| [.agents/backend-signing.md](.agents/backend-signing.md) | Backend OCI image signing (keyless cosign + sigstore-go) — producer-side CI setup, consumer-side gallery `verification:` block, strict mode (`LOCALAI_REQUIRE_BACKEND_INTEGRITY`), revocation via `not_before` |
| [.agents/preparing-a-release.md](.agents/preparing-a-release.md) | Cutting a release: PR labels, `RELEASE_NOTES_vX.Y.Z.md`, the blog post under `website/content/blog/`, and the demo clips under `website/static/media/` |
| [.agents/distributed-state.md](.agents/distributed-state.md) | Features that keep runtime state — how they must behave with several frontends (syncstate, advisory-lock leaders, fakebus tests) |
| [.impeccable.md](.impeccable.md) | Design context for UI/UX work — users, brand personality, aesthetic direction, and design principles |
## Quick Reference
@@ -49,3 +50,4 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants]
- **Backend OS coverage**: a new backend must target every OS it can build for, not just Linux. `.github/backend-matrix.yml` has two matrices — `include:` (Linux) and `includeDarwin:` (macOS / Apple Silicon). Most C/C++/GGML and many Python backends build on Darwin too — wire the `includeDarwin` entry + `backend/index.yaml` `metal:` entries, or say in the PR why an OS is unsupported. See the darwin checklist in [.agents/adding-backends.md](.agents/adding-backends.md).
- **Gallery variant ranking**: a gallery entry can declare `variants` (alternative builds of the same weights), and LocalAI ranks the ones a host can run by engine preference first, size second. A new backend that should be preferred on some hardware must be listed in `engineNamePreferenceRules` in `pkg/system/capabilities.go`; the sibling `backendBuildTagPreferenceRules` speaks build tags rather than engine names, and using the wrong table matches nothing without erroring. See [.agents/adding-backends.md](.agents/adding-backends.md).
- **UI**: The active UI is the React app in `core/http/react-ui/`. The older Alpine.js/HTML UI in `core/http/static/` is pending deprecation — all new UI work goes in the React UI
- **Distributed-aware state**: any feature that keeps runtime state (maps, caches, pins, schedulers, background loops) must choose shared (syncstate), single-runner (advisorylock), stateless, or documented per-instance behaviour for multi-frontend clusters. See [.agents/distributed-state.md](.agents/distributed-state.md).
+13 -4
View File
@@ -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
########################################################
+13
View File
@@ -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
+453
View File
@@ -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)
}
+489
View File
@@ -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 }
+376
View File
@@ -0,0 +1,376 @@
package main
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"github.com/mudler/xlog"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
// maxErrorBody caps the upstream body quoted in an error. It keeps gRPC
// status messages small while leaving room for LocalAI's JSON error text,
// which failover scans for request errors such as context overflows.
const maxErrorBody = 500
// postJSON sends body as JSON to path and decodes a 2xx JSON reply into out
// (skipped when out is nil). The request_timeout_seconds limit applies.
func (p *LocalAIProxy) postJSON(ctx context.Context, path string, body, out any) error {
cfg, err := p.config()
if err != nil {
return err
}
payload, err := json.Marshal(body)
if err != nil {
return status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
}
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
return p.do(req, path, out)
}
// postMultipart sends fields and, when fileField is set, the local file at
// filePath as a multipart form to path, decoding a 2xx JSON reply into out
// (skipped when out is nil). Core hands audio and images to backends as local
// paths, and LocalAI's upload endpoints take them as multipart files. The
// request_timeout_seconds limit applies.
func (p *LocalAIProxy) postMultipart(ctx context.Context, path string, fields map[string]string, fileField, filePath string, out any) error {
form := multipartForm{fields: url.Values{}}
for k, v := range fields {
form.fields.Set(k, v)
}
if fileField != "" {
form.files = []formFile{{field: fileField, path: filePath}}
}
return p.postForm(ctx, path, form, out)
}
// postForm uploads form to path and decodes the 2xx JSON reply into out. The
// request_timeout_seconds limit applies.
func (p *LocalAIProxy) postForm(ctx context.Context, path string, form multipartForm, out any) error {
cfg, err := p.config()
if err != nil {
return err
}
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
req, err := p.newMultipartRequest(ctx, cfg, path, form)
if err != nil {
return err
}
return p.do(req, path, out)
}
// multipartForm is an upload: repeated fields (timestamp_granularities[]) need
// url.Values, and audio transforms send two files.
type multipartForm struct {
fields url.Values
files []formFile
}
type formFile struct {
field string
path string
}
// newMultipartRequest builds a POST whose body streams form through a pipe,
// so large audio files are not buffered in memory. The writer goroutine owns
// the opened files and closes them when it finishes; the transport closes the
// pipe when the request ends, which unblocks the writer on every error path.
func (p *LocalAIProxy) newMultipartRequest(ctx context.Context, cfg *proxyConfig, path string, form multipartForm) (*http.Request, error) {
// Open before contacting the upstream so a bad path is reported as a
// request error, not as a failure of the remote host.
files := make([]*os.File, 0, len(form.files))
closeAll := func() {
for _, f := range files {
_ = f.Close()
}
}
for _, ff := range form.files {
f, err := os.Open(ff.path)
if err != nil {
closeAll()
return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: open %s: %v", ff.path, err)
}
files = append(files, f)
}
pr, pw := io.Pipe()
mw := multipart.NewWriter(pw)
go func() {
defer closeAll()
pw.CloseWithError(writeMultipart(mw, form, files))
}()
req, err := p.newRequest(ctx, cfg, http.MethodPost, path, pr)
if err != nil {
_ = pr.CloseWithError(err)
return nil, err
}
req.Header.Set("Content-Type", mw.FormDataContentType())
return req, nil
}
func writeMultipart(mw *multipart.Writer, form multipartForm, files []*os.File) error {
for k, vs := range form.fields {
for _, v := range vs {
if err := mw.WriteField(k, v); err != nil {
return err
}
}
}
for i, file := range files {
part, err := mw.CreateFormFile(form.files[i].field, filepath.Base(file.Name()))
if err != nil {
return err
}
if _, err := io.Copy(part, file); err != nil {
return err
}
}
return mw.Close()
}
// postStream sends body as JSON to path and returns the open response of a
// 2xx reply; the caller must close its body. No request_timeout_seconds
// limit applies: streams legitimately outlast it, and ctx bounds them.
func (p *LocalAIProxy) postStream(ctx context.Context, path string, body any) (*http.Response, error) {
cfg, err := p.config()
if err != nil {
return nil, err
}
payload, err := json.Marshal(body)
if err != nil {
return nil, status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
}
req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
return p.doStream(req, path)
}
// postMultipartStream uploads form and returns the open response of a 2xx
// reply; the caller must close its body. Like postStream, only ctx bounds it.
func (p *LocalAIProxy) postMultipartStream(ctx context.Context, path string, form multipartForm) (*http.Response, error) {
cfg, err := p.config()
if err != nil {
return nil, err
}
req, err := p.newMultipartRequest(ctx, cfg, path, form)
if err != nil {
return nil, err
}
return p.doStream(req, path)
}
// doStream runs req and returns the open response of a 2xx reply.
func (p *LocalAIProxy) doStream(req *http.Request, path string) (*http.Response, error) {
resp, err := p.client.Do(req)
if err != nil {
return nil, transportError(path, err)
}
if resp.StatusCode < 200 || resp.StatusCode > 299 {
defer func() { _ = resp.Body.Close() }()
return nil, statusError(path, resp)
}
return resp, nil
}
func (p *LocalAIProxy) newRequest(ctx context.Context, cfg *proxyConfig, method, path string, body io.Reader) (*http.Request, error) {
req, err := http.NewRequestWithContext(ctx, method, cfg.base+path, body)
if err != nil {
return nil, status.Errorf(codes.Internal, "localai-proxy: build %s request: %v", path, err)
}
if cfg.apiKey != "" {
req.Header.Set("Authorization", "Bearer "+cfg.apiKey)
}
return req, nil
}
// do runs req and decodes a 2xx JSON reply into out.
func (p *LocalAIProxy) do(req *http.Request, path string, out any) error {
resp, err := p.client.Do(req)
if err != nil {
return transportError(path, err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return statusError(path, resp)
}
if out == nil {
_, _ = io.Copy(io.Discard, resp.Body)
return nil
}
if err := json.NewDecoder(resp.Body).Decode(out); err != nil {
if ctxErr := req.Context().Err(); ctxErr != nil {
return transportError(path, ctxErr)
}
return status.Errorf(codes.Internal, "localai-proxy: decode %s response: %v", path, err)
}
return nil
}
func withTimeout(ctx context.Context, cfg *proxyConfig) (context.Context, context.CancelFunc) {
if cfg.timeout > 0 {
return context.WithTimeout(ctx, cfg.timeout)
}
return context.WithCancel(ctx)
}
// transportError maps a failed round trip to a gRPC status. A dead or
// unreachable upstream is Unavailable so failover moves to the next target.
func transportError(path string, err error) error {
code := codes.Unavailable
switch {
case errors.Is(err, context.DeadlineExceeded):
code = codes.DeadlineExceeded
case errors.Is(err, context.Canceled):
code = codes.Canceled
}
xlog.Warn("localai-proxy: upstream request failed", "path", path, "error", err)
return status.Errorf(code, "localai-proxy: upstream %s: %v", path, err)
}
// statusError maps a non-2xx upstream reply to a gRPC status. 5xx means the
// upstream is unhealthy (Unavailable, so failover retries elsewhere); 4xx
// means the request itself is wrong (InvalidArgument, so failover does not
// trip a healthy target over a client error). 429 is the exception, see
// below. 501 is the upstream saying it cannot serve this kind of request,
// which failover treats as a capability gap, like our own Unimplemented
// methods.
func statusError(path string, resp *http.Response) error {
raw, _ := io.ReadAll(io.LimitReader(resp.Body, maxErrorBody+1))
msg := strings.TrimSpace(string(raw))
if len(msg) > maxErrorBody {
msg = msg[:maxErrorBody] + "..."
}
// gRPC refuses to send a status message that is not valid UTF-8, and the
// cut above may split a rune.
msg = strings.ToValidUTF8(msg, "")
var code codes.Code
switch {
case resp.StatusCode == http.StatusNotImplemented:
code = codes.Unimplemented
case resp.StatusCode >= 500:
code = codes.Unavailable
case resp.StatusCode == http.StatusTooManyRequests:
// Rate limited: the request is fine but the upstream is out of
// capacity. Failover retries it on the next target and trips this
// one, so traffic moves off it for a while.
code = codes.ResourceExhausted
case resp.StatusCode >= 400:
code = codes.InvalidArgument
default:
// A 1xx/3xx here means a misbehaving upstream (redirects are refused
// by the client), not a bad request.
code = codes.Unavailable
}
xlog.Warn("localai-proxy: upstream error", "path", path, "status", resp.StatusCode)
return status.Error(code, fmt.Sprintf("localai-proxy: upstream %s returned %d: %s", path, resp.StatusCode, msg))
}
// postJSONToFile sends body as JSON to path and writes a 2xx reply's body,
// which is audio rather than JSON, to dst. The request_timeout_seconds limit
// applies.
func (p *LocalAIProxy) postJSONToFile(ctx context.Context, path string, body any, dst string) error {
cfg, err := p.config()
if err != nil {
return err
}
payload, err := json.Marshal(body)
if err != nil {
return status.Errorf(codes.InvalidArgument, "localai-proxy: encode %s request: %v", path, err)
}
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
req, err := p.newRequest(ctx, cfg, http.MethodPost, path, bytes.NewReader(payload))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
_, err = p.doToFile(req, path, dst)
return err
}
// postMultipartToFile uploads form to path and writes a 2xx reply's body to
// dst, returning the reply headers for endpoints that describe extra outputs
// there. The request_timeout_seconds limit applies.
func (p *LocalAIProxy) postMultipartToFile(ctx context.Context, path string, form multipartForm, dst string) (http.Header, error) {
cfg, err := p.config()
if err != nil {
return nil, err
}
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
req, err := p.newMultipartRequest(ctx, cfg, path, form)
if err != nil {
return nil, err
}
return p.doToFile(req, path, dst)
}
// getToFile downloads path to dst. The request_timeout_seconds limit applies.
func (p *LocalAIProxy) getToFile(ctx context.Context, path, dst string) error {
cfg, err := p.config()
if err != nil {
return err
}
ctx, cancel := withTimeout(ctx, cfg)
defer cancel()
req, err := p.newRequest(ctx, cfg, http.MethodGet, path, nil)
if err != nil {
return err
}
_, err = p.doToFile(req, path, dst)
return err
}
// doToFile runs req and writes a 2xx body to dst. A failed copy removes dst:
// core serves whatever file it finds there, and a truncated recording must
// not pass for a finished one.
func (p *LocalAIProxy) doToFile(req *http.Request, path, dst string) (http.Header, error) {
resp, err := p.doStream(req, path)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
// #nosec G304 -- dst is the output path core chose for this call (generated content dir), never a caller-supplied path
f, err := os.Create(filepath.Clean(dst))
if err != nil {
return nil, status.Errorf(codes.Internal, "localai-proxy: create %s: %v", dst, err)
}
_, copyErr := io.Copy(f, resp.Body)
closeErr := f.Close()
if copyErr != nil || closeErr != nil {
_ = os.Remove(dst)
if copyErr != nil {
if ctxErr := req.Context().Err(); ctxErr != nil {
return nil, transportError(path, ctxErr)
}
return nil, transportError(path, copyErr)
}
return nil, status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, closeErr)
}
return resp.Header, nil
}
@@ -0,0 +1,164 @@
package main
import (
"encoding/json"
"io"
"mime"
"mime/multipart"
"net/http"
"net/http/httptest"
"strings"
"sync"
. "github.com/onsi/gomega"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
// recordedRequest is what the fake upstream saw for one call. JSON bodies land
// in JSON; multipart bodies land in Fields and Files (field name to content).
type recordedRequest struct {
Method string
Path string
Auth string
JSON map[string]any
Fields map[string]string
Files map[string]string
}
// scriptedResponse is the reply for one path. SSE, when set, is written as
// "data: <frame>" events and wins over Body. Header adds response headers.
type scriptedResponse struct {
Status int
ContentType string
Header map[string]string
Body string
SSE []string
}
// fakeUpstream stands in for a remote LocalAI: it records every request and
// answers each path with the response scripted for it (404 otherwise).
type fakeUpstream struct {
*httptest.Server
mu sync.Mutex
requests []recordedRequest
responses map[string]scriptedResponse
}
func newFakeUpstream() *fakeUpstream {
f := &fakeUpstream{responses: map[string]scriptedResponse{}}
f.Server = httptest.NewServer(http.HandlerFunc(f.serve))
return f
}
// newFakeUpstreamWithHandler serves every request with h instead of the
// scripted responses, for tests that need to control timing.
func newFakeUpstreamWithHandler(h http.HandlerFunc) *fakeUpstream {
return &fakeUpstream{Server: httptest.NewServer(h), responses: map[string]scriptedResponse{}}
}
func (f *fakeUpstream) script(path string, r scriptedResponse) {
f.mu.Lock()
defer f.mu.Unlock()
f.responses[path] = r
}
// replyJSON scripts a 200 JSON response for path.
func (f *fakeUpstream) replyJSON(path string, body any) {
raw, err := json.Marshal(body)
Expect(err).NotTo(HaveOccurred())
f.script(path, scriptedResponse{Status: http.StatusOK, ContentType: "application/json", Body: string(raw)})
}
func (f *fakeUpstream) recorded() []recordedRequest {
f.mu.Lock()
defer f.mu.Unlock()
return append([]recordedRequest(nil), f.requests...)
}
// last returns the single most recent request, failing when none arrived.
func (f *fakeUpstream) last() recordedRequest {
reqs := f.recorded()
ExpectWithOffset(1, reqs).NotTo(BeEmpty(), "upstream received no request")
return reqs[len(reqs)-1]
}
func (f *fakeUpstream) serve(w http.ResponseWriter, r *http.Request) {
rec := recordedRequest{Method: r.Method, Path: r.URL.Path, Auth: r.Header.Get("Authorization")}
mediaType, params, _ := mime.ParseMediaType(r.Header.Get("Content-Type"))
switch {
case mediaType == "multipart/form-data":
rec.Fields, rec.Files = map[string]string{}, map[string]string{}
mr := multipart.NewReader(r.Body, params["boundary"])
for {
part, err := mr.NextPart()
if err != nil {
break
}
data, _ := io.ReadAll(part)
if part.FileName() != "" {
rec.Files[part.FormName()] = string(data)
} else {
rec.Fields[part.FormName()] = string(data)
}
}
default:
raw, _ := io.ReadAll(r.Body)
if len(raw) > 0 {
_ = json.Unmarshal(raw, &rec.JSON)
}
}
f.mu.Lock()
f.requests = append(f.requests, rec)
resp, ok := f.responses[r.URL.Path]
f.mu.Unlock()
if !ok {
http.Error(w, "no scripted response for "+r.URL.Path, http.StatusNotFound)
return
}
if resp.SSE != nil {
w.Header().Set("Content-Type", "text/event-stream")
w.WriteHeader(http.StatusOK)
flusher, _ := w.(http.Flusher)
for _, frame := range resp.SSE {
_, _ = io.WriteString(w, "data: "+frame+"\n\n")
if flusher != nil {
flusher.Flush()
}
}
return
}
if resp.ContentType != "" {
w.Header().Set("Content-Type", resp.ContentType)
}
for k, v := range resp.Header {
w.Header().Set(k, v)
}
w.WriteHeader(resp.Status)
_, _ = io.WriteString(w, resp.Body)
}
// loadProxy returns a proxy loaded against the fake upstream with the given
// proxy options merged over sane defaults.
func loadProxy(f *fakeUpstream, mutate func(*pb.ModelOptions)) *LocalAIProxy {
opts := &pb.ModelOptions{
Model: "local-name",
Proxy: &pb.ProxyOptions{UpstreamUrl: f.URL + "/", UpstreamModel: "remote-model"},
}
if mutate != nil {
mutate(opts)
}
p := NewLocalAIProxy()
ExpectWithOffset(1, p.Load(opts)).To(Succeed())
return p
}
// sseJSON marshals v for use as one SSE frame.
func sseJSON(v any) string {
raw, err := json.Marshal(v)
Expect(err).NotTo(HaveOccurred())
return strings.TrimSpace(string(raw))
}
+479
View File
@@ -0,0 +1,479 @@
package main
import (
"encoding/base64"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"math"
"net/http"
"net/url"
"strings"
"time"
"github.com/gorilla/websocket"
"github.com/mudler/xlog"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
const (
realtimePath = "/v1/realtime"
// defaultLiveSampleRate is what TranscriptLiveConfig.sample_rate 0 means.
defaultLiveSampleRate = 16000
// finalWait bounds how long a closing session waits for an utterance the
// upstream still has in flight. The upstream transcribes only after its
// VAD sees the speech end, so a slow model can finish after the client
// stops sending; waiting forever would pin the gRPC call on a hung
// upstream.
finalWait = 5 * time.Second
// commitGrace is how long a speech_stopped waits for its
// input_audio_buffer.committed. The upstream sends the two back to back,
// so a stop with no commit inside this window is a discarded turn and
// must not hold a closing session for the full finalWait.
commitGrace = 500 * time.Millisecond
// maxBacklogSeconds caps the audio held before the ready ack. Callers
// wait for the ack before streaming, so more than this means a client
// that ignores the contract, and holding it unbounded while a cold
// upstream loads models would grow memory without limit.
maxBacklogSeconds = 5
// liveWriteTimeout turns an upstream that stops reading into an error
// instead of a call blocked on a full socket buffer.
liveWriteTimeout = 10 * time.Second
// liveHandshakeTimeout bounds the WebSocket upgrade only. The upstream
// warms the pipeline models before it sends session.created, which can
// take minutes on a cold box, so the setup phase after the upgrade has
// its own, longer bound.
liveHandshakeTimeout = 30 * time.Second
// defaultLiveSetupTimeout bounds the session setup (upgrade to
// session.updated) when request_timeout_seconds is unset. Core waits for
// the ready ack with a plain Recv, so without a bound a hung upstream
// would hold the call, and the failover that should move the stage to
// the next target, forever. It is generous because a cold upstream loads
// the pipeline's VAD and transcription models before it answers.
defaultLiveSetupTimeout = 3 * time.Minute
)
// liveSetupTimeout is defaultLiveSetupTimeout, as a variable so tests can
// exercise the bound without waiting minutes.
var liveSetupTimeout = defaultLiveSetupTimeout
// realtimeEvent holds the fields the bridge reads from any upstream server
// event; unrelated events decode into it harmlessly.
type realtimeEvent struct {
Type string `json:"type"`
// ItemID is set on committed, delta, completed and failed events; the
// upstream uses the committed turn's id for its transcription events.
ItemID string `json:"item_id"`
Delta string `json:"delta"`
Transcript string `json:"transcript"`
Error *struct {
Message string `json:"message"`
Code string `json:"code"`
} `json:"error"`
}
// AudioTranscriptionLive bridges a live transcription session to the
// upstream's /v1/realtime transcription session, so a realtime pipeline on
// this box can use a remote LocalAI's streaming ASR. The upstream runs its own
// server VAD and transcribes per utterance: each completed utterance becomes
// Delta (whatever the streamed deltas did not already carry) plus Eou. Word
// timings and Eob have no upstream counterpart and stay empty.
//
// Contract (pkg/grpc/server.go): this method closes out, in closes when the
// client half-closes, and errors are returned at once because the caller
// blocks on the first Recv for the ready ack.
func (p *LocalAIProxy) AudioTranscriptionLive(in <-chan *pb.TranscriptLiveRequest, out chan<- *pb.TranscriptLiveResponse) error {
defer close(out)
cfg, err := p.config()
if err != nil {
return err
}
if cfg.realtimePipeline == "" {
return grpcerrors.LiveTranscriptionUnsupported(backendName, "set the realtime_pipeline backend option")
}
first, ok := <-in
if !ok {
return nil // the caller closed without sending anything
}
lc := first.GetConfig()
if lc == nil {
return status.Error(codes.InvalidArgument, "localai-proxy: the first live transcription message must carry a config")
}
rate := int(lc.GetSampleRate())
if rate == 0 {
rate = defaultLiveSampleRate
}
conn, err := p.dialRealtime(cfg)
if err != nil {
return err
}
defer func() { _ = conn.Close() }()
s := &liveSession{
conn: conn,
out: out,
pipeline: cfg.realtimePipeline,
language: lc.GetLanguage(),
rate: rate,
sent: map[string]string{},
committed: map[string]bool{},
events: make(chan realtimeEvent),
done: make(chan struct{}),
}
// Closing done releases the reader if it is blocked handing over an
// event; closing conn (deferred above, runs after this) unblocks its read.
defer close(s.done)
go s.readLoop()
setup := cfg.timeout
if setup <= 0 {
setup = liveSetupTimeout
}
return s.run(in, setup)
}
// dialRealtime opens the upstream WebSocket. A refused upgrade is mapped like
// any other upstream HTTP reply, so a 5xx trips failover and a 4xx does not.
func (p *LocalAIProxy) dialRealtime(cfg *proxyConfig) (*websocket.Conn, error) {
u, err := url.Parse(cfg.base + realtimePath)
if err != nil {
return nil, status.Errorf(codes.Internal, "localai-proxy: build %s URL: %v", realtimePath, err)
}
// Load only accepts http(s) bases.
if u.Scheme == "https" {
u.Scheme = "wss"
} else {
u.Scheme = "ws"
}
u.RawQuery = url.Values{"model": {cfg.realtimePipeline}}.Encode()
header := http.Header{}
if cfg.apiKey != "" {
header.Set("Authorization", "Bearer "+cfg.apiKey)
}
dialer := websocket.Dialer{
Proxy: http.ProxyFromEnvironment,
HandshakeTimeout: liveHandshakeTimeout,
}
conn, resp, err := dialer.Dial(u.String(), header)
if err != nil {
if resp != nil && resp.StatusCode != http.StatusSwitchingProtocols {
defer func() { _ = resp.Body.Close() }()
return nil, statusError(realtimePath, resp)
}
return nil, transportError(realtimePath, err)
}
return conn, nil
}
// liveSession is the state of one bridged session. Only run touches it; the
// reader goroutine only hands events over.
type liveSession struct {
conn *websocket.Conn
out chan<- *pb.TranscriptLiveResponse
pipeline string
language string
rate int
sent map[string]string // text already sent as deltas, per upstream item
final []string // completed transcripts, in order
// In-flight tracking, so a closing session waits only for an utterance
// the upstream will still transcribe. The upstream VAD emits
// speech_started, then either speech_stopped plus committed (a turn it
// will transcribe) or nothing at all (a turn it discarded as no speech).
speaking bool // between speech_started and speech_stopped
stopping bool // speech_stopped seen, its committed not yet
committed map[string]bool // committed items not yet completed
events chan realtimeEvent
readErr error // set before events is closed
done chan struct{}
}
// readLoop decodes upstream events until the socket fails. It is the only
// reader, as gorilla/websocket requires.
func (s *liveSession) readLoop() {
defer close(s.events)
for {
_, msg, err := s.conn.ReadMessage()
if err != nil {
s.readErr = err
return
}
var ev realtimeEvent
if err := json.Unmarshal(msg, &ev); err != nil {
xlog.Debug("localai-proxy: skipping undecodable realtime event", "error", err)
continue
}
select {
case s.events <- ev:
case <-s.done:
return
}
}
}
// run drives the session: setup, then audio forwarding and event mapping,
// then the drain once the client closes its side. It is also the only writer
// on the socket, so frames never interleave.
func (s *liveSession) run(in <-chan *pb.TranscriptLiveRequest, setupTimeout time.Duration) error {
var (
created, ready bool
inClosed bool
backlog []*pb.TranscriptLiveRequest
backlogSamples int
drainTimer <-chan time.Time
graceTimer <-chan time.Time
)
setup := time.NewTimer(setupTimeout)
defer setup.Stop()
setupTimer := setup.C
drain := time.NewTimer(finalWait)
drain.Stop()
defer drain.Stop()
grace := time.NewTimer(commitGrace)
grace.Stop()
defer grace.Stop()
for {
select {
case req, ok := <-in:
if !ok {
in, inClosed = nil, true
if !ready {
// The caller gave up waiting for the ready ack (core
// closes its send side on failure or cancel). Canceled
// rather than nil: there is no session to report as
// complete, and failover neither retries nor trips a
// target on a canceled call.
return status.Error(codes.Canceled, "localai-proxy: live transcription closed before the upstream session was ready")
}
if s.inFlight() {
drain.Reset(finalWait)
drainTimer = drain.C
}
} else if !ready {
// Hold audio sent before the ready ack rather than drop it,
// up to a bound.
backlogSamples += len(req.GetAudio().GetPcm())
if backlogSamples > maxBacklogSeconds*s.rate {
return status.Errorf(codes.InvalidArgument,
"localai-proxy: more than %d s of audio sent before the live transcription session was ready", maxBacklogSeconds)
}
backlog = append(backlog, req)
} else if err := s.forward(req); err != nil {
return err
}
case ev, ok := <-s.events:
if !ok {
return s.upstreamGone()
}
switch ev.Type {
case "session.created":
if created {
continue
}
created = true
if err := s.write(s.sessionUpdate()); err != nil {
return err
}
case "session.updated":
if ready || !created {
continue
}
ready = true
setupTimer = nil
s.out <- &pb.TranscriptLiveResponse{Ready: true}
for _, req := range backlog {
if err := s.forward(req); err != nil {
return err
}
}
backlog = nil
case "error":
return s.upstreamError("error", ev)
case "conversation.item.input_audio_transcription.failed":
return s.upstreamError("transcription failed", ev)
case "input_audio_buffer.speech_started":
s.speaking, s.stopping = true, false
case "input_audio_buffer.speech_stopped":
if s.speaking {
s.speaking, s.stopping = false, true
grace.Reset(commitGrace)
graceTimer = grace.C
}
case "input_audio_buffer.committed":
s.stopping, graceTimer = false, nil
if ev.ItemID != "" {
s.committed[ev.ItemID] = true
}
case "conversation.item.input_audio_transcription.delta":
s.delta(ev)
case "conversation.item.input_audio_transcription.completed":
s.completed(ev)
}
case <-graceTimer:
// A stop the upstream never committed: the turn was discarded.
s.stopping, graceTimer = false, nil
case <-setupTimer:
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s did not set up the transcription session within %s", realtimePath, setupTimeout)
case <-drainTimer:
xlog.Warn("localai-proxy: upstream did not finish the last utterance in time; finalizing without it",
"pipeline", s.pipeline, "wait", finalWait)
return s.finish()
}
if inClosed && !s.inFlight() {
return s.finish()
}
}
}
// inFlight reports an utterance the upstream is still expected to
// transcribe.
func (s *liveSession) inFlight() bool {
return s.speaking || s.stopping || len(s.committed) > 0
}
func (s *liveSession) sessionUpdate() map[string]any {
return map[string]any{
"type": "session.update",
"session": map[string]any{
"type": "transcription",
"audio": map[string]any{
"input": map[string]any{
"format": map[string]any{"type": "audio/pcm", "rate": s.rate},
"transcription": map[string]any{"model": s.pipeline, "language": s.language},
"turn_detection": map[string]any{"type": "server_vad"},
},
},
},
}
}
// forward sends one client message upstream. The upstream session rate is
// fixed at setup, so a second config cannot be honoured.
func (s *liveSession) forward(req *pb.TranscriptLiveRequest) error {
if req.GetConfig() != nil {
return status.Error(codes.InvalidArgument, "localai-proxy: a live transcription config is only accepted as the first message")
}
pcm := req.GetAudio().GetPcm()
if len(pcm) == 0 {
return nil
}
return s.write(map[string]any{
"type": "input_audio_buffer.append",
"audio": base64.StdEncoding.EncodeToString(pcm16LE(pcm)),
})
}
// pcm16LE converts float samples in [-1, 1] to the little-endian PCM16 the
// realtime API takes. Out-of-range samples are clipped rather than wrapped.
func pcm16LE(pcm []float32) []byte {
buf := make([]byte, len(pcm)*2)
for i, f := range pcm {
v := float64(f)
if math.IsNaN(v) {
// A NaN would convert to an arbitrary int16; silence is the
// only neutral value.
v = 0
}
v = math.Max(-1, math.Min(1, v))
// #nosec G115 -- two's-complement reinterpretation for little-endian PCM16 encoding, value range already clamped
binary.LittleEndian.PutUint16(buf[i*2:], uint16(int16(v*math.MaxInt16)))
}
return buf
}
func (s *liveSession) delta(ev realtimeEvent) {
if ev.Delta == "" {
return
}
s.sent[ev.ItemID] += ev.Delta
s.out <- &pb.TranscriptLiveResponse{Delta: ev.Delta}
}
// completed ends one utterance. Deltas only arrive when the upstream pipeline
// streams its transcription, so the completion carries whatever text the
// deltas did not; if the final transcript diverged from them there is no way
// to retract, and FinalResult carries the authoritative text.
func (s *liveSession) completed(ev realtimeEvent) {
sent := s.sent[ev.ItemID]
delete(s.sent, ev.ItemID)
rest, ok := strings.CutPrefix(ev.Transcript, sent)
if !ok {
rest = ""
}
delete(s.committed, ev.ItemID)
if t := strings.TrimSpace(ev.Transcript); t != "" {
s.final = append(s.final, t)
}
s.out <- &pb.TranscriptLiveResponse{Delta: rest, Eou: true}
}
// finish sends the final transcript and closes the upstream session cleanly.
func (s *liveSession) finish() error {
s.out <- &pb.TranscriptLiveResponse{FinalResult: &pb.TranscriptResult{Text: strings.Join(s.final, " ")}}
_ = s.conn.WriteControl(websocket.CloseMessage,
websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""), time.Now().Add(time.Second))
return nil
}
func (s *liveSession) write(v any) error {
if err := s.conn.SetWriteDeadline(time.Now().Add(liveWriteTimeout)); err != nil {
return s.writeFailed(err)
}
if err := s.conn.WriteJSON(v); err != nil {
return s.writeFailed(err)
}
return nil
}
func (s *liveSession) writeFailed(err error) error {
xlog.Warn("localai-proxy: realtime write failed", "pipeline", s.pipeline, "error", err)
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s: %v", realtimePath, err)
}
// upstreamGone maps a socket that stopped delivering events. It is
// Unavailable even for a clean close: the client did not end the session, so
// the upstream dropped it, and failover should reopen on the next target.
func (s *liveSession) upstreamGone() error {
err := s.readErr
if err == nil {
err = errors.New("connection closed")
}
xlog.Warn("localai-proxy: realtime upstream disconnected", "pipeline", s.pipeline, "error", err)
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s disconnected: %v", realtimePath, err)
}
func (s *liveSession) upstreamError(what string, ev realtimeEvent) error {
msg := "no details"
if ev.Error != nil && ev.Error.Message != "" {
msg = ev.Error.Message
if ev.Error.Code != "" {
msg = fmt.Sprintf("%s (%s)", msg, ev.Error.Code)
}
}
xlog.Warn("localai-proxy: realtime upstream reported an error", "pipeline", s.pipeline, "event", what, "error", msg)
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s %s: %s", realtimePath, what, msg)
}
+533
View File
@@ -0,0 +1,533 @@
package main
import (
"encoding/base64"
"encoding/binary"
"encoding/json"
"math"
"net/http"
"runtime"
"strings"
"time"
"github.com/gorilla/websocket"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
// wsUpstream is a fake upstream /v1/realtime endpoint. Each spec scripts the
// server side of the session in script, which runs on the upgraded socket.
func wsUpstream(script func(c *websocket.Conn, r *http.Request)) *fakeUpstream {
upgrader := websocket.Upgrader{}
return newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/v1/realtime" {
http.NotFound(w, r)
return
}
c, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
defer func() { _ = c.Close() }()
script(c, r)
})
}
func wsSend(c *websocket.Conn, v any) {
defer GinkgoRecover()
Expect(c.WriteJSON(v)).To(Succeed())
}
// wsRecv reads the next client event; ok is false once the client is gone.
func wsRecv(c *websocket.Conn) (map[string]any, bool) {
var ev map[string]any
if err := c.ReadJSON(&ev); err != nil {
return nil, false
}
return ev, true
}
// wsHandshake plays the upstream's session setup and returns the client's
// session.update event.
func wsHandshake(c *websocket.Conn) map[string]any {
defer GinkgoRecover()
wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}})
upd, ok := wsRecv(c)
Expect(ok).To(BeTrue())
Expect(upd["type"]).To(Equal("session.update"))
wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}})
return upd
}
// wsDrain reads client events until the client closes, returning the
// decoded PCM16 samples of every input_audio_buffer.append.
func wsDrain(c *websocket.Conn) [][]int16 {
var frames [][]int16
for {
ev, ok := wsRecv(c)
if !ok {
return frames
}
if ev["type"] != "input_audio_buffer.append" {
continue
}
raw, err := base64.StdEncoding.DecodeString(ev["audio"].(string))
if err != nil {
return frames
}
samples := make([]int16, len(raw)/2)
for i := range samples {
samples[i] = int16(binary.LittleEndian.Uint16(raw[i*2:]))
}
frames = append(frames, samples)
}
}
type liveCall struct {
in chan *pb.TranscriptLiveRequest
out chan *pb.TranscriptLiveResponse
errc chan error
}
// startLive runs AudioTranscriptionLive the way pkg/grpc/server.go does:
// buffered channels, the caller owning in and the backend owning out.
func startLive(p *LocalAIProxy) *liveCall {
lc := &liveCall{
in: make(chan *pb.TranscriptLiveRequest, 4),
out: make(chan *pb.TranscriptLiveResponse, 4),
errc: make(chan error, 1),
}
go func() { lc.errc <- p.AudioTranscriptionLive(lc.in, lc.out) }()
return lc
}
func (lc *liveCall) config(lang string, rate int32) {
lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Config{
Config: &pb.TranscriptLiveConfig{Language: lang, SampleRate: rate},
}}
}
func (lc *liveCall) audio(pcm ...float32) {
lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Audio{
Audio: &pb.TranscriptLiveAudio{Pcm: pcm},
}}
}
func (lc *liveCall) next() *pb.TranscriptLiveResponse {
var r *pb.TranscriptLiveResponse
EventuallyWithOffset(1, lc.out, 2*time.Second).Should(Receive(&r))
return r
}
// finish asserts the call returned within 2 s and out is closed, and
// returns the call's error.
func (lc *liveCall) finish() error {
var err error
EventuallyWithOffset(1, lc.errc, 2*time.Second).Should(Receive(&err))
for range lc.out {
}
return err
}
// liveGoroutines counts goroutines still running code from live.go, so a
// spec can prove the bridge left nothing blocked behind.
func liveGoroutines() int {
buf := make([]byte, 1<<20)
buf = buf[:runtime.Stack(buf, true)]
n := 0
for _, g := range strings.Split(string(buf), "\n\n") {
if strings.Contains(g, "localai-proxy/live.go:") {
n++
}
}
return n
}
func ev(typ string) map[string]any { return map[string]any{"type": typ} }
func item(typ, id string) map[string]any { return map[string]any{"type": typ, "item_id": id} }
func completed(id, transcript string) map[string]any {
return map[string]any{"type": "conversation.item.input_audio_transcription.completed", "item_id": id, "transcript": transcript}
}
// stallAfterUpdate plays session.created, reads the session.update and then
// never answers, as a hung upstream would, until the spec ends.
func stallAfterUpdate(c *websocket.Conn, gone chan<- struct{}) {
wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}})
wsRecv(c)
for {
if _, ok := wsRecv(c); !ok {
close(gone)
return
}
}
}
func withPipeline(o *pb.ModelOptions) {
o.Options = append(o.Options, "realtime_pipeline:remote-pipe")
}
var _ = Describe("AudioTranscriptionLive", func() {
It("reports live transcription unsupported without realtime_pipeline", func() {
up := newFakeUpstream()
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, nil))
err := lc.finish()
Expect(grpcerrors.IsLiveTranscriptionUnsupported(err)).To(BeTrue())
Expect(err.Error()).To(ContainSubstring("realtime_pipeline"))
Expect(up.recorded()).To(BeEmpty())
})
It("rejects a first message that is not a config", func() {
up := wsUpstream(func(*websocket.Conn, *http.Request) {})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.audio(0.1)
Expect(status.Code(lc.finish())).To(Equal(codes.InvalidArgument))
})
It("opens the pipeline session and acks ready only after session.updated", func() {
gotURL := make(chan string, 1)
gotAuth := make(chan string, 1)
gotUpdate := make(chan map[string]any, 1)
release := make(chan struct{})
up := wsUpstream(func(c *websocket.Conn, r *http.Request) {
gotURL <- r.URL.RequestURI()
gotAuth <- r.Header.Get("Authorization")
wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}})
upd, _ := wsRecv(c)
gotUpdate <- upd
<-release
wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}})
wsDrain(c)
})
DeferCleanup(up.Close)
keyFile := writeInput("key", "sekret\n")
p := loadProxy(up, func(o *pb.ModelOptions) {
withPipeline(o)
o.Proxy.ApiKeyFile = keyFile
})
lc := startLive(p)
lc.config("it", 24000)
Eventually(gotURL, 2*time.Second).Should(Receive(Equal("/v1/realtime?model=remote-pipe")))
Expect(gotAuth).To(Receive(Equal("Bearer sekret")))
var upd map[string]any
Eventually(gotUpdate, 2*time.Second).Should(Receive(&upd))
raw, err := json.Marshal(upd)
Expect(err).NotTo(HaveOccurred())
Expect(raw).To(MatchJSON(`{"type":"session.update","session":{"type":"transcription","audio":{"input":{
"format":{"type":"audio/pcm","rate":24000},
"transcription":{"model":"remote-pipe","language":"it"},
"turn_detection":{"type":"server_vad"}}}}}`))
Consistently(lc.out, 200*time.Millisecond).ShouldNot(Receive())
close(release)
Expect(lc.next().GetReady()).To(BeTrue())
close(lc.in)
Expect(lc.finish()).To(Succeed())
})
It("defaults the session rate to 16000", func() {
gotUpdate := make(chan map[string]any, 1)
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
gotUpdate <- wsHandshake(c)
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("", 0)
Expect(lc.next().GetReady()).To(BeTrue())
var upd map[string]any
Expect(gotUpdate).To(Receive(&upd))
format := upd["session"].(map[string]any)["audio"].(map[string]any)["input"].(map[string]any)["format"]
Expect(format).To(HaveKeyWithValue("rate", BeNumerically("==", 16000)))
close(lc.in)
Expect(lc.finish()).To(Succeed())
})
It("forwards audio as base64 PCM16 appends", func() {
frames := make(chan [][]int16, 1)
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
frames <- wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
lc.audio(0, 0.5, -1, 1, 2)
lc.audio(-0.25)
lc.audio(float32(math.NaN()))
close(lc.in)
Expect(lc.finish()).To(Succeed())
Eventually(frames, 2*time.Second).Should(Receive(Equal([][]int16{
{0, 16383, -32767, 32767, 32767},
{-8191},
{0},
})))
})
It("maps deltas and completions to Delta and Eou, and finishes with the full text", func() {
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
wsSend(c, ev("input_audio_buffer.speech_started"))
wsSend(c, ev("input_audio_buffer.speech_stopped"))
wsSend(c, item("input_audio_buffer.committed", "a"))
wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "hel"})
wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "lo"})
wsSend(c, completed("a", "hello world"))
wsSend(c, ev("input_audio_buffer.speech_started"))
wsSend(c, ev("input_audio_buffer.speech_stopped"))
wsSend(c, item("input_audio_buffer.committed", "b"))
wsSend(c, completed("b", "again"))
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
Expect(lc.next()).To(SatisfyAll(
WithTransform((*pb.TranscriptLiveResponse).GetDelta, Equal("hel")),
WithTransform((*pb.TranscriptLiveResponse).GetEou, BeFalse())))
Expect(lc.next().GetDelta()).To(Equal("lo"))
r := lc.next()
Expect(r.GetDelta()).To(Equal(" world"))
Expect(r.GetEou()).To(BeTrue())
r = lc.next()
Expect(r.GetDelta()).To(Equal("again"))
Expect(r.GetEou()).To(BeTrue())
close(lc.in)
Expect(lc.next().GetFinalResult().GetText()).To(Equal("hello world again"))
Expect(lc.finish()).To(Succeed())
})
It("waits for a committed utterance before the final result", func() {
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
wsSend(c, ev("input_audio_buffer.speech_started"))
wsSend(c, ev("input_audio_buffer.speech_stopped"))
wsSend(c, item("input_audio_buffer.committed", "a"))
wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.delta", "item_id": "a", "delta": "late"})
// The client closes its side once it sees the delta; the
// transcription completes later.
time.Sleep(300 * time.Millisecond)
wsSend(c, completed("a", "late words"))
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
Expect(lc.next().GetDelta()).To(Equal("late"))
close(lc.in)
r := lc.next()
Expect(r.GetDelta()).To(Equal(" words"))
Expect(r.GetEou()).To(BeTrue())
Expect(lc.next().GetFinalResult().GetText()).To(Equal("late words"))
Expect(lc.finish()).To(Succeed())
})
It("does not hold the close for a turn the upstream discarded", func() {
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
wsSend(c, ev("input_audio_buffer.speech_started"))
wsSend(c, ev("input_audio_buffer.speech_stopped"))
wsSend(c, item("input_audio_buffer.committed", "a"))
wsSend(c, completed("a", "kept"))
// A stop that is never committed, then nothing more.
wsSend(c, ev("input_audio_buffer.speech_started"))
wsSend(c, ev("input_audio_buffer.speech_stopped"))
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
Expect(lc.next().GetDelta()).To(Equal("kept"))
close(lc.in)
start := time.Now()
Expect(lc.next().GetFinalResult().GetText()).To(Equal("kept"))
Expect(lc.finish()).To(Succeed())
Expect(time.Since(start)).To(BeNumerically("<", finalWait/2))
})
It("ends with Unavailable on an upstream error event during setup", func() {
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}})
wsRecv(c)
wsSend(c, map[string]any{"type": "error", "error": map[string]any{
"type": "invalid_request_error", "code": "session_update_error",
"message": "model is not a valid pipeline model: remote-pipe",
}})
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
err := lc.finish()
Expect(status.Code(err)).To(Equal(codes.Unavailable))
Expect(err.Error()).To(ContainSubstring("not a valid pipeline model"))
})
It("ends with Unavailable when a transcription fails", func() {
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
wsSend(c, map[string]any{"type": "conversation.item.input_audio_transcription.failed", "item_id": "a",
"error": map[string]any{"message": "backend crashed"}})
wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
err := lc.finish()
Expect(status.Code(err)).To(Equal(codes.Unavailable))
Expect(err.Error()).To(ContainSubstring("backend crashed"))
})
It("maps a refused upgrade to the upstream status", func() {
up := newFakeUpstreamWithHandler(func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "loading", http.StatusServiceUnavailable)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable))
})
It("upstream disconnect ends the stream", func() {
before := liveGoroutines()
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsHandshake(c)
// Drop the socket mid-session, as a crashed upstream would.
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(lc.next().GetReady()).To(BeTrue())
// The caller keeps its send side open; the bridge must not wait on it.
Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable))
Eventually(liveGoroutines, time.Second).Should(Equal(before))
close(lc.in)
})
It("gives up with Canceled when the caller closes before the ready ack", func() {
before := liveGoroutines()
gone := make(chan struct{})
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) })
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
lc.audio(0.1)
close(lc.in)
Expect(status.Code(lc.finish())).To(Equal(codes.Canceled))
Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open")
Eventually(liveGoroutines, time.Second).Should(Equal(before))
})
It("bounds a hung setup by request_timeout_seconds", func() {
before := liveGoroutines()
gone := make(chan struct{})
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) })
DeferCleanup(up.Close)
p := loadProxy(up, func(o *pb.ModelOptions) {
withPipeline(o)
o.Proxy.RequestTimeoutSeconds = 1
})
lc := startLive(p)
lc.config("en", 16000)
var err error
Eventually(lc.errc, 3*time.Second).Should(Receive(&err))
Expect(status.Code(err)).To(Equal(codes.Unavailable))
Expect(lc.out).To(BeClosed())
Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open")
Eventually(liveGoroutines, time.Second).Should(Equal(before))
close(lc.in)
})
It("bounds a hung setup by default when request_timeout_seconds is unset", func() {
Expect(defaultLiveSetupTimeout).To(BeNumerically(">=", 2*time.Minute))
saved := liveSetupTimeout
liveSetupTimeout = 300 * time.Millisecond
DeferCleanup(func() { liveSetupTimeout = saved })
before := liveGoroutines()
gone := make(chan struct{})
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) })
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
Expect(status.Code(lc.finish())).To(Equal(codes.Unavailable))
Eventually(gone, 2*time.Second).Should(BeClosed(), "upstream socket left open")
Eventually(liveGoroutines, time.Second).Should(Equal(before))
close(lc.in)
})
It("forwards audio sent before the ready ack once ready", func() {
release := make(chan struct{})
frames := make(chan [][]int16, 1)
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) {
wsSend(c, map[string]any{"type": "session.created", "session": map[string]any{}})
wsRecv(c)
<-release
wsSend(c, map[string]any{"type": "session.updated", "session": map[string]any{}})
frames <- wsDrain(c)
})
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
lc.audio(0.5)
close(release)
Expect(lc.next().GetReady()).To(BeTrue())
close(lc.in)
Expect(lc.finish()).To(Succeed())
Eventually(frames, 2*time.Second).Should(Receive(Equal([][]int16{{16383}})))
})
It("refuses more than the backlog cap of audio before the ready ack", func() {
gone := make(chan struct{})
up := wsUpstream(func(c *websocket.Conn, _ *http.Request) { stallAfterUpdate(c, gone) })
DeferCleanup(up.Close)
lc := startLive(loadProxy(up, withPipeline))
lc.config("en", 16000)
second := make([]float32, 16000)
for range maxBacklogSeconds + 1 {
select {
case lc.in <- &pb.TranscriptLiveRequest{Payload: &pb.TranscriptLiveRequest_Audio{Audio: &pb.TranscriptLiveAudio{Pcm: second}}}:
case <-time.After(2 * time.Second):
Fail("bridge stopped reading audio before the cap")
}
}
err := lc.finish()
Expect(status.Code(err)).To(Equal(codes.InvalidArgument))
Expect(err.Error()).To(ContainSubstring("before the live transcription session was ready"))
close(lc.in)
})
})
@@ -0,0 +1,17 @@
package main
import (
"testing"
"github.com/mudler/xlog"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
func TestLocalAIProxy(t *testing.T) {
RegisterFailHandler(Fail)
// The specs drive upstream failures on purpose; their warnings are
// expected and would only bury real failures in the output.
xlog.SetLogger(xlog.NewLogger(xlog.LogLevelError, xlog.TextFormat))
RunSpecs(t, "localai-proxy specs")
}
+36
View File
@@ -0,0 +1,36 @@
package main
// localai-proxy is a LocalAI backend that serves backend gRPC methods by
// calling the REST API of another LocalAI instance. It lets a model config
// (and so a failover chain target) live on a remote LocalAI while callers
// keep using the local backend interface for every modality, not only chat.
import (
"flag"
"os"
grpc "github.com/mudler/LocalAI/pkg/grpc"
"github.com/mudler/xlog"
"golang.org/x/term"
)
var addr = flag.String("addr", "localhost:50051", "the address to listen on")
func main() {
// xlog's default handler emits ANSI color codes, which are unreadable once
// LocalAI captures the backend's stdout into a log file. Force plain text
// when LOCALAI_LOG_FORMAT is unset and stdout is not a terminal.
format := os.Getenv("LOCALAI_LOG_FORMAT")
if format == "" && !term.IsTerminal(int(os.Stdout.Fd())) {
format = xlog.TextFormat
}
xlog.SetLogger(xlog.NewLogger(xlog.LogLevel(os.Getenv("LOCALAI_LOG_LEVEL")), format))
flag.Parse()
if err := grpc.StartServer(*addr, NewLocalAIProxy()); err != nil {
panic(err)
}
}
// The gRPC server only passes the call's context to backends that implement
// this, so keep the proxy on it.
var _ grpc.AIModelRichContext = (*LocalAIProxy)(nil)
+862
View File
@@ -0,0 +1,862 @@
package main
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
// fileToBase64 reads path and base64-encodes its contents, or returns "" for
// an empty path. Core stages image/video/3D conditioning inputs to local
// files before calling the backend, but the REST endpoints on the other side
// take the same media inline as base64 (or a URL/data-URI, neither of which
// a local path is), so every media method needs this same conversion.
func fileToBase64(path string) (string, error) {
if path == "" {
return "", nil
}
// #nosec G304 -- path is a staging file core wrote for this call, never a caller-supplied path
data, err := os.ReadFile(filepath.Clean(path))
if err != nil {
return "", status.Errorf(codes.InvalidArgument, "localai-proxy: read %s: %v", path, err)
}
return base64.StdEncoding.EncodeToString(data), nil
}
// toDataURI wraps a base64 payload as a data URI. Detect/Depth/FaceVerify/
// FaceAnalyze's REST endpoints decode their image field with
// utils.GetContentURIAsBase64, which only accepts an http(s) URL or a
// `data:<mime>;base64,<payload>` string — never bare base64. That is exactly
// what core hands the backend for these methods (see
// core/http/endpoints/localai/images.go's decodeImageInput, which already
// stripped any data: prefix off before we ever see it), so forwarding it
// unwrapped 400s on every real call. The MIME type isn't carried alongside
// the payload, so it's sniffed from the decoded bytes.
func toDataURI(b64 string) (string, error) {
if b64 == "" {
return "", nil
}
data, err := base64.StdEncoding.DecodeString(b64)
if err != nil {
return "", status.Errorf(codes.InvalidArgument, "localai-proxy: decode base64 image: %v", err)
}
mime := http.DetectContentType(data)
if mime == "application/octet-stream" {
mime = "image/png"
}
return "data:" + mime + ";base64," + b64, nil
}
// genItem is one schema.Item as the image/video/3D generation endpoints
// return it: either inline base64 data, or a URL to download the asset from.
type genItem struct {
URL string `json:"url,omitempty"`
B64JSON string `json:"b64_json,omitempty"`
}
type genResponse struct {
Data []genItem `json:"data"`
}
// writeGenItem writes the first item of a generation reply to dst: it
// base64-decodes inline data directly, or downloads the URL relative to the
// configured upstream (same client, same auth) when the upstream sent one
// instead.
func (p *LocalAIProxy) writeGenItem(ctx context.Context, path string, items []genItem, dst string) error {
if len(items) == 0 {
return status.Errorf(codes.Internal, "localai-proxy: upstream %s returned no data", path)
}
item := items[0]
if item.B64JSON != "" {
data, err := base64.StdEncoding.DecodeString(item.B64JSON)
if err != nil {
return status.Errorf(codes.Internal, "localai-proxy: decode %s b64_json: %v", path, err)
}
// 0o600: core runs as the same user and serves the file itself, so no
// one else needs to read generated media.
if err := os.WriteFile(dst, data, 0o600); err != nil {
_ = os.Remove(dst)
return status.Errorf(codes.Internal, "localai-proxy: write %s: %v", dst, err)
}
return nil
}
if item.URL == "" {
return status.Errorf(codes.Internal, "localai-proxy: upstream %s returned neither b64_json nor url", path)
}
rel, err := generatedContentPath(path, item.URL)
if err != nil {
return err
}
return p.getToFile(ctx, rel, dst)
}
// generatedContentPrefixes are the static paths a LocalAI instance serves
// generated media under (core/http/app.go's e.Static calls). A generation
// reply's URL is only ever safe to re-fetch through this proxy's own
// authenticated client when it resolves to one of these.
var generatedContentPrefixes = []string{"/generated-images/", "/generated-videos/", "/generated-audio/", "/generated-3d/"}
// generatedContentPath extracts the path to re-download raw from, ignoring
// whatever host is in it. A literal-prefix strip of the configured upstream
// base (the previous approach) breaks the moment the upstream advertises a
// different host than the one this proxy is configured with — LOCALAI_BASE_URL,
// a reverse proxy, or X-Forwarded-Host can all change it — so only the path is
// trusted, and only when it is one this proxy's upstream is actually known to
// serve generated media under; anything else could point anywhere, and taking
// it on faith would let a compromised or misconfigured upstream make this
// proxy fetch (with its bearer key) whatever URL it likes.
func generatedContentPath(callPath, raw string) (string, error) {
u, err := url.Parse(raw)
if err != nil {
return "", status.Errorf(codes.Internal, "localai-proxy: upstream %s returned an invalid url %q: %v", callPath, raw, err)
}
for _, prefix := range generatedContentPrefixes {
if strings.HasPrefix(u.Path, prefix) {
if u.RawQuery != "" {
return u.Path + "?" + u.RawQuery, nil
}
return u.Path, nil
}
}
return "", status.Errorf(codes.InvalidArgument, "localai-proxy: upstream %s returned an unexpected url %q", callPath, raw)
}
// --- Images ---------------------------------------------------------
type imageGenerationRequest struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
NegativePrompt string `json:"negative_prompt,omitempty"`
Size string `json:"size,omitempty"`
Step int32 `json:"step,omitempty"`
Seed int32 `json:"seed,omitempty"`
ResponseFormat string `json:"response_format,omitempty"`
File string `json:"file,omitempty"`
RefImages []string `json:"ref_images,omitempty"`
}
// GenerateImage sends Src and RefImages (local paths) as base64, asks for a
// b64_json reply, and writes the result to Dst.
func (p *LocalAIProxy) GenerateImage(req *pb.GenerateImageRequest) error {
ctx := context.Background()
file, err := fileToBase64(req.GetSrc())
if err != nil {
return err
}
var refImages []string
for _, ref := range req.GetRefImages() {
encoded, err := fileToBase64(ref)
if err != nil {
return err
}
refImages = append(refImages, encoded)
}
body := imageGenerationRequest{
Model: p.model(""),
Prompt: req.GetPositivePrompt(),
NegativePrompt: req.GetNegativePrompt(),
Size: fmt.Sprintf("%dx%d", req.GetWidth(), req.GetHeight()),
Step: req.GetStep(),
Seed: req.GetSeed(),
ResponseFormat: "b64_json",
File: file,
RefImages: refImages,
}
var resp genResponse
if err := p.postJSON(ctx, "/v1/images/generations", body, &resp); err != nil {
return err
}
return p.writeGenItem(ctx, "/v1/images/generations", resp.Data, req.GetDst())
}
// UpscaleImage uploads Src as a multipart file, matching the REST endpoint's
// upload-only contract, and writes the (always URL) reply to Dst.
func (p *LocalAIProxy) UpscaleImage(req *pb.UpscaleImageRequest) error {
ctx := context.Background()
fields := url.Values{}
fields.Set("model", p.model(""))
if req.GetScale() > 0 {
fields.Set("scale", strconv.Itoa(int(req.GetScale())))
}
form := multipartForm{fields: fields, files: []formFile{{field: "image", path: req.GetSrc()}}}
var resp genResponse
if err := p.postForm(ctx, "/v1/images/upscale", form, &resp); err != nil {
return err
}
return p.writeGenItem(ctx, "/v1/images/upscale", resp.Data, req.GetDst())
}
// --- Video ------------------------------------------------------------
type videoGenerationRequest struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
NegativePrompt string `json:"negative_prompt,omitempty"`
StartImage string `json:"start_image,omitempty"`
EndImage string `json:"end_image,omitempty"`
Audio string `json:"audio,omitempty"`
Width int32 `json:"width,omitempty"`
Height int32 `json:"height,omitempty"`
NumFrames int32 `json:"num_frames,omitempty"`
FPS int32 `json:"fps,omitempty"`
Seed int32 `json:"seed,omitempty"`
CFGScale float32 `json:"cfg_scale,omitempty"`
Step int32 `json:"step,omitempty"`
ResponseFormat string `json:"response_format,omitempty"`
Params map[string]string `json:"params,omitempty"`
}
// GenerateVideo sends the staged local media (StartImage, EndImage, Audio) as
// base64, asks for a b64_json reply, and writes the result to Dst.
func (p *LocalAIProxy) GenerateVideo(req *pb.GenerateVideoRequest) error {
ctx := context.Background()
startImage, err := fileToBase64(req.GetStartImage())
if err != nil {
return err
}
endImage, err := fileToBase64(req.GetEndImage())
if err != nil {
return err
}
audio, err := fileToBase64(req.GetAudio())
if err != nil {
return err
}
body := videoGenerationRequest{
Model: p.model(""),
Prompt: req.GetPrompt(),
NegativePrompt: req.GetNegativePrompt(),
StartImage: startImage,
EndImage: endImage,
Audio: audio,
Width: req.GetWidth(),
Height: req.GetHeight(),
NumFrames: req.GetNumFrames(),
FPS: req.GetFps(),
Seed: req.GetSeed(),
CFGScale: req.GetCfgScale(),
Step: req.GetStep(),
ResponseFormat: "b64_json",
Params: req.GetParams(),
}
var resp genResponse
if err := p.postJSON(ctx, "/video", body, &resp); err != nil {
return err
}
return p.writeGenItem(ctx, "/video", resp.Data, req.GetDst())
}
// --- 3D -----------------------------------------------------------------
type model3DRequest struct {
Model string `json:"model"`
Image string `json:"image"`
Seed int32 `json:"seed,omitempty"`
Step int32 `json:"step,omitempty"`
CFGScale float32 `json:"cfg_scale,omitempty"`
TextureSteps int32 `json:"texture_steps,omitempty"`
Quality string `json:"quality,omitempty"`
Background string `json:"background,omitempty"`
ResponseFormat string `json:"response_format,omitempty"`
Params map[string]string `json:"params,omitempty"`
}
// Generate3D sends Src (the staged conditioning image, a local path) as
// base64, asks for a b64_json reply, and writes the result to Dst.
func (p *LocalAIProxy) Generate3D(req *pb.Generate3DRequest) error {
ctx := context.Background()
image, err := fileToBase64(req.GetSrc())
if err != nil {
return err
}
body := model3DRequest{
Model: p.model(""),
Image: image,
Seed: req.GetSeed(),
Step: req.GetStep(),
CFGScale: req.GetCfgScale(),
TextureSteps: req.GetTextureSteps(),
Quality: req.GetQuality(),
Background: req.GetBackground(),
ResponseFormat: "b64_json",
Params: req.GetParams(),
}
var resp genResponse
if err := p.postJSON(ctx, "/3d/generations", body, &resp); err != nil {
return err
}
return p.writeGenItem(ctx, "/3d/generations", resp.Data, req.GetDst())
}
type animationInputBody struct {
Type string `json:"type"`
Data string `json:"data"`
}
type animate3DRequestBody struct {
Model string `json:"model"`
Inputs map[string]animationInputBody `json:"inputs"`
Params map[string]string `json:"params,omitempty"`
ResponseFormat string `json:"response_format,omitempty"`
}
type animate3DResponseBody struct {
genResponse
Metadata json.RawMessage `json:"metadata,omitempty"`
}
// animate3D does the shared work behind Animate3D and Animate3DWithMetadata:
// non-text inputs are staged local paths, so they go over the wire as
// base64, like every other media field; text inputs travel verbatim.
func (p *LocalAIProxy) animate3D(ctx context.Context, req *pb.Animate3DRequest) (json.RawMessage, error) {
inputs := make(map[string]animationInputBody, len(req.GetInputs()))
for name, in := range req.GetInputs() {
data := in.GetData()
if in.GetType() != "text" {
encoded, err := fileToBase64(data)
if err != nil {
return nil, err
}
data = encoded
}
inputs[name] = animationInputBody{Type: in.GetType(), Data: data}
}
body := animate3DRequestBody{
Model: p.model(""),
Inputs: inputs,
Params: req.GetParams(),
ResponseFormat: "b64_json",
}
var resp animate3DResponseBody
if err := p.postJSON(ctx, "/3d/animate", body, &resp); err != nil {
return nil, err
}
if err := p.writeGenItem(ctx, "/3d/animate", resp.Data, req.GetDst()); err != nil {
return nil, err
}
return resp.Metadata, nil
}
func (p *LocalAIProxy) Animate3D(req *pb.Animate3DRequest) error {
_, err := p.animate3D(context.Background(), req)
return err
}
// Animate3DWithMetadata implements grpc.AnimationMetadataModel: the REST
// reply's metadata field carries whatever the animation backend reported, and
// the gRPC server prefers this method over Animate3D when it is implemented.
func (p *LocalAIProxy) Animate3DWithMetadata(req *pb.Animate3DRequest) ([]byte, error) {
return p.animate3D(context.Background(), req)
}
// --- Vision (detection, depth) ------------------------------------------
type detectRequestBody struct {
Model string `json:"model"`
Image string `json:"image"`
Prompt string `json:"prompt,omitempty"`
Points []float32 `json:"points,omitempty"`
Boxes []float32 `json:"boxes,omitempty"`
Threshold float32 `json:"threshold,omitempty"`
}
type detectionBody struct {
X float32 `json:"x"`
Y float32 `json:"y"`
Width float32 `json:"width"`
Height float32 `json:"height"`
ClassName string `json:"class_name"`
Confidence float32 `json:"confidence,omitempty"`
Mask string `json:"mask,omitempty"`
}
type detectResponseBody struct {
Detections []detectionBody `json:"detections"`
}
// Detect posts Src, which core already carries as a bare base64 payload (the
// same convention DetectionEndpoint uses to call this method locally) — but
// the upstream's own endpoint only accepts a URL or a data URI, so it is
// re-wrapped as one (see toDataURI) — and maps each detection, decoding its
// PNG mask.
func (p *LocalAIProxy) Detect(req *pb.DetectOptions) (pb.DetectResponse, error) {
image, err := toDataURI(req.GetSrc())
if err != nil {
return pb.DetectResponse{}, err
}
body := detectRequestBody{
Model: p.model(""),
Image: image,
Prompt: req.GetPrompt(),
Points: req.GetPoints(),
Boxes: req.GetBoxes(),
Threshold: req.GetThreshold(),
}
var resp detectResponseBody
if err := p.postJSON(context.Background(), "/v1/detection", body, &resp); err != nil {
return pb.DetectResponse{}, err
}
var detections []*pb.Detection
for _, d := range resp.Detections {
det := &pb.Detection{
X: d.X, Y: d.Y, Width: d.Width, Height: d.Height,
Confidence: d.Confidence, ClassName: d.ClassName,
}
if d.Mask != "" {
mask, err := base64.StdEncoding.DecodeString(d.Mask)
if err != nil {
return pb.DetectResponse{}, status.Errorf(codes.Internal, "localai-proxy: decode /v1/detection mask: %v", err)
}
det.Mask = mask
}
detections = append(detections, det)
}
// A composite literal here, not a variable of type pb.DetectResponse,
// avoids copying the protobuf message's embedded lock on return.
return pb.DetectResponse{Detections: detections}, nil
}
// depthRequestBody has no Dst/Exports fields: Depth refuses those requests
// before building the body (see the Depth doc comment), so they never reach
// the upstream.
type depthRequestBody struct {
Model string `json:"model"`
Image string `json:"image"`
IncludeDepth bool `json:"include_depth,omitempty"`
IncludeConfidence bool `json:"include_confidence,omitempty"`
IncludePose bool `json:"include_pose,omitempty"`
IncludeSky bool `json:"include_sky,omitempty"`
IncludePoints bool `json:"include_points,omitempty"`
PointsConfThresh float32 `json:"points_conf_thresh,omitempty"`
}
type depthResponseBody struct {
Width int32 `json:"width"`
Height int32 `json:"height"`
Depth []float32 `json:"depth,omitempty"`
Confidence []float32 `json:"confidence,omitempty"`
Sky []float32 `json:"sky,omitempty"`
Extrinsics []float32 `json:"extrinsics,omitempty"`
Intrinsics []float32 `json:"intrinsics,omitempty"`
NumPoints int32 `json:"num_points,omitempty"`
Points []float32 `json:"points,omitempty"`
PointColors string `json:"point_colors,omitempty"`
ExportPaths []string `json:"export_paths,omitempty"`
IsMetric bool `json:"is_metric"`
}
// Depth posts Src (a bare base64 payload, per the same convention as Detect,
// re-wrapped as a data URI via toDataURI) and maps the full response,
// decoding the point-cloud color bytes. A request for exports or a dst
// directory is refused: those would be written to the upstream's own local
// disk, and ExportPaths would name files this proxy (and whatever asked it
// for them) can never reach.
func (p *LocalAIProxy) Depth(req *pb.DepthRequest) (pb.DepthResponse, error) {
if req.GetDst() != "" || len(req.GetExports()) > 0 {
return pb.DepthResponse{}, unimplemented("Depth exports (written to the upstream's own local disk, unreachable from here)")
}
image, err := toDataURI(req.GetSrc())
if err != nil {
return pb.DepthResponse{}, err
}
body := depthRequestBody{
Model: p.model(""),
Image: image,
IncludeDepth: req.GetIncludeDepth(),
IncludeConfidence: req.GetIncludeConfidence(),
IncludePose: req.GetIncludePose(),
IncludeSky: req.GetIncludeSky(),
IncludePoints: req.GetIncludePoints(),
PointsConfThresh: req.GetPointsConfThresh(),
}
var resp depthResponseBody
if err := p.postJSON(context.Background(), "/v1/depth", body, &resp); err != nil {
return pb.DepthResponse{}, err
}
var colors []byte
if resp.PointColors != "" {
var err error
colors, err = base64.StdEncoding.DecodeString(resp.PointColors)
if err != nil {
return pb.DepthResponse{}, status.Errorf(codes.Internal, "localai-proxy: decode /v1/depth point_colors: %v", err)
}
}
// A composite literal here, not a variable of type pb.DepthResponse,
// avoids copying the protobuf message's embedded lock on return.
return pb.DepthResponse{
Width: resp.Width, Height: resp.Height, Depth: resp.Depth, Confidence: resp.Confidence,
Sky: resp.Sky, Extrinsics: resp.Extrinsics, Intrinsics: resp.Intrinsics,
NumPoints: resp.NumPoints, Points: resp.Points, ExportPaths: resp.ExportPaths, IsMetric: resp.IsMetric,
PointColors: colors,
}, nil
}
// --- Face recognition -----------------------------------------------------
type facialAreaBody struct {
X float32 `json:"x"`
Y float32 `json:"y"`
W float32 `json:"w"`
H float32 `json:"h"`
}
func (a facialAreaBody) toProto() *pb.FacialArea {
return &pb.FacialArea{X: a.X, Y: a.Y, W: a.W, H: a.H}
}
type faceVerifyRequestBody struct {
Model string `json:"model"`
Img1 string `json:"img1"`
Img2 string `json:"img2"`
Threshold float32 `json:"threshold,omitempty"`
AntiSpoofing bool `json:"anti_spoofing,omitempty"`
}
type faceVerifyResponseBody struct {
Verified bool `json:"verified"`
Distance float32 `json:"distance"`
Threshold float32 `json:"threshold"`
Confidence float32 `json:"confidence"`
Model string `json:"model"`
Img1Area facialAreaBody `json:"img1_area"`
Img2Area facialAreaBody `json:"img2_area"`
ProcessingTimeMs float32 `json:"processing_time_ms,omitempty"`
Img1IsReal *bool `json:"img1_is_real,omitempty"`
Img1AntispoofScore *float32 `json:"img1_antispoof_score,omitempty"`
Img2IsReal *bool `json:"img2_is_real,omitempty"`
Img2AntispoofScore *float32 `json:"img2_antispoof_score,omitempty"`
}
// FaceVerify posts Img1/Img2, which core already carries as bare base64 (the
// same convention FaceVerifyEndpoint uses to call this method locally) —
// re-wrapped as data URIs via toDataURI, since the upstream's own endpoint
// only accepts a URL or a data URI.
func (p *LocalAIProxy) FaceVerify(req *pb.FaceVerifyRequest) (pb.FaceVerifyResponse, error) {
img1, err := toDataURI(req.GetImg1())
if err != nil {
return pb.FaceVerifyResponse{}, err
}
img2, err := toDataURI(req.GetImg2())
if err != nil {
return pb.FaceVerifyResponse{}, err
}
body := faceVerifyRequestBody{
Model: p.model(""), Img1: img1, Img2: img2,
Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(),
}
var resp faceVerifyResponseBody
if err := p.postJSON(context.Background(), "/v1/face/verify", body, &resp); err != nil {
return pb.FaceVerifyResponse{}, err
}
// A composite literal here, not a variable of type pb.FaceVerifyResponse,
// avoids copying the protobuf message's embedded lock on return.
return pb.FaceVerifyResponse{
Verified: resp.Verified, Distance: resp.Distance, Threshold: resp.Threshold,
Confidence: resp.Confidence, Model: resp.Model,
Img1Area: resp.Img1Area.toProto(), Img2Area: resp.Img2Area.toProto(),
ProcessingTimeMs: resp.ProcessingTimeMs,
Img1IsReal: boolValue(resp.Img1IsReal),
Img1AntispoofScore: float32Value(resp.Img1AntispoofScore),
Img2IsReal: boolValue(resp.Img2IsReal),
Img2AntispoofScore: float32Value(resp.Img2AntispoofScore),
}, nil
}
// boolValue and float32Value read the liveness fields the REST endpoints only
// populate when anti_spoofing was requested; proto keeps them as plain
// bool/float32 (no "not checked" state), so an absent pointer becomes zero.
func boolValue(b *bool) bool {
return b != nil && *b
}
func float32Value(v *float32) float32 {
if v == nil {
return 0
}
return *v
}
type faceAnalyzeRequestBody struct {
Model string `json:"model"`
Img string `json:"img"`
Actions []string `json:"actions,omitempty"`
AntiSpoofing bool `json:"anti_spoofing,omitempty"`
}
type faceAnalysisBody struct {
Region facialAreaBody `json:"region"`
FaceConfidence float32 `json:"face_confidence"`
Age float32 `json:"age,omitempty"`
DominantGender string `json:"dominant_gender,omitempty"`
Gender map[string]float32 `json:"gender,omitempty"`
DominantEmotion string `json:"dominant_emotion,omitempty"`
Emotion map[string]float32 `json:"emotion,omitempty"`
DominantRace string `json:"dominant_race,omitempty"`
Race map[string]float32 `json:"race,omitempty"`
IsReal *bool `json:"is_real,omitempty"`
AntispoofScore *float32 `json:"antispoof_score,omitempty"`
}
type faceAnalyzeResponseBody struct {
Faces []faceAnalysisBody `json:"faces"`
}
// FaceAnalyze posts Img, which core already carries as bare base64 —
// re-wrapped as a data URI via toDataURI, since the upstream's own endpoint
// only accepts a URL or a data URI.
func (p *LocalAIProxy) FaceAnalyze(req *pb.FaceAnalyzeRequest) (pb.FaceAnalyzeResponse, error) {
img, err := toDataURI(req.GetImg())
if err != nil {
return pb.FaceAnalyzeResponse{}, err
}
body := faceAnalyzeRequestBody{
Model: p.model(""), Img: img, Actions: req.GetActions(), AntiSpoofing: req.GetAntiSpoofing(),
}
var resp faceAnalyzeResponseBody
if err := p.postJSON(context.Background(), "/v1/face/analyze", body, &resp); err != nil {
return pb.FaceAnalyzeResponse{}, err
}
var faces []*pb.FaceAnalysis
for _, f := range resp.Faces {
faces = append(faces, &pb.FaceAnalysis{
Region: f.Region.toProto(), FaceConfidence: f.FaceConfidence, Age: f.Age,
DominantGender: f.DominantGender, Gender: f.Gender,
DominantEmotion: f.DominantEmotion, Emotion: f.Emotion,
DominantRace: f.DominantRace, Race: f.Race,
IsReal: boolValue(f.IsReal), AntispoofScore: float32Value(f.AntispoofScore),
})
}
// A composite literal here, not a variable of type pb.FaceAnalyzeResponse,
// avoids copying the protobuf message's embedded lock on return.
return pb.FaceAnalyzeResponse{Faces: faces}, nil
}
// --- Voice (speaker) recognition -----------------------------------------
type voiceVerifyRequestBody struct {
Model string `json:"model"`
Audio1 string `json:"audio1"`
Audio2 string `json:"audio2"`
Threshold float32 `json:"threshold,omitempty"`
AntiSpoofing bool `json:"anti_spoofing,omitempty"`
}
type voiceVerifyResponseBody struct {
Verified bool `json:"verified"`
Distance float32 `json:"distance"`
Threshold float32 `json:"threshold"`
Confidence float32 `json:"confidence"`
Model string `json:"model"`
ProcessingTimeMs float32 `json:"processing_time_ms,omitempty"`
}
// VoiceVerify sends Audio1/Audio2 (staged local paths) as base64, matching
// VoiceVerifyRequest's URL/base64/data-URI contract.
func (p *LocalAIProxy) VoiceVerify(req *pb.VoiceVerifyRequest) (pb.VoiceVerifyResponse, error) {
audio1, err := fileToBase64(req.GetAudio1())
if err != nil {
return pb.VoiceVerifyResponse{}, err
}
audio2, err := fileToBase64(req.GetAudio2())
if err != nil {
return pb.VoiceVerifyResponse{}, err
}
body := voiceVerifyRequestBody{
Model: p.model(""), Audio1: audio1, Audio2: audio2,
Threshold: req.GetThreshold(), AntiSpoofing: req.GetAntiSpoofing(),
}
var resp voiceVerifyResponseBody
if err := p.postJSON(context.Background(), "/v1/voice/verify", body, &resp); err != nil {
return pb.VoiceVerifyResponse{}, err
}
return pb.VoiceVerifyResponse{
Verified: resp.Verified, Distance: resp.Distance, Threshold: resp.Threshold,
Confidence: resp.Confidence, Model: resp.Model, ProcessingTimeMs: resp.ProcessingTimeMs,
}, nil
}
type voiceAnalyzeRequestBody struct {
Model string `json:"model"`
Audio string `json:"audio"`
Actions []string `json:"actions,omitempty"`
}
type voiceAnalysisBody struct {
Start float32 `json:"start"`
End float32 `json:"end"`
Age float32 `json:"age,omitempty"`
DominantGender string `json:"dominant_gender,omitempty"`
Gender map[string]float32 `json:"gender,omitempty"`
DominantEmotion string `json:"dominant_emotion,omitempty"`
Emotion map[string]float32 `json:"emotion,omitempty"`
}
type voiceAnalyzeResponseBody struct {
Segments []voiceAnalysisBody `json:"segments"`
}
// VoiceAnalyze sends Audio (a staged local path) as base64.
func (p *LocalAIProxy) VoiceAnalyze(req *pb.VoiceAnalyzeRequest) (pb.VoiceAnalyzeResponse, error) {
audio, err := fileToBase64(req.GetAudio())
if err != nil {
return pb.VoiceAnalyzeResponse{}, err
}
body := voiceAnalyzeRequestBody{Model: p.model(""), Audio: audio, Actions: req.GetActions()}
var resp voiceAnalyzeResponseBody
if err := p.postJSON(context.Background(), "/v1/voice/analyze", body, &resp); err != nil {
return pb.VoiceAnalyzeResponse{}, err
}
var segments []*pb.VoiceAnalysis
for _, s := range resp.Segments {
segments = append(segments, &pb.VoiceAnalysis{
Start: s.Start, End: s.End, Age: s.Age,
DominantGender: s.DominantGender, Gender: s.Gender,
DominantEmotion: s.DominantEmotion, Emotion: s.Emotion,
})
}
// A composite literal here, not a variable of type pb.VoiceAnalyzeResponse,
// avoids copying the protobuf message's embedded lock on return.
return pb.VoiceAnalyzeResponse{Segments: segments}, nil
}
type voiceEmbedRequestBody struct {
Model string `json:"model"`
Audio string `json:"audio"`
}
type voiceEmbedResponseBody struct {
Embedding []float32 `json:"embedding"`
Model string `json:"model,omitempty"`
}
// VoiceEmbed sends Audio (a staged local path) as base64.
func (p *LocalAIProxy) VoiceEmbed(req *pb.VoiceEmbedRequest) (pb.VoiceEmbedResponse, error) {
audio, err := fileToBase64(req.GetAudio())
if err != nil {
return pb.VoiceEmbedResponse{}, err
}
body := voiceEmbedRequestBody{Model: p.model(""), Audio: audio}
var resp voiceEmbedResponseBody
if err := p.postJSON(context.Background(), "/v1/voice/embed", body, &resp); err != nil {
return pb.VoiceEmbedResponse{}, err
}
return pb.VoiceEmbedResponse{Embedding: resp.Embedding, Model: resp.Model}, nil
}
// --- Stores ---------------------------------------------------------------
// storeKeysToFloats and storeFloatsToKeys convert between the proto's boxed
// StoresKey/StoresValue slices and the plain [][]float32 / []string the REST
// stores endpoints take, per schema.StoresSet and friends.
func storeKeysToFloats(keys []*pb.StoresKey) [][]float32 {
out := make([][]float32, len(keys))
for i, k := range keys {
out[i] = k.GetFloats()
}
return out
}
func storeFloatsToKeys(keys [][]float32) []*pb.StoresKey {
out := make([]*pb.StoresKey, len(keys))
for i, k := range keys {
out[i] = &pb.StoresKey{Floats: k}
}
return out
}
func storeValuesToStrings(values []*pb.StoresValue) []string {
out := make([]string, len(values))
for i, v := range values {
out[i] = string(v.GetBytes())
}
return out
}
func storeStringsToValues(values []string) []*pb.StoresValue {
out := make([]*pb.StoresValue, len(values))
for i, v := range values {
out[i] = &pb.StoresValue{Bytes: []byte(v)}
}
return out
}
type storesSetRequestBody struct {
Store string `json:"store,omitempty"`
Keys [][]float32 `json:"keys"`
Values []string `json:"values"`
}
// StoresSet uses the configured upstream model name as the store name: the
// proxy is loaded per store, the same way it is loaded per model for every
// other method.
func (p *LocalAIProxy) StoresSet(req *pb.StoresSetOptions) error {
body := storesSetRequestBody{
Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys()), Values: storeValuesToStrings(req.GetValues()),
}
return p.postJSON(context.Background(), "/stores/set", body, nil)
}
type storesDeleteRequestBody struct {
Store string `json:"store,omitempty"`
Keys [][]float32 `json:"keys"`
}
func (p *LocalAIProxy) StoresDelete(req *pb.StoresDeleteOptions) error {
body := storesDeleteRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())}
return p.postJSON(context.Background(), "/stores/delete", body, nil)
}
type storesGetRequestBody struct {
Store string `json:"store,omitempty"`
Keys [][]float32 `json:"keys"`
}
type storesGetResponseBody struct {
Keys [][]float32 `json:"keys"`
Values []string `json:"values"`
}
func (p *LocalAIProxy) StoresGet(req *pb.StoresGetOptions) (pb.StoresGetResult, error) {
body := storesGetRequestBody{Store: p.model(""), Keys: storeKeysToFloats(req.GetKeys())}
var resp storesGetResponseBody
if err := p.postJSON(context.Background(), "/stores/get", body, &resp); err != nil {
return pb.StoresGetResult{}, err
}
return pb.StoresGetResult{Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values)}, nil
}
type storesFindRequestBody struct {
Store string `json:"store,omitempty"`
Key []float32 `json:"key"`
Topk int `json:"topk,omitempty"`
}
type storesFindResponseBody struct {
Keys [][]float32 `json:"keys"`
Values []string `json:"values"`
Similarities []float32 `json:"similarities"`
}
func (p *LocalAIProxy) StoresFind(req *pb.StoresFindOptions) (pb.StoresFindResult, error) {
body := storesFindRequestBody{Store: p.model(""), Key: req.GetKey().GetFloats(), Topk: int(req.GetTopK())}
var resp storesFindResponseBody
if err := p.postJSON(context.Background(), "/stores/find", body, &resp); err != nil {
return pb.StoresFindResult{}, err
}
return pb.StoresFindResult{
Keys: storeFloatsToKeys(resp.Keys), Values: storeStringsToValues(resp.Values), Similarities: resp.Similarities,
}, nil
}
+525
View File
@@ -0,0 +1,525 @@
package main
import (
"encoding/base64"
"encoding/json"
"net/http"
"os"
"path/filepath"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"google.golang.org/grpc/codes"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
"github.com/mudler/LocalAI/pkg/utils"
)
// b64 is the base64 encoding a spec expects the proxy to have produced from
// a local file's content, matching what the REST endpoints take inline.
func b64(content string) string {
return base64.StdEncoding.EncodeToString([]byte(content))
}
var _ = Describe("media methods", func() {
var up *fakeUpstream
BeforeEach(func() {
up = newFakeUpstream()
DeferCleanup(up.Close)
})
Describe("GenerateImage", func() {
It("posts base64 src/ref_images and decodes data[0].b64_json into Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/images/generations", map[string]any{
"data": []map[string]any{{"b64_json": b64("png-bytes")}},
})
src := writeInput("src.png", "src-bytes")
ref := writeInput("ref.png", "ref-bytes")
dst := filepath.Join(GinkgoT().TempDir(), "out.png")
Expect(p.GenerateImage(&pb.GenerateImageRequest{
Width: 64, Height: 32, Step: 20, Seed: 7,
PositivePrompt: "a cat", NegativePrompt: "blurry",
Src: src, RefImages: []string{ref}, Dst: dst,
})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("png-bytes"))
req := up.last()
Expect(req.Path).To(Equal("/v1/images/generations"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
Expect(req.JSON).To(HaveKeyWithValue("prompt", "a cat"))
Expect(req.JSON).To(HaveKeyWithValue("negative_prompt", "blurry"))
Expect(req.JSON).To(HaveKeyWithValue("size", "64x32"))
Expect(req.JSON).To(HaveKeyWithValue("step", BeNumerically("==", 20)))
Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 7)))
Expect(req.JSON).To(HaveKeyWithValue("response_format", "b64_json"))
Expect(req.JSON).To(HaveKeyWithValue("file", b64("src-bytes")))
Expect(req.JSON).To(HaveKeyWithValue("ref_images", ConsistOf(b64("ref-bytes"))))
})
It("downloads a URL reply relative to the upstream into Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/images/generations", map[string]any{
"data": []map[string]any{{"url": up.URL + "/generated-images/out.png"}},
})
up.script("/generated-images/out.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "downloaded-bytes"})
dst := filepath.Join(GinkgoT().TempDir(), "out.png")
Expect(p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("downloaded-bytes"))
})
It("maps an upstream failure and leaves no partial file", func() {
p := loadProxy(up, nil)
up.script("/v1/images/generations", scriptedResponse{Status: http.StatusInternalServerError, Body: "boom"})
dst := filepath.Join(GinkgoT().TempDir(), "out.png")
err := p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "x", Dst: dst})
Expect(codeOf(err)).To(Equal(codes.Unavailable))
Expect(dst).NotTo(BeAnExistingFile())
})
It("downloads a reply URL from the configured upstream even when its host does not match", func() {
// A reverse proxy, LOCALAI_BASE_URL, or X-Forwarded-Host can all make
// the upstream hand back a URL on a different host than the one this
// proxy is configured with. Only the path should matter: the proxy
// must still fetch it from cfg.base (this fake upstream), with its
// own bearer key, never from the host named in the URL.
p := loadProxy(up, nil)
up.replyJSON("/v1/images/generations", map[string]any{
"data": []map[string]any{{"url": "http://mismatched-host.invalid:1/generated-images/out.png"}},
})
up.script("/generated-images/out.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "downloaded-bytes"})
dst := filepath.Join(GinkgoT().TempDir(), "out.png")
Expect(p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("downloaded-bytes"))
})
It("refuses a reply URL whose path is not a known generated-content path", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/images/generations", map[string]any{
"data": []map[string]any{{"url": up.URL + "/etc/passwd"}},
})
dst := filepath.Join(GinkgoT().TempDir(), "out.png")
err := p.GenerateImage(&pb.GenerateImageRequest{PositivePrompt: "a cat", Dst: dst})
Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
Expect(dst).NotTo(BeAnExistingFile())
})
})
Describe("UpscaleImage", func() {
It("uploads the image as multipart and downloads the returned URL into Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/images/upscale", map[string]any{
"data": []map[string]any{{"url": up.URL + "/generated-images/big.png"}},
})
up.script("/generated-images/big.png", scriptedResponse{Status: http.StatusOK, ContentType: "image/png", Body: "big-bytes"})
src := writeInput("small.png", "small-bytes")
dst := filepath.Join(GinkgoT().TempDir(), "big.png")
Expect(p.UpscaleImage(&pb.UpscaleImageRequest{Src: src, Dst: dst, Scale: 4})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("big-bytes"))
req := up.recorded()[0]
Expect(req.Path).To(Equal("/v1/images/upscale"))
Expect(req.Fields).To(HaveKeyWithValue("model", "remote-model"))
Expect(req.Fields).To(HaveKeyWithValue("scale", "4"))
Expect(req.Files).To(HaveKeyWithValue("image", "small-bytes"))
})
})
Describe("GenerateVideo", func() {
It("posts base64 media fields and decodes data[0].b64_json into Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/video", map[string]any{
"data": []map[string]any{{"b64_json": b64("mp4-bytes")}},
})
start := writeInput("start.png", "start-bytes")
dst := filepath.Join(GinkgoT().TempDir(), "out.mp4")
Expect(p.GenerateVideo(&pb.GenerateVideoRequest{
Prompt: "a dog running", NegativePrompt: "static", StartImage: start,
Width: 512, Height: 288, NumFrames: 24, Fps: 8, Seed: 3, CfgScale: 5.5, Step: 10,
Dst: dst, Params: map[string]string{"motion": "high"},
})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("mp4-bytes"))
req := up.last()
Expect(req.Path).To(Equal("/video"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
Expect(req.JSON).To(HaveKeyWithValue("prompt", "a dog running"))
Expect(req.JSON).To(HaveKeyWithValue("start_image", b64("start-bytes")))
Expect(req.JSON).To(HaveKeyWithValue("num_frames", BeNumerically("==", 24)))
Expect(req.JSON).To(HaveKeyWithValue("fps", BeNumerically("==", 8)))
Expect(req.JSON).To(HaveKeyWithValue("params", HaveKeyWithValue("motion", "high")))
})
})
Describe("Generate3D", func() {
It("posts the base64 conditioning image and decodes data[0].b64_json into Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/3d/generations", map[string]any{
"data": []map[string]any{{"b64_json": b64("glb-bytes")}},
})
src := writeInput("cond.png", "cond-bytes")
dst := filepath.Join(GinkgoT().TempDir(), "out.glb")
Expect(p.Generate3D(&pb.Generate3DRequest{
Src: src, Dst: dst, Seed: 1, Step: 12, CfgScale: 7.5, TextureSteps: 12, Quality: "auto", Background: "keep",
})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("glb-bytes"))
req := up.last()
Expect(req.Path).To(Equal("/3d/generations"))
Expect(req.JSON).To(HaveKeyWithValue("image", b64("cond-bytes")))
Expect(req.JSON).To(HaveKeyWithValue("quality", "auto"))
Expect(req.JSON).To(HaveKeyWithValue("background", "keep"))
})
})
Describe("Animate3D", func() {
It("base64-encodes non-text inputs, keeps text raw, and writes the reply to Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/3d/animate", map[string]any{
"data": []map[string]any{{"b64_json": b64("anim-bytes")}},
"metadata": map[string]any{"frames": float64(30)},
})
mesh := writeInput("mesh.glb", "mesh-bytes")
dst := filepath.Join(GinkgoT().TempDir(), "out.glb")
Expect(p.Animate3D(&pb.Animate3DRequest{
Inputs: map[string]*pb.AnimationInput{
"mesh": {Type: "mesh", Data: mesh},
"prompt": {Type: "text", Data: "wave hello"},
},
Dst: dst,
})).To(Succeed())
got, err := os.ReadFile(dst)
Expect(err).NotTo(HaveOccurred())
Expect(string(got)).To(Equal("anim-bytes"))
req := up.last()
Expect(req.Path).To(Equal("/3d/animate"))
Expect(req.JSON["inputs"]).To(HaveKeyWithValue("mesh", map[string]any{"type": "mesh", "data": b64("mesh-bytes")}))
Expect(req.JSON["inputs"]).To(HaveKeyWithValue("prompt", map[string]any{"type": "text", "data": "wave hello"}))
})
It("Animate3DWithMetadata returns the upstream metadata and writes Dst", func() {
p := loadProxy(up, nil)
up.replyJSON("/3d/animate", map[string]any{
"data": []map[string]any{{"b64_json": b64("anim-bytes")}},
"metadata": map[string]any{"frames": float64(30)},
})
dst := filepath.Join(GinkgoT().TempDir(), "out.glb")
metadata, err := p.Animate3DWithMetadata(&pb.Animate3DRequest{
Inputs: map[string]*pb.AnimationInput{"prompt": {Type: "text", Data: "wave"}},
Dst: dst,
})
Expect(err).NotTo(HaveOccurred())
Expect(string(metadata)).To(MatchJSON(`{"frames":30}`))
Expect(os.ReadFile(dst)).To(BeEquivalentTo("anim-bytes"))
})
})
// pngBytes is a minimal payload with a real PNG magic number, so
// http.DetectContentType (which toDataURI uses) sniffs "image/png" the
// way it would for a real image, and decodeAndCheckImage below can tell
// this test wrote a valid data URI apart from one that merely echoes bare
// base64.
pngBytes := append([]byte("\x89PNG\r\n\x1a\n"), []byte("fake-png-body")...)
// decodeAndCheckImage extracts field from an upstream request body and
// decodes it exactly the way the real REST handlers do — with
// utils.GetContentURIAsBase64 (core/http/endpoints/localai/images.go's
// decodeImageInput calls the same function) — so a spec here fails the
// same way a live upstream would if the proxy ever regressed to sending
// bare base64, which that function rejects outright.
decodeAndCheckImage := func(body map[string]any, field string, want []byte) {
raw, _ := body[field].(string)
decoded, err := utils.GetContentURIAsBase64(raw)
ExpectWithOffset(1, err).NotTo(HaveOccurred(), "the real upstream would reject %s the same way", field)
got, err := base64.StdEncoding.DecodeString(decoded)
ExpectWithOffset(1, err).NotTo(HaveOccurred())
ExpectWithOffset(1, got).To(Equal(want))
}
Describe("Detect", func() {
It("wraps the bare base64 image as a data URI the real upstream decoder accepts, and maps detections", func() {
up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
decodeAndCheckImage(body, "image", pngBytes)
Expect(body["prompt"]).To(Equal("cat"))
mask := base64.StdEncoding.EncodeToString([]byte("png-mask"))
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"detections": []map[string]any{
{"x": 1.0, "y": 2.0, "width": 3.0, "height": 4.0, "confidence": 0.9, "class_name": "cat", "mask": mask},
},
})
})
DeferCleanup(up.Close)
p := loadProxy(up, nil)
res, err := p.Detect(&pb.DetectOptions{
Src: base64.StdEncoding.EncodeToString(pngBytes), Prompt: "cat", Threshold: 0.5,
})
Expect(err).NotTo(HaveOccurred())
Expect(res.Detections).To(HaveLen(1))
Expect(res.Detections[0].ClassName).To(Equal("cat"))
Expect(res.Detections[0].Confidence).To(BeNumerically("~", 0.9, 1e-6))
Expect(res.Detections[0].Mask).To(Equal([]byte("png-mask")))
})
})
Describe("Depth", func() {
It("wraps the bare base64 image as a data URI and maps the full depth response", func() {
up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
decodeAndCheckImage(body, "image", pngBytes)
Expect(body["include_depth"]).To(Equal(true))
colors := base64.StdEncoding.EncodeToString([]byte("rgb"))
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"width": 2, "height": 1, "depth": []float64{0.1, 0.2},
"point_colors": colors, "is_metric": true,
})
})
DeferCleanup(up.Close)
p := loadProxy(up, nil)
res, err := p.Depth(&pb.DepthRequest{Src: base64.StdEncoding.EncodeToString(pngBytes), IncludeDepth: true})
Expect(err).NotTo(HaveOccurred())
Expect(res.Width).To(Equal(int32(2)))
Expect(res.Height).To(Equal(int32(1)))
Expect(res.Depth).To(Equal([]float32{0.1, 0.2}))
Expect(res.PointColors).To(Equal([]byte("rgb")))
Expect(res.IsMetric).To(BeTrue())
})
It("refuses a request for exports or a dst directory without calling the upstream, so failover moves on", func() {
p := loadProxy(up, nil)
_, exportsErr := p.Depth(&pb.DepthRequest{Src: "irrelevant", Exports: []string{"glb"}})
Expect(codeOf(exportsErr)).To(Equal(codes.Unimplemented))
_, dstErr := p.Depth(&pb.DepthRequest{Src: "irrelevant", Dst: "/some/output/dir"})
Expect(codeOf(dstErr)).To(Equal(codes.Unimplemented))
Expect(up.recorded()).To(BeEmpty(), "an exports/dst request must never reach the upstream")
})
})
Describe("FaceVerify", func() {
It("wraps both bare base64 images as data URIs and maps the response, including liveness fields", func() {
img2Bytes := append([]byte("\x89PNG\r\n\x1a\n"), []byte("other-face")...)
up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
decodeAndCheckImage(body, "img1", pngBytes)
decodeAndCheckImage(body, "img2", img2Bytes)
Expect(body["anti_spoofing"]).To(Equal(true))
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"verified": true, "distance": 0.1, "threshold": 0.4, "confidence": 92.0, "model": "buffalo_l",
"img1_area": map[string]any{"x": 1.0, "y": 2.0, "w": 3.0, "h": 4.0},
"img2_area": map[string]any{"x": 5.0, "y": 6.0, "w": 7.0, "h": 8.0},
"img1_is_real": true, "img1_antispoof_score": 0.99,
"img2_is_real": false, "img2_antispoof_score": 0.1,
})
})
DeferCleanup(up.Close)
p := loadProxy(up, nil)
res, err := p.FaceVerify(&pb.FaceVerifyRequest{
Img1: base64.StdEncoding.EncodeToString(pngBytes), Img2: base64.StdEncoding.EncodeToString(img2Bytes),
AntiSpoofing: true,
})
Expect(err).NotTo(HaveOccurred())
Expect(res.Verified).To(BeTrue())
Expect(res.Model).To(Equal("buffalo_l"))
Expect(res.Img1Area.W).To(BeNumerically("==", 3))
Expect(res.Img1IsReal).To(BeTrue())
Expect(res.Img2IsReal).To(BeFalse())
})
})
Describe("FaceAnalyze", func() {
It("wraps the bare base64 image as a data URI and maps per-face demographic attributes", func() {
up = newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
decodeAndCheckImage(body, "img", pngBytes)
Expect(body["actions"]).To(ConsistOf("age", "gender"))
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"faces": []map[string]any{
{
"region": map[string]any{"x": 1.0, "y": 2.0, "w": 3.0, "h": 4.0},
"face_confidence": 0.95, "age": 30.0, "dominant_gender": "Man",
"gender": map[string]any{"Man": 0.9, "Woman": 0.1},
},
},
})
})
DeferCleanup(up.Close)
p := loadProxy(up, nil)
res, err := p.FaceAnalyze(&pb.FaceAnalyzeRequest{
Img: base64.StdEncoding.EncodeToString(pngBytes), Actions: []string{"age", "gender"},
})
Expect(err).NotTo(HaveOccurred())
Expect(res.Faces).To(HaveLen(1))
Expect(res.Faces[0].DominantGender).To(Equal("Man"))
Expect(res.Faces[0].Gender).To(HaveKeyWithValue("Man", Equal(float32(0.9))))
})
})
Describe("VoiceVerify", func() {
It("base64-encodes both audio clips and maps the response", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/voice/verify", map[string]any{
"verified": true, "distance": 0.2, "threshold": 0.5, "confidence": 88.0, "model": "ecapa",
})
a1 := writeInput("a1.wav", "audio-one")
a2 := writeInput("a2.wav", "audio-two")
res, err := p.VoiceVerify(&pb.VoiceVerifyRequest{Audio1: a1, Audio2: a2, Threshold: 0.5})
Expect(err).NotTo(HaveOccurred())
Expect(res.Verified).To(BeTrue())
Expect(res.Model).To(Equal("ecapa"))
req := up.last()
Expect(req.Path).To(Equal("/v1/voice/verify"))
Expect(req.JSON).To(HaveKeyWithValue("audio1", b64("audio-one")))
Expect(req.JSON).To(HaveKeyWithValue("audio2", b64("audio-two")))
})
})
Describe("VoiceAnalyze", func() {
It("base64-encodes the audio clip and maps demographic segments", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/voice/analyze", map[string]any{
"segments": []map[string]any{
{"start": 0.0, "end": 1.5, "age": 25.0, "dominant_gender": "Woman"},
},
})
audio := writeInput("clip.wav", "clip-bytes")
res, err := p.VoiceAnalyze(&pb.VoiceAnalyzeRequest{Audio: audio, Actions: []string{"age"}})
Expect(err).NotTo(HaveOccurred())
Expect(res.Segments).To(HaveLen(1))
Expect(res.Segments[0].DominantGender).To(Equal("Woman"))
req := up.last()
Expect(req.Path).To(Equal("/v1/voice/analyze"))
Expect(req.JSON).To(HaveKeyWithValue("audio", b64("clip-bytes")))
})
})
Describe("VoiceEmbed", func() {
It("base64-encodes the audio clip and returns the embedding", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/voice/embed", map[string]any{
"embedding": []float64{0.1, 0.2, 0.3}, "dim": 3.0, "model": "ecapa",
})
audio := writeInput("clip.wav", "clip-bytes")
res, err := p.VoiceEmbed(&pb.VoiceEmbedRequest{Audio: audio})
Expect(err).NotTo(HaveOccurred())
Expect(res.Embedding).To(Equal([]float32{0.1, 0.2, 0.3}))
Expect(res.Model).To(Equal("ecapa"))
req := up.last()
Expect(req.Path).To(Equal("/v1/voice/embed"))
Expect(req.JSON).To(HaveKeyWithValue("audio", b64("clip-bytes")))
})
})
Describe("Stores", func() {
It("StoresSet posts the store name, keys and values", func() {
p := loadProxy(up, nil)
up.script("/stores/set", scriptedResponse{Status: http.StatusOK})
Expect(p.StoresSet(&pb.StoresSetOptions{
Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}},
Values: []*pb.StoresValue{{Bytes: []byte("v1")}},
})).To(Succeed())
req := up.last()
Expect(req.Path).To(Equal("/stores/set"))
Expect(req.JSON).To(HaveKeyWithValue("store", "remote-model"))
Expect(req.JSON).To(HaveKeyWithValue("values", ConsistOf("v1")))
})
It("StoresDelete posts the store name and keys", func() {
p := loadProxy(up, nil)
up.script("/stores/delete", scriptedResponse{Status: http.StatusOK})
Expect(p.StoresDelete(&pb.StoresDeleteOptions{Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}}})).To(Succeed())
req := up.last()
Expect(req.Path).To(Equal("/stores/delete"))
Expect(req.JSON).To(HaveKeyWithValue("store", "remote-model"))
})
It("StoresGet posts keys and maps the returned keys/values", func() {
p := loadProxy(up, nil)
up.replyJSON("/stores/get", map[string]any{
"keys": [][]float64{{1, 2}}, "values": []string{"v1"},
})
res, err := p.StoresGet(&pb.StoresGetOptions{Keys: []*pb.StoresKey{{Floats: []float32{1, 2}}}})
Expect(err).NotTo(HaveOccurred())
Expect(res.Keys).To(HaveLen(1))
Expect(res.Keys[0].Floats).To(Equal([]float32{1, 2}))
Expect(res.Values).To(HaveLen(1))
Expect(res.Values[0].Bytes).To(Equal([]byte("v1")))
})
It("StoresFind posts the query key/topk and maps keys/values/similarities", func() {
p := loadProxy(up, nil)
up.replyJSON("/stores/find", map[string]any{
"keys": [][]float64{{1, 2}}, "values": []string{"v1"}, "similarities": []float64{0.9},
})
res, err := p.StoresFind(&pb.StoresFindOptions{Key: &pb.StoresKey{Floats: []float32{1, 2}}, TopK: 5})
Expect(err).NotTo(HaveOccurred())
Expect(res.Similarities).To(Equal([]float32{0.9}))
req := up.last()
Expect(req.Path).To(Equal("/stores/find"))
Expect(req.JSON).To(HaveKeyWithValue("topk", BeNumerically("==", 5)))
Expect(req.JSON).To(HaveKeyWithValue("key", ConsistOf(BeNumerically("==", 1), BeNumerically("==", 2))))
})
})
})
+13
View File
@@ -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/
+239
View File
@@ -0,0 +1,239 @@
package main
import (
"context"
"errors"
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"sync/atomic"
"time"
"github.com/mudler/xlog"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/mudler/LocalAI/pkg/grpc/base"
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
"github.com/mudler/LocalAI/pkg/httpclient"
)
const (
backendName = "localai-proxy"
// realtimePipelineOption names the upstream realtime pipeline that serves
// live transcription sessions (options: ["realtime_pipeline:<name>"]).
realtimePipelineOption = "realtime_pipeline:"
)
// LocalAIProxy serves backend methods by calling a remote LocalAI's REST API.
// base.SingleThread is not embedded: every call is an independent HTTP
// request, so serialising them would only add latency.
type LocalAIProxy struct {
base.Base
cfg atomic.Pointer[proxyConfig]
client *http.Client
}
type proxyConfig struct {
base string // upstream base URL without a trailing slash
upstreamModel string // model name sent upstream
apiKey string
realtimePipeline string
timeout time.Duration // per-request limit for non-streaming calls; 0 = none
}
func NewLocalAIProxy() *LocalAIProxy {
// httpclient.New refuses redirects: the upstream is one configured
// LocalAI, so a 3xx means misconfiguration or a hijacked host, and
// following it would replay the bearer key to an unvetted host. It also
// sets no body deadline, so long SSE streams are not cut short.
return &LocalAIProxy{client: httpclient.New()}
}
// Load refuses a model without proxy options so greedy backend probing,
// which tries every installed backend on a model file, never selects it.
func (p *LocalAIProxy) Load(opts *pb.ModelOptions) error {
po := opts.GetProxy()
if po == nil {
return errors.New("localai-proxy: Load requires proxy options (proxy.upstream_url)")
}
raw := po.GetUpstreamUrl()
if raw == "" {
return errors.New("localai-proxy: proxy.upstream_url is required")
}
u, err := url.ParseRequestURI(raw)
if err != nil {
return fmt.Errorf("localai-proxy: proxy.upstream_url %q invalid: %w", raw, err)
}
if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
return fmt.Errorf("localai-proxy: proxy.upstream_url %q must be an http(s) URL with a host", raw)
}
// Every request path starts with /v1, so upstream_url is the server's
// root. A URL copied from an OpenAI-style config ends in /v1 (or a full
// endpoint): cut the path at /v1 as the failover prober does, or the
// prober reports the target healthy while every request 404s.
base := strings.TrimRight(raw, "/")
if i := strings.Index(u.Path, "/v1"); i >= 0 {
base = strings.TrimRight(u.Scheme+"://"+u.Host+u.Path[:i], "/")
xlog.Warn("localai-proxy: proxy.upstream_url should be the server root; ignoring its /v1 path",
"upstream_url", raw, "using", base)
}
// There is no translate mode: the upstream always speaks LocalAI's API.
if po.GetMode() != "" || po.GetProvider() != "" {
xlog.Warn("localai-proxy: proxy.mode and proxy.provider are ignored",
"mode", po.GetMode(), "provider", po.GetProvider())
}
key, err := resolveAPIKey(po.GetApiKeyEnv(), po.GetApiKeyFile())
if err != nil {
return err
}
model := po.GetUpstreamModel()
if model == "" {
model = opts.GetModel()
}
if model == "" {
xlog.Warn("localai-proxy: no upstream model name; set proxy.upstream_model")
}
var pipeline string
for _, o := range opts.GetOptions() {
if v, ok := strings.CutPrefix(o, realtimePipelineOption); ok {
pipeline = strings.TrimSpace(v)
}
}
var timeout time.Duration
if s := po.GetRequestTimeoutSeconds(); s > 0 {
timeout = time.Duration(s) * time.Second
}
p.cfg.Store(&proxyConfig{
base: base,
upstreamModel: model,
apiKey: key,
realtimePipeline: pipeline,
timeout: timeout,
})
xlog.Info("localai-proxy: ready", "upstream", base, "upstream_model", model,
"has_key", key != "", "realtime_pipeline", pipeline)
return nil
}
// config returns the loaded configuration, or the typed not-loaded error so
// callers see FailedPrecondition instead of a nil dereference.
func (p *LocalAIProxy) config() (*proxyConfig, error) {
cfg := p.cfg.Load()
if cfg == nil {
return nil, grpcerrors.ModelNotLoaded(backendName)
}
return cfg, nil
}
// model returns the model name to send upstream. The configured name wins so
// every method targets the same upstream model; req (a model named by the
// request itself) is only a fallback for configs that resolved no name.
func (p *LocalAIProxy) model(req string) string {
if cfg := p.cfg.Load(); cfg != nil && cfg.upstreamModel != "" {
return cfg.upstreamModel
}
return req
}
// resolveAPIKey mirrors config.ProxyConfig.ResolveAPIKey (and cloud-proxy's
// copy). Duplicated so the backend binary does not depend on core's layout.
func resolveAPIKey(envName, filePath string) (string, error) {
if envName != "" {
v := os.Getenv(envName)
if v == "" {
return "", fmt.Errorf("localai-proxy: api_key_env %q is unset", envName)
}
return v, nil
}
if filePath != "" {
// #nosec G304 -- api_key_file comes from the operator's model config (passed by core as a backend option), not from a request
b, err := os.ReadFile(filepath.Clean(filePath))
if err != nil {
return "", fmt.Errorf("localai-proxy: read api_key_file %q: %w", filePath, err)
}
return strings.TrimSpace(string(b)), nil
}
return "", nil
}
// unimplemented is the error for methods LocalAI's REST API cannot serve.
// Failover reads gRPC Unimplemented as a capability gap and moves to the next
// target without marking this one unhealthy.
func unimplemented(method string) error {
return status.Errorf(codes.Unimplemented, "localai-proxy: %s has no upstream counterpart", method)
}
func (p *LocalAIProxy) AudioEncode(*pb.AudioEncodeRequest) (*pb.AudioEncodeResult, error) {
return nil, unimplemented("AudioEncode")
}
func (p *LocalAIProxy) AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error) {
return nil, unimplemented("AudioDecode")
}
// AudioToAudioStream closes out because the gRPC server drains it until
// closed; leaving it open would hang the call.
func (p *LocalAIProxy) AudioToAudioStream(_ <-chan *pb.AudioToAudioRequest, out chan<- *pb.AudioToAudioResponse) error {
close(out)
return unimplemented("AudioToAudioStream")
}
func (p *LocalAIProxy) TokenClassify(context.Context, *pb.TokenClassifyRequest) (*pb.TokenClassifyResponse, error) {
return nil, unimplemented("TokenClassify")
}
func (p *LocalAIProxy) ModelMetadata(*pb.ModelOptions) (*pb.ModelMetadataResponse, error) {
return nil, unimplemented("ModelMetadata")
}
func (p *LocalAIProxy) StartFineTune(*pb.FineTuneRequest) (*pb.FineTuneJobResult, error) {
return nil, unimplemented("StartFineTune")
}
// FineTuneProgress closes the channel: the gRPC server waits for it to close
// before returning, and base.Base leaves it open.
func (p *LocalAIProxy) FineTuneProgress(_ *pb.FineTuneProgressRequest, updates chan *pb.FineTuneProgressUpdate) error {
close(updates)
return unimplemented("FineTuneProgress")
}
func (p *LocalAIProxy) StopFineTune(*pb.FineTuneStopRequest) error {
return unimplemented("StopFineTune")
}
func (p *LocalAIProxy) ListCheckpoints(*pb.ListCheckpointsRequest) (*pb.ListCheckpointsResponse, error) {
return nil, unimplemented("ListCheckpoints")
}
func (p *LocalAIProxy) ExportModel(*pb.ExportModelRequest) error {
return unimplemented("ExportModel")
}
func (p *LocalAIProxy) StartQuantization(*pb.QuantizationRequest) (*pb.QuantizationJobResult, error) {
return nil, unimplemented("StartQuantization")
}
// QuantizationProgress closes the channel for the same reason as
// FineTuneProgress.
func (p *LocalAIProxy) QuantizationProgress(_ *pb.QuantizationProgressRequest, updates chan *pb.QuantizationProgressUpdate) error {
close(updates)
return unimplemented("QuantizationProgress")
}
func (p *LocalAIProxy) StopQuantization(*pb.QuantizationStopRequest) error {
return unimplemented("StopQuantization")
}
+6
View File
@@ -0,0 +1,6 @@
#!/bin/bash
set -ex
CURDIR=$(dirname "$(realpath "$0")")
exec "$CURDIR"/localai-proxy "$@"
+447
View File
@@ -0,0 +1,447 @@
package main
import (
"bufio"
"context"
"encoding/json"
"math"
"strings"
"github.com/mudler/xlog"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
// textRequest is the body for /v1/chat/completions (Messages) and
// /v1/completions (Prompt). Zero sampling values are omitted so the upstream
// model's own config defaults apply, as they would for a direct caller.
// Temperature is the exception: 0 is a real choice (greedy decoding), and
// core always fills it from the model config, so it is always sent.
type textRequest struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages,omitempty"`
Prompt string `json:"prompt,omitempty"`
Stream bool `json:"stream,omitempty"`
MaxTokens int32 `json:"max_tokens,omitempty"`
Temperature float32 `json:"temperature"`
TopP float32 `json:"top_p,omitempty"`
TopK int32 `json:"top_k,omitempty"`
Seed int32 `json:"seed,omitempty"`
Stop []string `json:"stop,omitempty"`
Tools json.RawMessage `json:"tools,omitempty"`
ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
// StreamOptions is set on streamed requests: the upstream sends the
// usage trailer only when include_usage asks for it.
StreamOptions *streamOptions `json:"stream_options,omitempty"`
}
type streamOptions struct {
IncludeUsage bool `json:"include_usage"`
}
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
Name string `json:"name,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
ToolCalls []toolCall `json:"tool_calls,omitempty"`
}
type toolCall struct {
Index int `json:"index"`
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
} `json:"function"`
}
// textChoice covers both endpoints and both shapes: chat replies fill Message
// (or Delta when streaming), completions fill Text.
type textChoice struct {
Text string `json:"text"`
Message choiceDelta `json:"message"`
Delta choiceDelta `json:"delta"`
}
type choiceDelta struct {
Content string `json:"content"`
// LocalAI names the field "reasoning"; other OpenAI-compatible servers
// use "reasoning_content".
Reasoning string `json:"reasoning"`
ReasoningContent string `json:"reasoning_content"`
ToolCalls []toolCall `json:"tool_calls"`
}
type textResponse struct {
// Error is set on the frame LocalAI sends when generation fails after
// the stream has started (followed by [DONE]).
Error *struct {
Message string `json:"message"`
} `json:"error"`
Choices []textChoice `json:"choices"`
Usage *struct {
PromptTokens int32 `json:"prompt_tokens"`
CompletionTokens int32 `json:"completion_tokens"`
} `json:"usage"`
}
// textRequest picks the chat endpoint when core sent structured messages
// (the model uses the tokenizer template, so the upstream must template
// too); otherwise core already rendered the prompt and completions takes it
// verbatim.
func (p *LocalAIProxy) textRequest(opts *pb.PredictOptions, stream bool) (string, textRequest) {
if dropped := unforwardedFields(opts); len(dropped) > 0 {
xlog.Warn("localai-proxy: request fields are not forwarded upstream", "fields", dropped)
}
req := textRequest{
Model: p.model(""),
Stream: stream,
MaxTokens: opts.GetTokens(),
Temperature: opts.GetTemperature(),
TopP: opts.GetTopP(),
TopK: opts.GetTopK(),
Seed: opts.GetSeed(),
Stop: opts.GetStopPrompts(),
Tools: rawJSON(opts.GetTools()),
ToolChoice: rawJSON(opts.GetToolChoice()),
}
if stream {
req.StreamOptions = &streamOptions{IncludeUsage: true}
}
if len(opts.GetMessages()) == 0 {
req.Prompt = opts.GetPrompt()
return "/v1/completions", req
}
for _, m := range opts.GetMessages() {
msg := chatMessage{
Role: m.GetRole(),
Content: m.GetContent(),
Name: m.GetName(),
ToolCallID: m.GetToolCallId(),
}
// A previous assistant turn carries its tool calls as a JSON string.
if tc := m.GetToolCalls(); tc != "" {
if err := json.Unmarshal([]byte(tc), &msg.ToolCalls); err != nil {
xlog.Debug("localai-proxy: drop malformed tool_calls on message", "error", err)
}
}
req.Messages = append(req.Messages, msg)
}
return "/v1/chat/completions", req
}
// unforwardedFields names the request inputs the REST text endpoints cannot
// carry from here: a grammar core compiled locally, and media that core hands
// over as local paths or base64 outside the messages. They are dropped, so
// say so instead of letting the answer silently ignore them.
func unforwardedFields(opts *pb.PredictOptions) []string {
var out []string
if opts.GetGrammar() != "" {
out = append(out, "grammar")
}
if len(opts.GetImages()) > 0 {
out = append(out, "images")
}
if len(opts.GetAudios()) > 0 {
out = append(out, "audios")
}
if len(opts.GetVideos()) > 0 {
out = append(out, "videos")
}
return out
}
// rawJSON passes a JSON string through untouched, or omits it when it is
// empty or invalid rather than failing the whole request.
func rawJSON(s string) json.RawMessage {
if s == "" || !json.Valid([]byte(s)) {
return nil
}
return json.RawMessage(s)
}
// replyFromChoice builds the Reply for one choice. Reasoning and tool calls
// travel as ChatDeltas, the same shape the llama.cpp autoparser emits, so
// core handles them without parsing the text again.
func replyFromChoice(c textChoice, streaming bool) *pb.Reply {
d := c.Message
if streaming {
d = c.Delta
}
content := d.Content
if content == "" {
content = c.Text
}
reasoning := d.Reasoning
if reasoning == "" {
reasoning = d.ReasoningContent
}
reply := &pb.Reply{Message: []byte(content)}
// Streaming chunks always carry a delta, like the autoparser's. A
// complete reply only needs one for reasoning or tool calls; its
// content must then be in the delta too, because core takes the
// content from the deltas once they carry anything.
if reasoning == "" && len(d.ToolCalls) == 0 && (!streaming || content == "") {
return reply
}
delta := &pb.ChatDelta{Content: content, ReasoningContent: reasoning}
for _, tc := range d.ToolCalls {
delta.ToolCalls = append(delta.ToolCalls, &pb.ToolCallDelta{
Index: clampInt32(tc.Index),
Id: tc.ID,
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
})
}
reply.ChatDeltas = []*pb.ChatDelta{delta}
return reply
}
func (p *LocalAIProxy) PredictRich(opts *pb.PredictOptions) (*pb.Reply, error) {
return p.PredictRichContext(context.Background(), opts)
}
// PredictRichContext is PredictRich bound to the gRPC call: when the caller
// goes away, the upstream request is cancelled too.
func (p *LocalAIProxy) PredictRichContext(ctx context.Context, opts *pb.PredictOptions) (*pb.Reply, error) {
path, body := p.textRequest(opts, false)
var resp textResponse
if err := p.postJSON(ctx, path, body, &resp); err != nil {
return nil, err
}
if len(resp.Choices) == 0 {
return nil, status.Errorf(codes.Internal, "localai-proxy: upstream %s returned no choices", path)
}
reply := replyFromChoice(resp.Choices[0], false)
if resp.Usage != nil {
reply.PromptTokens = resp.Usage.PromptTokens
reply.Tokens = resp.Usage.CompletionTokens
}
return reply, nil
}
// PredictStreamRich sends one Reply per upstream SSE delta. It does not close
// results: the gRPC server does, after this returns.
func (p *LocalAIProxy) PredictStreamRich(opts *pb.PredictOptions, results chan<- *pb.Reply) error {
return p.PredictStreamRichContext(context.Background(), opts, results)
}
// PredictStreamRichContext is PredictStreamRich bound to the gRPC stream: a
// client that disconnects, or a failover that abandons this target, stops the
// upstream generation instead of letting it run (and bill) to the end.
func (p *LocalAIProxy) PredictStreamRichContext(ctx context.Context, opts *pb.PredictOptions, results chan<- *pb.Reply) error {
path, body := p.textRequest(opts, true)
resp, err := p.postStream(ctx, path, body)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
scanner := bufio.NewScanner(resp.Body)
// A single frame can carry a long tool-call argument chunk.
scanner.Buffer(make([]byte, 0, 64*1024), 4<<20)
for scanner.Scan() {
payload, ok := strings.CutPrefix(scanner.Text(), "data:")
if !ok {
continue
}
payload = strings.TrimSpace(payload)
if payload == "[DONE]" {
return nil
}
if payload == "" {
continue
}
var chunk textResponse
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
xlog.Debug("localai-proxy: skip malformed SSE frame", "path", path, "error", err)
continue
}
if chunk.Error != nil {
// The upstream failed mid-generation. Returning nil would turn a
// cut-off answer into a success; Unavailable lets failover and
// the client see the failure.
xlog.Warn("localai-proxy: upstream stream error", "path", path, "error", chunk.Error.Message)
return status.Errorf(codes.Unavailable, "localai-proxy: upstream %s stream failed: %s", path, chunk.Error.Message)
}
if chunk.Usage != nil && len(chunk.Choices) == 0 {
results <- &pb.Reply{PromptTokens: chunk.Usage.PromptTokens, Tokens: chunk.Usage.CompletionTokens}
continue
}
for _, c := range chunk.Choices {
reply := replyFromChoice(c, true)
if len(reply.GetMessage()) == 0 && len(reply.GetChatDeltas()) == 0 {
continue // role-only or finish frames carry nothing to emit
}
results <- reply
}
}
if err := scanner.Err(); err != nil {
return transportError(path, err)
}
return nil
}
// Predict is the legacy string path; the gRPC server prefers PredictRich.
func (p *LocalAIProxy) Predict(opts *pb.PredictOptions) (string, error) {
reply, err := p.PredictRich(opts)
if err != nil {
return "", err
}
return string(reply.GetMessage()), nil
}
// PredictStream is the legacy string stream. Unlike PredictStreamRich it
// owns and closes results, per the AIModel contract.
func (p *LocalAIProxy) PredictStream(opts *pb.PredictOptions, results chan string) error {
defer close(results)
rich := make(chan *pb.Reply)
errCh := make(chan error, 1)
go func() {
errCh <- p.PredictStreamRich(opts, rich)
close(rich)
}()
for reply := range rich {
if msg := reply.GetMessage(); len(msg) > 0 {
results <- string(msg)
}
}
return <-errCh
}
func (p *LocalAIProxy) Embeddings(opts *pb.PredictOptions) ([]float32, error) {
var resp struct {
Data []struct {
Embedding []float32 `json:"embedding"`
} `json:"data"`
}
body := map[string]any{"model": p.model(""), "input": opts.GetEmbeddings()}
// Core sends tokenized input in EmbeddingTokens and leaves Embeddings
// empty; a list of token lists is how the REST API takes tokens.
if tokens := opts.GetEmbeddingTokens(); len(tokens) > 0 {
body["input"] = [][]int32{tokens}
}
if err := p.postJSON(context.Background(), "/v1/embeddings", body, &resp); err != nil {
return nil, err
}
if len(resp.Data) == 0 || len(resp.Data[0].Embedding) == 0 {
return nil, status.Error(codes.Internal, "localai-proxy: upstream /v1/embeddings returned no embedding")
}
return resp.Data[0].Embedding, nil
}
func (p *LocalAIProxy) Rerank(ctx context.Context, in *pb.RerankRequest) (*pb.RerankResult, error) {
body := map[string]any{
"model": p.model(""),
"query": in.GetQuery(),
"documents": in.GetDocuments(),
}
// TopN 0 means "score every document" (the router's reranker sends it);
// upstream rejects top_n < 1, and an absent top_n means the same thing.
if n := in.GetTopN(); n > 0 {
body["top_n"] = n
}
var resp struct {
Usage struct {
TotalTokens int32 `json:"total_tokens"`
PromptTokens int32 `json:"prompt_tokens"`
} `json:"usage"`
Results []struct {
Index int32 `json:"index"`
Document struct {
Text string `json:"text"`
} `json:"document"`
RelevanceScore float32 `json:"relevance_score"`
} `json:"results"`
}
if err := p.postJSON(ctx, "/v1/rerank", body, &resp); err != nil {
return nil, err
}
out := &pb.RerankResult{Usage: &pb.Usage{TotalTokens: resp.Usage.TotalTokens, PromptTokens: resp.Usage.PromptTokens}}
for _, r := range resp.Results {
out.Results = append(out.Results, &pb.DocumentResult{Index: r.Index, Text: r.Document.Text, RelevanceScore: r.RelevanceScore})
}
return out, nil
}
// TokenizeString uses the upstream model's tokenizer, which is the one that
// matters: the upstream is where the tokens will be spent.
func (p *LocalAIProxy) TokenizeString(opts *pb.PredictOptions) (pb.TokenizationResponse, error) {
var resp struct {
Tokens []int32 `json:"tokens"`
}
body := map[string]any{"model": p.model(""), "content": opts.GetPrompt()}
if err := p.postJSON(context.Background(), "/v1/tokenize", body, &resp); err != nil {
return pb.TokenizationResponse{}, err
}
return pb.TokenizationResponse{Length: clampInt32(len(resp.Tokens)), Tokens: resp.Tokens}, nil
}
func (p *LocalAIProxy) Detokenize(in *pb.DetokenizeRequest) (pb.DetokenizeResponse, error) {
var resp struct {
Content string `json:"content"`
}
body := map[string]any{"model": p.model(""), "tokens": in.GetTokens()}
if err := p.postJSON(context.Background(), "/v1/detokenize", body, &resp); err != nil {
return pb.DetokenizeResponse{}, err
}
return pb.DetokenizeResponse{Content: resp.Content}, nil
}
// Score forwards plain candidate scoring to /api/score. That endpoint has no
// decision-pipeline fields, so a question_type request would silently become
// plain scoring upstream; refuse it as a capability gap instead.
func (p *LocalAIProxy) Score(ctx context.Context, in *pb.ScoreRequest) (*pb.ScoreResponse, error) {
if in.GetQuestionType() != "" {
return nil, unimplemented("Score with question_type")
}
body := map[string]any{
"model": p.model(""),
"prompt": in.GetPrompt(),
"candidates": in.GetCandidates(),
"include_token_logprobs": in.GetIncludeTokenLogprobs(),
"length_normalize": in.GetLengthNormalize(),
}
var resp struct {
Candidates []struct {
LogProb float64 `json:"log_prob"`
LengthNormalizedLogProb float64 `json:"length_normalized_log_prob"`
NumTokens int32 `json:"num_tokens"`
Tokens []struct {
Token string `json:"token"`
LogProb float64 `json:"log_prob"`
} `json:"tokens"`
} `json:"candidates"`
}
if err := p.postJSON(ctx, "/api/score", body, &resp); err != nil {
return nil, err
}
out := &pb.ScoreResponse{}
for _, c := range resp.Candidates {
cs := &pb.CandidateScore{LogProb: c.LogProb, LengthNormalizedLogProb: c.LengthNormalizedLogProb, NumTokens: c.NumTokens}
for _, t := range c.Tokens {
cs.Tokens = append(cs.Tokens, &pb.TokenLogProb{Token: t.Token, LogProb: t.LogProb})
}
out.Candidates = append(out.Candidates, cs)
}
return out, nil
}
// clampInt32 narrows an upstream-supplied int (token counts, tool-call
// indexes) to the int32 the gRPC protocol carries. The upstream is another
// server, so an absurd value must saturate rather than wrap to a negative or
// small number that core would take at face value.
func clampInt32(n int) int32 {
switch {
case n > math.MaxInt32:
return math.MaxInt32
case n < math.MinInt32:
return math.MinInt32
}
return int32(n)
}
+629
View File
@@ -0,0 +1,629 @@
package main
import (
"context"
"math"
"net/http"
"os"
"path/filepath"
"time"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
grpc "github.com/mudler/LocalAI/pkg/grpc"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
func codeOf(err error) codes.Code {
st, ok := status.FromError(err)
if !ok {
return codes.Unknown
}
return st.Code()
}
var _ = Describe("localai-proxy", func() {
var up *fakeUpstream
BeforeEach(func() {
up = newFakeUpstream()
DeferCleanup(up.Close)
})
Describe("Load", func() {
It("refuses a model without proxy options", func() {
err := NewLocalAIProxy().Load(&pb.ModelOptions{Model: "m"})
Expect(err).To(MatchError(ContainSubstring("proxy")))
})
It("refuses a missing or invalid upstream_url", func() {
p := NewLocalAIProxy()
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{}})).NotTo(Succeed())
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "not a url"}})).NotTo(Succeed())
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{UpstreamUrl: "ftp://host"}})).NotTo(Succeed())
})
It("refuses an api_key_env that is unset", func() {
err := NewLocalAIProxy().Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{
UpstreamUrl: up.URL, ApiKeyEnv: "LOCALAI_PROXY_TEST_UNSET_KEY",
}})
Expect(err).To(MatchError(ContainSubstring("LOCALAI_PROXY_TEST_UNSET_KEY")))
})
It("parses realtime_pipeline, strips the trailing slash and keeps the timeout", func() {
p := loadProxy(up, func(o *pb.ModelOptions) {
o.Options = []string{"other:1", "realtime_pipeline:my-pipe"}
o.Proxy.RequestTimeoutSeconds = 7
})
cfg := p.cfg.Load()
Expect(cfg.realtimePipeline).To(Equal("my-pipe"))
Expect(cfg.base).To(Equal(up.URL))
Expect(cfg.timeout).To(Equal(7 * time.Second))
})
DescribeTable("strips an OpenAI-style /v1 suffix, as the failover prober does",
func(suffix string) {
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamUrl = up.URL + suffix })
Expect(p.cfg.Load().base).To(Equal(up.URL))
},
Entry("/v1", "/v1"),
Entry("/v1/", "/v1/"),
Entry("a full endpoint path", "/v1/chat/completions"),
)
It("keeps a path prefix in front of /v1", func() {
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamUrl = up.URL + "/localai/v1" })
Expect(p.cfg.Load().base).To(Equal(up.URL + "/localai"))
})
It("falls back to the model name when upstream_model is unset", func() {
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.UpstreamModel = "" })
Expect(p.model("")).To(Equal("local-name"))
})
})
Describe("PredictRich", func() {
It("sends messages to /v1/chat/completions with the upstream model and key", func() {
GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-test")
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" })
up.replyJSON("/v1/chat/completions", map[string]any{
"choices": []any{map[string]any{"message": map[string]any{"role": "assistant", "content": "hello back"}}},
"usage": map[string]any{"prompt_tokens": 3, "completion_tokens": 2},
})
reply, err := p.PredictRich(&pb.PredictOptions{
Messages: []*pb.Message{{Role: "user", Content: "hello"}},
Tokens: 32,
Temperature: 0.5,
TopK: 40,
StopPrompts: []string{"</s>"},
Seed: 9,
})
Expect(err).NotTo(HaveOccurred())
Expect(string(reply.GetMessage())).To(Equal("hello back"))
Expect(reply.GetPromptTokens()).To(Equal(int32(3)))
Expect(reply.GetTokens()).To(Equal(int32(2)))
req := up.last()
Expect(req.Method).To(Equal(http.MethodPost))
Expect(req.Path).To(Equal("/v1/chat/completions"))
Expect(req.Auth).To(Equal("Bearer sk-test"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
Expect(req.JSON).To(HaveKeyWithValue("max_tokens", BeNumerically("==", 32)))
Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0.5)))
Expect(req.JSON).To(HaveKeyWithValue("top_k", BeNumerically("==", 40)))
Expect(req.JSON).To(HaveKeyWithValue("seed", BeNumerically("==", 9)))
Expect(req.JSON).To(HaveKeyWithValue("stop", ConsistOf("</s>")))
Expect(req.JSON).NotTo(HaveKey("stream"))
Expect(req.JSON["messages"]).To(ConsistOf(HaveKeyWithValue("content", "hello")))
})
It("returns upstream tool calls as chat deltas", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/chat/completions", map[string]any{
"choices": []any{map[string]any{"message": map[string]any{
"role": "assistant",
"tool_calls": []any{map[string]any{
"id": "call_1", "type": "function",
"function": map[string]any{"name": "get_weather", "arguments": `{"city":"Rome"}`},
}},
}}},
})
reply, err := p.PredictRich(&pb.PredictOptions{
Messages: []*pb.Message{{Role: "user", Content: "weather?"}},
Tools: `[{"type":"function","function":{"name":"get_weather"}}]`,
})
Expect(err).NotTo(HaveOccurred())
Expect(reply.GetChatDeltas()).To(HaveLen(1))
tc := reply.GetChatDeltas()[0].GetToolCalls()
Expect(tc).To(HaveLen(1))
Expect(tc[0].GetName()).To(Equal("get_weather"))
Expect(tc[0].GetArguments()).To(Equal(`{"city":"Rome"}`))
Expect(up.last().JSON).To(HaveKey("tools"))
})
It("sends a bare prompt to /v1/completions", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/completions", map[string]any{
"choices": []any{map[string]any{"text": "completed"}},
})
reply, err := p.PredictRich(&pb.PredictOptions{Prompt: "once upon"})
Expect(err).NotTo(HaveOccurred())
Expect(string(reply.GetMessage())).To(Equal("completed"))
req := up.last()
Expect(req.Path).To(Equal("/v1/completions"))
Expect(req.JSON).To(HaveKeyWithValue("prompt", "once upon"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
})
It("maps a 5xx upstream to Unavailable with the body in the message", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusServiceUnavailable, Body: "backend is down"})
_, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}})
Expect(codeOf(err)).To(Equal(codes.Unavailable))
Expect(err.Error()).To(ContainSubstring("backend is down"))
})
It("maps a 4xx upstream to InvalidArgument", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{Status: http.StatusBadRequest, Body: "bad request"})
_, err := p.PredictRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "x"}}})
Expect(codeOf(err)).To(Equal(codes.InvalidArgument))
Expect(err.Error()).To(ContainSubstring("bad request"))
})
It("truncates a long upstream error body", func() {
p := loadProxy(up, nil)
long := make([]byte, 2000)
for i := range long {
long[i] = 'a'
}
up.script("/v1/completions", scriptedResponse{Status: http.StatusInternalServerError, Body: string(long)})
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
Expect(codeOf(err)).To(Equal(codes.Unavailable))
Expect(len(status.Convert(err).Message())).To(BeNumerically("<", 700))
})
It("maps a 429 upstream to ResourceExhausted so failover moves to the next target", func() {
p := loadProxy(up, nil)
up.script("/v1/completions", scriptedResponse{Status: http.StatusTooManyRequests, Body: "slow down"})
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
Expect(codeOf(err)).To(Equal(codes.ResourceExhausted))
Expect(err.Error()).To(ContainSubstring("slow down"))
})
It("always sends temperature, even 0, so greedy decoding survives", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "t"}}})
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
Expect(err).NotTo(HaveOccurred())
req := up.last()
Expect(req.JSON).To(HaveKeyWithValue("temperature", BeNumerically("==", 0)))
Expect(req.JSON).NotTo(HaveKey("top_p"))
Expect(req.JSON).NotTo(HaveKey("top_k"))
})
It("maps an unreachable upstream to Unavailable", func() {
p := loadProxy(up, nil)
up.Close()
_, err := p.PredictRich(&pb.PredictOptions{Prompt: "x"})
Expect(codeOf(err)).To(Equal(codes.Unavailable))
})
It("reports an unloaded proxy as FailedPrecondition", func() {
_, err := NewLocalAIProxy().PredictRich(&pb.PredictOptions{Prompt: "x"})
Expect(codeOf(err)).To(Equal(codes.FailedPrecondition))
})
})
Describe("PredictStreamRich", func() {
It("streams SSE deltas in order and leaves the channel open", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"role": "assistant"}}}}),
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "Hel"}}}}),
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "lo"}}}}),
"[DONE]",
}})
results := make(chan *pb.Reply, 10)
err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)
Expect(err).NotTo(HaveOccurred())
var got []string
for len(results) > 0 {
got = append(got, string((<-results).GetMessage()))
}
Expect(got).To(Equal([]string{"Hel", "lo"}))
// The gRPC server closes the channel; closing it here must not panic.
close(results)
Expect(up.last().JSON).To(HaveKeyWithValue("stream", true))
})
It("asks the upstream for the usage trailer and reports its token counts", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "hi"}}}}),
sseJSON(map[string]any{"choices": []any{}, "usage": map[string]any{"prompt_tokens": 7, "completion_tokens": 1}}),
"[DONE]",
}})
results := make(chan *pb.Reply, 10)
Expect(p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)).To(Succeed())
// LocalAI (like OpenAI) only sends the usage trailer on request.
Expect(up.last().JSON).To(HaveKeyWithValue("stream_options", HaveKeyWithValue("include_usage", true)))
Expect(results).To(HaveLen(2))
<-results
usage := <-results
Expect(usage.GetPromptTokens()).To(Equal(int32(7)))
Expect(usage.GetTokens()).To(Equal(int32(1)))
})
It("streams /v1/completions text for a bare prompt", func() {
p := loadProxy(up, nil)
up.script("/v1/completions", scriptedResponse{SSE: []string{
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "a"}}}),
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "b"}}}),
"[DONE]",
}})
results := make(chan *pb.Reply, 10)
Expect(p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, results)).To(Succeed())
Expect(results).To(HaveLen(2))
Expect(string((<-results).GetMessage())).To(Equal("a"))
Expect(string((<-results).GetMessage())).To(Equal("b"))
})
It("stops the upstream request when the caller cancels", func() {
upstreamGone := make(chan struct{})
slow := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = w.Write([]byte("data: " + sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "a"}}}}) + "\n\n"))
w.(http.Flusher).Flush()
// A generation that outlasts the spec unless the proxy hangs up.
select {
case <-r.Context().Done():
close(upstreamGone)
case <-time.After(8 * time.Second):
}
})
DeferCleanup(slow.Close)
p := loadProxy(slow, nil)
addr := "test://localai-proxy-cancel"
grpc.Provide(addr, p)
client := grpc.NewClient(addr, true, nil, false)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
errCh := make(chan error, 1)
first := make(chan struct{}, 1)
go func() {
errCh <- client.PredictStream(ctx, &pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, func(*pb.Reply) {
select {
case first <- struct{}{}:
default:
}
})
}()
Eventually(first, 5*time.Second).Should(Receive())
cancel()
Eventually(upstreamGone, 5*time.Second).Should(BeClosed(), "the upstream generation must stop with the caller")
Eventually(errCh, 5*time.Second).Should(Receive(HaveOccurred()))
})
It("returns a mid-stream upstream error frame as Unavailable", func() {
p := loadProxy(up, nil)
up.script("/v1/chat/completions", scriptedResponse{SSE: []string{
sseJSON(map[string]any{"choices": []any{map[string]any{"delta": map[string]any{"content": "partial"}}}}),
sseJSON(map[string]any{"error": map[string]any{"message": "backend crashed", "type": "server_error", "code": "server_error"}}),
"[DONE]",
}})
results := make(chan *pb.Reply, 10)
err := p.PredictStreamRich(&pb.PredictOptions{Messages: []*pb.Message{{Role: "user", Content: "hi"}}}, results)
Expect(codeOf(err)).To(Equal(codes.Unavailable))
Expect(err.Error()).To(ContainSubstring("backend crashed"))
Expect(results).To(HaveLen(1))
})
It("maps a failing upstream to a gRPC code", func() {
p := loadProxy(up, nil)
up.script("/v1/completions", scriptedResponse{Status: http.StatusBadGateway, Body: "gateway"})
err := p.PredictStreamRich(&pb.PredictOptions{Prompt: "go"}, make(chan *pb.Reply, 1))
Expect(codeOf(err)).To(Equal(codes.Unavailable))
})
})
Describe("legacy Predict and PredictStream", func() {
It("wrap the rich variants", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/completions", map[string]any{"choices": []any{map[string]any{"text": "plain"}}})
out, err := p.Predict(&pb.PredictOptions{Prompt: "x"})
Expect(err).NotTo(HaveOccurred())
Expect(out).To(Equal("plain"))
up.script("/v1/completions", scriptedResponse{SSE: []string{
sseJSON(map[string]any{"choices": []any{map[string]any{"text": "s1"}}}),
"[DONE]",
}})
results := make(chan string, 10)
Expect(p.PredictStream(&pb.PredictOptions{Prompt: "x"}, results)).To(Succeed())
var got []string
for s := range results { // PredictStream closes the channel
got = append(got, s)
}
Expect(got).To(Equal([]string{"s1"}))
})
})
It("names the request fields it cannot forward", func() {
Expect(unforwardedFields(&pb.PredictOptions{Prompt: "x"})).To(BeEmpty())
Expect(unforwardedFields(&pb.PredictOptions{
Grammar: "root ::= x", Images: []string{"i"}, Audios: []string{"a"}, Videos: []string{"v"},
})).To(Equal([]string{"grammar", "images", "audios", "videos"}))
})
Describe("Embeddings", func() {
It("posts the input to /v1/embeddings and returns the first vector", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/embeddings", map[string]any{
"data": []any{map[string]any{"embedding": []float32{0.1, 0.2, 0.3}}},
})
vec, err := p.Embeddings(&pb.PredictOptions{Embeddings: "embed me"})
Expect(err).NotTo(HaveOccurred())
Expect(vec).To(Equal([]float32{0.1, 0.2, 0.3}))
req := up.last()
Expect(req.Path).To(Equal("/v1/embeddings"))
Expect(req.JSON).To(HaveKeyWithValue("input", "embed me"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
})
It("sends token input as a token array, not as an empty string", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/embeddings", map[string]any{
"data": []any{map[string]any{"embedding": []float32{0.5}}},
})
_, err := p.Embeddings(&pb.PredictOptions{EmbeddingTokens: []int32{1, 2, 3}})
Expect(err).NotTo(HaveOccurred())
Expect(up.last().JSON).To(HaveKeyWithValue("input", []any{[]any{1.0, 2.0, 3.0}}))
})
It("fails when the upstream returns no vector", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/embeddings", map[string]any{"data": []any{}})
_, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"})
Expect(err).To(HaveOccurred())
})
})
Describe("Rerank", func() {
It("posts to /v1/rerank and maps the results", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/rerank", map[string]any{
"model": "remote-model",
"usage": map[string]any{"total_tokens": 12, "prompt_tokens": 10},
"results": []any{
map[string]any{"index": 1, "document": map[string]any{"text": "b"}, "relevance_score": 0.9},
map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.1},
},
})
res, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}, TopN: 2})
Expect(err).NotTo(HaveOccurred())
Expect(res.GetUsage().GetTotalTokens()).To(Equal(int32(12)))
Expect(res.GetUsage().GetPromptTokens()).To(Equal(int32(10)))
Expect(res.GetResults()).To(HaveLen(2))
Expect(res.GetResults()[0].GetIndex()).To(Equal(int32(1)))
Expect(res.GetResults()[0].GetText()).To(Equal("b"))
Expect(res.GetResults()[0].GetRelevanceScore()).To(BeNumerically("~", 0.9, 1e-6))
req := up.last()
Expect(req.Path).To(Equal("/v1/rerank"))
Expect(req.JSON).To(HaveKeyWithValue("query", "q"))
Expect(req.JSON).To(HaveKeyWithValue("top_n", BeNumerically("==", 2)))
Expect(req.JSON).To(HaveKeyWithValue("documents", ConsistOf("a", "b")))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
})
})
It("Rerank omits top_n when it is 0, which means score every document", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/rerank", map[string]any{"results": []any{}})
_, err := p.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a", "b"}})
Expect(err).NotTo(HaveOccurred())
Expect(up.last().JSON).NotTo(HaveKey("top_n"))
})
Describe("TokenizeString and Detokenize", func() {
It("posts the prompt to /v1/tokenize", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/tokenize", map[string]any{"tokens": []int32{5, 6, 7}})
res, err := p.TokenizeString(&pb.PredictOptions{Prompt: "abc"})
Expect(err).NotTo(HaveOccurred())
Expect(res.GetTokens()).To(Equal([]int32{5, 6, 7}))
Expect(res.GetLength()).To(Equal(int32(3)))
req := up.last()
Expect(req.Path).To(Equal("/v1/tokenize"))
Expect(req.JSON).To(HaveKeyWithValue("content", "abc"))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
})
It("posts tokens to /v1/detokenize", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/detokenize", map[string]any{"content": "abc"})
res, err := p.Detokenize(&pb.DetokenizeRequest{Tokens: []int32{5, 6}})
Expect(err).NotTo(HaveOccurred())
Expect(res.GetContent()).To(Equal("abc"))
Expect(up.last().JSON).To(HaveKeyWithValue("tokens", ConsistOf(BeNumerically("==", 5), BeNumerically("==", 6))))
})
})
Describe("Score", func() {
It("posts to /api/score and maps the candidates", func() {
p := loadProxy(up, nil)
up.replyJSON("/api/score", map[string]any{
"model": "remote-model",
"candidates": []any{map[string]any{
"log_prob": -1.5, "length_normalized_log_prob": -0.75, "num_tokens": 2,
"tokens": []any{map[string]any{"token": "yes", "log_prob": -1.5}},
}},
})
res, err := p.Score(context.Background(), &pb.ScoreRequest{
Prompt: "p", Candidates: []string{"yes"}, IncludeTokenLogprobs: true, LengthNormalize: true,
})
Expect(err).NotTo(HaveOccurred())
Expect(res.GetCandidates()).To(HaveLen(1))
c := res.GetCandidates()[0]
Expect(c.GetLogProb()).To(Equal(-1.5))
Expect(c.GetLengthNormalizedLogProb()).To(Equal(-0.75))
Expect(c.GetNumTokens()).To(Equal(int32(2)))
Expect(c.GetTokens()).To(HaveLen(1))
Expect(c.GetTokens()[0].GetToken()).To(Equal("yes"))
req := up.last()
Expect(req.Path).To(Equal("/api/score"))
Expect(req.JSON).To(HaveKeyWithValue("prompt", "p"))
Expect(req.JSON).To(HaveKeyWithValue("include_token_logprobs", true))
Expect(req.JSON).To(HaveKeyWithValue("length_normalize", true))
Expect(req.JSON).To(HaveKeyWithValue("model", "remote-model"))
})
It("refuses decision-pipeline requests it cannot forward", func() {
p := loadProxy(up, nil)
_, err := p.Score(context.Background(), &pb.ScoreRequest{Prompt: "{}", QuestionType: "systemone"})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
Expect(up.recorded()).To(BeEmpty())
})
})
Describe("helpers", func() {
It("postMultipart sends fields and the file with auth", func() {
GinkgoT().Setenv("LOCALAI_PROXY_TEST_KEY", "sk-mp")
p := loadProxy(up, func(o *pb.ModelOptions) { o.Proxy.ApiKeyEnv = "LOCALAI_PROXY_TEST_KEY" })
up.replyJSON("/v1/audio/transcriptions", map[string]any{"text": "ok"})
path := filepath.Join(GinkgoT().TempDir(), "a.wav")
Expect(os.WriteFile(path, []byte("RIFFDATA"), 0o600)).To(Succeed())
var out struct {
Text string `json:"text"`
}
err := p.postMultipart(context.Background(), "/v1/audio/transcriptions",
map[string]string{"model": "remote-model", "language": "it"}, "file", path, &out)
Expect(err).NotTo(HaveOccurred())
Expect(out.Text).To(Equal("ok"))
req := up.last()
Expect(req.Auth).To(Equal("Bearer sk-mp"))
Expect(req.Fields).To(Equal(map[string]string{"model": "remote-model", "language": "it"}))
Expect(req.Files).To(HaveKeyWithValue("file", "RIFFDATA"))
})
It("postMultipart reports a missing local file without calling upstream", func() {
p := loadProxy(up, nil)
err := p.postMultipart(context.Background(), "/v1/audio/transcriptions", nil, "file", "/nonexistent/a.wav", nil)
Expect(err).To(HaveOccurred())
Expect(up.recorded()).To(BeEmpty())
})
It("applies request_timeout_seconds to non-streaming calls", func() {
slow := make(chan struct{})
hang := newFakeUpstreamWithHandler(func(w http.ResponseWriter, r *http.Request) { <-slow })
// Cleanups run last-in first-out: release the handler before
// Close, which waits for in-flight requests.
DeferCleanup(hang.Close)
DeferCleanup(func() { close(slow) })
p := NewLocalAIProxy()
Expect(p.Load(&pb.ModelOptions{Proxy: &pb.ProxyOptions{
UpstreamUrl: hang.URL, UpstreamModel: "m", RequestTimeoutSeconds: 1,
}})).To(Succeed())
_, err := p.Embeddings(&pb.PredictOptions{Embeddings: "x"})
Expect(codeOf(err)).To(Equal(codes.DeadlineExceeded))
})
It("postStream returns the open response for a 2xx", func() {
p := loadProxy(up, nil)
up.script("/tts", scriptedResponse{Status: http.StatusOK, ContentType: "audio/wav", Body: "WAVBYTES"})
resp, err := p.postStream(context.Background(), "/tts", map[string]any{"input": "hi"})
Expect(err).NotTo(HaveOccurred())
defer func() { _ = resp.Body.Close() }()
Expect(resp.Header.Get("Content-Type")).To(Equal("audio/wav"))
})
})
Describe("through the gRPC server", func() {
It("dispatches Rerank and keeps the Unimplemented code end to end", func() {
p := loadProxy(up, nil)
up.replyJSON("/v1/rerank", map[string]any{"results": []any{
map[string]any{"index": 0, "document": map[string]any{"text": "a"}, "relevance_score": 0.5},
}})
addr := "test://localai-proxy-grpc"
grpc.Provide(addr, p)
client := grpc.NewClient(addr, true, nil, false)
res, err := client.Rerank(context.Background(), &pb.RerankRequest{Query: "q", Documents: []string{"a"}})
Expect(err).NotTo(HaveOccurred())
Expect(res.GetResults()).To(HaveLen(1))
_, err = client.AudioEncode(context.Background(), &pb.AudioEncodeRequest{})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart"))
})
})
Describe("methods with no upstream counterpart", func() {
It("return Unimplemented with the exact message", func() {
p := loadProxy(up, nil)
_, err := p.AudioEncode(&pb.AudioEncodeRequest{})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
Expect(status.Convert(err).Message()).To(Equal("localai-proxy: AudioEncode has no upstream counterpart"))
_, err = p.TokenClassify(context.Background(), &pb.TokenClassifyRequest{})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
_, err = p.ModelMetadata(&pb.ModelOptions{})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
_, err = p.StartFineTune(&pb.FineTuneRequest{})
Expect(codeOf(err)).To(Equal(codes.Unimplemented))
})
It("close the output channel of streaming stubs so the server does not hang", func() {
p := loadProxy(up, nil)
updates := make(chan *pb.FineTuneProgressUpdate)
Expect(codeOf(p.FineTuneProgress(&pb.FineTuneProgressRequest{}, updates))).To(Equal(codes.Unimplemented))
Eventually(updates).Should(BeClosed())
out := make(chan *pb.AudioToAudioResponse)
Expect(codeOf(p.AudioToAudioStream(make(chan *pb.AudioToAudioRequest), out))).To(Equal(codes.Unimplemented))
Eventually(out).Should(BeClosed())
})
})
})
var _ = Describe("clampInt32", func() {
It("passes in-range values through and saturates the rest", func() {
Expect(clampInt32(0)).To(Equal(int32(0)))
Expect(clampInt32(42)).To(Equal(int32(42)))
Expect(clampInt32(-7)).To(Equal(int32(-7)))
Expect(clampInt32(math.MaxInt32 + 1)).To(Equal(int32(math.MaxInt32)))
Expect(clampInt32(math.MinInt32 - 1)).To(Equal(int32(math.MinInt32)))
})
})
+40
View File
@@ -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"
+30
View File
@@ -13,9 +13,12 @@ import (
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/http/auth"
mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp"
"github.com/mudler/LocalAI/core/services/advisorylock"
"github.com/mudler/LocalAI/core/services/agentpool"
"github.com/mudler/LocalAI/core/services/cloudproxy/mitm"
"github.com/mudler/LocalAI/core/services/facerecognition"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/failover/distsync"
"github.com/mudler/LocalAI/core/services/galleryop"
"github.com/mudler/LocalAI/core/services/monitoring"
"github.com/mudler/LocalAI/core/services/nodes"
@@ -82,6 +85,7 @@ type Application struct {
routerRegistry *router.Registry
routerCorpus *corpus.Manager
admissionLimiter *admission.Limiter
failoverManager *failover.Manager
watchdogMutex sync.Mutex
watchdogStop chan bool
p2pMutex sync.Mutex
@@ -92,6 +96,13 @@ type Application struct {
// Distributed mode services (nil when not in distributed mode)
distributed *DistributedServices
// failoverSync shares failover state between frontends; nil in
// standalone mode or when it could not start.
failoverSync *distsync.Sync
// failoverLock is the probe-leader lock this frontend holds or tries
// for; nil when failoverSync is nil.
failoverLock *advisorylock.HeldLock
// Upgrade checker (background service for detecting backend upgrades)
upgradeChecker *UpgradeChecker
@@ -99,6 +110,13 @@ type Application struct {
// is set; otherwise initialised in start() after galleryService.
localAIAssistant *mcpTools.LocalAIAssistantHolder
// assistantClient is the concrete inproc client backing localAIAssistant.
// start() constructs it before failoverManager exists (see New() in
// startup.go), so New() sets assistantClient.Failover once the manager
// is built, using this field to reach back into the already-registered
// MCP tool set. nil when DisableLocalAIAssistant is set.
assistantClient *localaiInproc.Client
// startupComplete flips to true once New() has finished its whole startup
// sequence. It backs the /readyz probe.
//
@@ -478,6 +496,9 @@ func (a *Application) AdmissionLimiter() *admission.Limiter {
return a.admissionLimiter
}
// FailoverManager serves failover chains. Never nil after New.
func (a *Application) FailoverManager() *failover.Manager { return a.failoverManager }
// StartupConfig returns the original startup configuration (from env vars, before file loading)
func (a *Application) StartupConfig() *config.ApplicationConfig {
return a.startupConfig
@@ -503,6 +524,9 @@ func (a *Application) IsDistributed() bool {
func (a *Application) Shutdown() error {
var err error
a.shutdownOnce.Do(func() {
// Before distributed shutdown: the sync's subscriptions live on the
// NATS connection that closes there.
a.stopFailoverDistributed()
a.distributed.Shutdown()
if a.modelLoader != nil {
err = a.modelLoader.StopAllGRPC()
@@ -593,6 +617,12 @@ func (a *Application) start() error {
assistantClient.RouterEmbedder = a.Embedder
assistantClient.RouterEmbedderFingerprint = a.EmbedderFingerprint
assistantClient.RouterVectorStore = a.VectorStore
// Failover chains: failoverManager does not exist yet at this point
// in startup (New() in startup.go builds it after start() returns),
// so it can't be wired here like the fields above. New() sets
// assistantClient.Failover directly once the manager is built;
// stash the client so it can reach back into it.
a.assistantClient = assistantClient
if err := holder.Initialize(a.applicationConfig.Context, assistantClient, localaitools.Options{}); err != nil {
// Why log+continue instead of fail: the assistant is an optional
// feature; a failure here must not take down the whole server.
+6 -1
View File
@@ -77,7 +77,9 @@ func (ds *DistributedServices) Shutdown() {
// Returns nil if distributed mode is not enabled.
// configLoader is used by the SmartRouter to compute concurrency-group
// anti-affinity at placement time (#9659); it may be nil in tests.
func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader) (*DistributedServices, error) {
// pinned, when set, replaces configLoader as the source of models the router
// and reconciler must keep loaded (it adds warm failover targets).
func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoader *config.ModelConfigLoader, pinned nodes.PinnedModelResolver) (*DistributedServices, error) {
if !cfg.Distributed.Enabled {
return nil, nil
}
@@ -387,6 +389,9 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade
modelFiles = declaredModelFiles(configLoader, cfg.SystemState.Model.ModelsPath)
}
}
if pinned != nil {
pinnedResolver = pinned
}
modelCleanup := nodes.NewModelCleanupService(registry, remoteUnloader)
router := nodes.NewSmartRouter(registry, nodes.SmartRouterOptions{
Unloader: remoteUnloader,
+59
View File
@@ -0,0 +1,59 @@
package application
import (
"github.com/mudler/LocalAI/core/backend"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/pkg/grpc"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/xlog"
)
// preloadModelByName is a seam over backend.PreloadModelByName so tests can
// substitute a controllable loader instead of touching real models/disk.
var preloadModelByName = backend.PreloadModelByName
// applyFailoverWarmTargets pins warm failover targets in the watchdog and
// loads them, so a switch does not wait for a cold load.
//
// SyncPinnedModelsToWatchdog runs synchronously: it is a cheap in-memory
// update, and the pin must land before the watchdog can evict a target that
// is about to become (or stay) a chain's active path. Preloading is not
// cheap — it can download or load a multi-GB model — and this callback runs
// on the failover manager's single scheduler goroutine (Sync, called from
// Tick, called from Run), before that tick's probes fire. A slow or hung
// load here would freeze probing and fail-back for every chain, so it runs
// in its own goroutine instead of blocking the scheduler loop.
//
// In distributed mode only the probe leader preloads: every frontend would
// otherwise ask the workers for the same load. Followers still pin, so their
// own watchdog never evicts a warm target they happen to hold.
func (a *Application) applyFailoverWarmTargets(warm []string) {
a.SyncPinnedModelsToWatchdog()
if a.IsDistributed() && !a.failoverManager.IsLeader() {
return
}
go func() {
for _, name := range warm {
if _, err := preloadModelByName(a.ApplicationConfig().Context, a.ModelConfigLoader(), a.ModelLoader(), a.ApplicationConfig(), name); err != nil {
xlog.Warn("failover: could not preload warm target", "model", name, "error", err)
}
}
}()
}
// failoverLoadedBackend gives the failover prober the running backend of a
// local target without ever loading it. CheckIsLoaded may run the loader's
// own health check and drop a dead process; the target is then "not loaded"
// and the next real request loads and judges it.
func failoverLoadedBackend(ml *model.ModelLoader) failover.LoadedFunc {
return func(cfg config.ModelConfig) grpc.Backend {
m := ml.CheckIsLoaded(cfg.ModelID())
if m == nil {
return nil
}
// Load always enables parallel requests; match it in case this is
// the first client built for the model.
return m.GRPC(true, ml.GetWatchDog())
}
}
+94
View File
@@ -0,0 +1,94 @@
package application
import (
"context"
"github.com/mudler/LocalAI/core/services/advisorylock"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/failover/distsync"
"github.com/mudler/LocalAI/core/services/nodes"
"github.com/mudler/xlog"
)
// failoverLeaderGate elects the frontend that probes failover targets and
// decides chains, so N frontends do not probe every target N times or
// disagree on which target is active.
//
// Leadership is sticky: the leader keeps the lock across ticks and loses it
// only when its database session dies or it shuts down. A lock taken per tick
// would pass between frontends on almost every tick, and each change of
// leader redelivers the warm set and republishes all state. A lock error
// counts as "not leader": skipping one tick is safe, two leaders are not.
func failoverLeaderGate(l *advisorylock.HeldLock) failover.LeaderGate {
return func(ctx context.Context, fn func()) bool {
if !l.Held() || !l.Verify(ctx) {
ok, err := l.TryAcquire(ctx)
if err != nil {
xlog.Warn("failover: could not take the prober leader lock", "error", err)
return false
}
if !ok {
return false
}
}
fn()
return true
}
}
// failoverPinnedResolver adds warm failover targets to the config-pinned
// models, so the SmartRouter and ReplicaReconciler keep them loaded on the
// workers the same way the local watchdog keeps them loaded in standalone mode.
type failoverPinnedResolver struct {
base nodes.PinnedModelResolver
fm *failover.Manager
}
func (r *failoverPinnedResolver) GetPinnedModelNames() []string {
var pinned []string
if r.base != nil {
pinned = r.base.GetPinnedModelNames()
}
return failover.MergePinned(pinned, r.fm.WarmTargets())
}
// startFailoverDistributed makes the failover manager cluster-aware: one
// frontend (the advisory-lock holder) probes and decides, and pins, target
// health and chain state are shared over NATS. A sync failure is logged and
// the manager keeps probing on its own, as in standalone mode.
func (a *Application) startFailoverDistributed(ctx context.Context) {
db := a.distributedDB()
pins, err := distsync.NewPinStore(db)
if err != nil {
xlog.Error("failover: pins will not persist, could not prepare the pin store", "error", err)
pins = nil // distsync.New treats a nil store as "no durable pins"
}
s, err := distsync.New(ctx, a.distributed.Nats, pins, a.failoverManager)
if err != nil {
xlog.Error("failover: state will not be shared between frontends", "error", err)
return
}
a.failoverSync = s
// Gate only once state is shared: a follower learns target health and
// chain decisions solely from the leader's publishes, so gating without
// the sync would leave every follower's chains frozen.
a.failoverLock = advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber)
a.failoverManager.SetLeaderGate(failoverLeaderGate(a.failoverLock))
}
// stopFailoverDistributed detaches the manager from the sync before closing
// it, so nothing publishes into closed maps, and gives up leadership at once
// instead of making another frontend wait for this session to time out.
func (a *Application) stopFailoverDistributed() {
if a.failoverSync != nil {
a.failoverManager.SetStateSync(nil)
if err := a.failoverSync.Close(); err != nil {
xlog.Warn("failover: closing state sync", "error", err)
}
}
if a.failoverLock != nil {
// Close, not Release: Run may still be ticking and would take the
// lock straight back.
a.failoverLock.Close()
}
}
@@ -0,0 +1,134 @@
package application
import (
"context"
"time"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/services/advisorylock"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/pkg/model"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
// staticPins is a nodes.PinnedModelResolver with a fixed list.
type staticPins []string
func (s staticPins) GetPinnedModelNames() []string { return s }
// failoverSource is a failover.ConfigSource over a fixed set of configs.
type failoverSource map[string]config.ModelConfig
func (s failoverSource) GetModelConfig(name string) (config.ModelConfig, bool) {
c, ok := s[name]
return c, ok
}
func (s failoverSource) GetAllModelsConfigs() []config.ModelConfig {
out := make([]config.ModelConfig, 0, len(s))
for _, c := range s {
out = append(out, c)
}
return out
}
// warmChainSource has one chain whose two local targets are warm.
func warmChainSource() failoverSource {
return failoverSource{
"b": {Name: "b", Backend: "llama-cpp"},
"c": {Name: "c", Backend: "llama-cpp"},
"chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{
{Model: "b", Warm: true}, {Model: "c", Warm: true},
}}},
}
}
var _ = Describe("failoverPinnedResolver", func() {
It("merges config pins and warm failover targets without duplicates", func() {
fm := failover.New(warmChainSource())
fm.Sync()
Expect(fm.WarmTargets()).To(Equal([]string{"b", "c"}))
r := &failoverPinnedResolver{base: staticPins{"pinned-a", "b"}, fm: fm}
Expect(r.GetPinnedModelNames()).To(Equal([]string{"pinned-a", "b", "c"}))
})
})
var _ = Describe("failoverLeaderGate", func() {
It("keeps leadership with the first frontend until it releases", func() {
// Not PostgreSQL, so advisorylock falls back to its in-process lock,
// which has the same try-lock semantics as pg_try_advisory_lock.
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
Expect(err).ToNot(HaveOccurred())
firstLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber)
secondLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber)
DeferCleanup(firstLock.Release)
DeferCleanup(secondLock.Release)
first, second := failoverLeaderGate(firstLock), failoverLeaderGate(secondLock)
ctx := context.Background()
var firstRuns, secondRuns int
for range 5 {
Expect(first(ctx, func() { firstRuns++ })).To(BeTrue())
Expect(second(ctx, func() { secondRuns++ })).To(BeFalse(), "leadership must not flip between ticks")
}
Expect(firstRuns).To(Equal(5))
Expect(secondRuns).To(BeZero())
firstLock.Release()
Expect(second(ctx, func() { secondRuns++ })).To(BeTrue(), "a released leadership passes to the next frontend")
Expect(secondRuns).To(Equal(1))
Expect(first(ctx, func() { firstRuns++ })).To(BeFalse())
})
})
var _ = Describe("applyFailoverWarmTargets in distributed mode", func() {
It("pins but does not preload on a frontend that is not the probe leader", func() {
preloaded := make(chan string, 1)
orig := preloadModelByName
preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) {
preloaded <- name
return nil, nil
}
DeferCleanup(func() { preloadModelByName = orig })
// A gate that never grants leadership: this frontend is a follower.
fm := failover.New(warmChainSource(),
failover.WithLeaderGate(func(context.Context, func()) bool { return false }))
app := &Application{
applicationConfig: &config.ApplicationConfig{Context: context.Background()},
distributed: &DistributedServices{},
failoverManager: fm,
}
app.applyFailoverWarmTargets([]string{"b"})
Consistently(preloaded, 200*time.Millisecond).ShouldNot(Receive(), "only the probe leader preloads warm targets")
})
It("preloads warm targets on the probe leader", func() {
preloaded := make(chan string, 4)
orig := preloadModelByName
preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) {
preloaded <- name
return nil, nil
}
DeferCleanup(func() { preloadModelByName = orig })
app := &Application{
applicationConfig: &config.ApplicationConfig{Context: context.Background()},
distributed: &DistributedServices{},
}
// A gate that always grants: this frontend is the leader. The warm
// set is delivered on the first tick that wins it.
app.failoverManager = failover.New(warmChainSource(),
failover.WithLeaderGate(func(_ context.Context, fn func()) bool { fn(); return true }),
failover.WithOnWarmChanged(app.applyFailoverWarmTargets))
app.failoverManager.Tick(context.Background())
Eventually(preloaded, time.Second).Should(Receive(Equal("b")))
Eventually(preloaded, time.Second).Should(Receive(Equal("c")))
})
})
+63
View File
@@ -0,0 +1,63 @@
package application
import (
"context"
"time"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/pkg/grpc"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/system"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
var _ = Describe("applyFailoverWarmTargets", func() {
It("returns promptly even when the preload call blocks", func() {
// Guards against a regression to a synchronous preload loop: onWarm
// runs on the failover manager's single scheduler goroutine, so a
// blocking loader here must not block the caller.
started := make(chan struct{})
release := make(chan struct{})
orig := preloadModelByName
preloadModelByName = func(ctx context.Context, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, name string) ([]string, error) {
close(started)
<-release // never released within the test's timeout
return nil, nil
}
DeferCleanup(func() { preloadModelByName = orig })
DeferCleanup(func() { close(release) })
app := &Application{applicationConfig: &config.ApplicationConfig{Context: context.Background()}}
callReturned := make(chan struct{})
go func() {
defer GinkgoRecover()
app.applyFailoverWarmTargets([]string{"warm-a"})
close(callReturned)
}()
Eventually(callReturned, time.Second).Should(BeClosed(), "applyFailoverWarmTargets must not wait on the preload goroutine")
Eventually(started, time.Second).Should(BeClosed(), "the preload goroutine should still run in the background")
})
})
type healthyBackend struct{ grpc.Backend }
func (healthyBackend) HealthCheck(context.Context) (bool, error) { return true, nil }
var _ = Describe("failoverLoadedBackend", func() {
It("returns the running backend and never loads a model that is not loaded", func() {
ml := model.NewModelLoader(&system.SystemState{Model: system.Model{ModelsPath: GinkgoT().TempDir()}})
store := model.NewInMemoryModelStore()
ml.SetModelStore(store)
loaded := failoverLoadedBackend(ml)
Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(BeNil())
Expect(ml.ListLoadedModels()).To(BeEmpty(), "the lookup must not start a load")
client := healthyBackend{}
store.Set("gemma", model.NewModelWithClient("gemma", "127.0.0.1:0", client))
Expect(loaded(config.ModelConfig{Name: "gemma", Backend: "llama-cpp"})).To(Equal(client))
})
})
+26 -1
View File
@@ -12,6 +12,7 @@ import (
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/gallery"
"github.com/mudler/LocalAI/core/http/auth"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/galleryop"
"github.com/mudler/LocalAI/core/services/jobs"
"github.com/mudler/LocalAI/core/services/messaging"
@@ -250,6 +251,20 @@ func New(opts ...config.AppOption) (*Application, error) {
// the embedding-cache stats endpoint sees a single source of truth.
application.routerRegistry = router.NewRegistry()
// Failover chains: probe targets and track which one is active per
// chain. WithOnWarmChanged pins and preloads warm local targets so a
// switch to them does not wait for a cold load.
application.failoverManager = failover.New(application.ModelConfigLoader(),
failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()), options.ProxyAPIKeyEnvLookup)),
failover.WithOnWarmChanged(application.applyFailoverWarmTargets),
)
// The assistant client was built in start() (above), before this
// manager existed; wire it now so list_failover_chains /
// pin_failover_target / unpin_failover_target see real chains.
if application.assistantClient != nil {
application.assistantClient.Failover = application.failoverManager
}
// Subsystem 5: admission control. Limiter is always wired so a
// model that gains a limits: block via gallery install or YAML
// edit takes effect on the next restart without conditional plumbing.
@@ -271,12 +286,16 @@ func New(opts ...config.AppOption) (*Application, error) {
// the model configs are loaded, so it is declared out here.
var revisionStore modeladmin.RevisionStore
distSvc, err := initDistributed(options, application.authDB, application.ModelConfigLoader())
distSvc, err := initDistributed(options, application.authDB, application.ModelConfigLoader(),
&failoverPinnedResolver{base: application.ModelConfigLoader(), fm: application.failoverManager})
if err != nil {
return nil, fmt.Errorf("distributed mode initialization failed: %w", err)
}
if distSvc != nil {
application.distributed = distSvc
// Before failoverManager.Run starts below: the gate and sync must be
// in place for its first tick.
application.startFailoverDistributed(options.Context)
// Wire remote model unloader so ShutdownModel works for remote nodes
// Uses NATS to tell serve-backend nodes to Free + kill their backend process
application.modelLoader.SetRemoteUnloader(distSvc.Unloader)
@@ -548,6 +567,12 @@ func New(opts ...config.AppOption) (*Application, error) {
}
}
// Start the failover scheduler: it syncs chains from config, runs
// liveness/recovery probes and dwell-based fail-back. Run is the only
// caller of Sync in production so onWarm callbacks stay ordered.
failover.RegisterMetrics(application.failoverManager)
go application.failoverManager.Run(options.Context)
// Watch the configuration directory
startWatcher(options)
+4
View File
@@ -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))
}
+9 -1
View File
@@ -542,7 +542,7 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions {
Tokenizer: c.Tokenizer,
}
if c.Backend == "cloud-proxy" {
if c.IsRemoteProxy() {
opts.Proxy = &pb.ProxyOptions{
UpstreamUrl: c.Proxy.UpstreamURL,
Mode: c.Proxy.Mode,
@@ -553,6 +553,14 @@ func grpcModelOpts(c config.ModelConfig, modelPath string) *pb.ModelOptions {
RequestTimeoutSeconds: int32(c.Proxy.RequestTimeoutSeconds),
CachePrompt: c.Proxy.CachePrompt,
}
// localai-proxy calls a LocalAI that knows the model by name, so an
// unset upstream_model means this config's name, the same derivation
// failover.UpstreamModel uses. Not for cloud-proxy: its translate mode
// falls back to parameters.model and passthrough keeps the client's
// model when upstream_model is empty.
if c.Backend == model.LocalAIProxyBackend && opts.Proxy.UpstreamModel == "" {
opts.Proxy.UpstreamModel = c.Name
}
}
if c.MMProj != "" {
+52
View File
@@ -44,6 +44,58 @@ var _ = Describe("grpcModelOpts EngineArgs", func() {
})
})
var _ = Describe("grpcModelOpts Proxy options", func() {
It("builds Proxy options for localai-proxy, the same as cloud-proxy", func() {
threads := 1
cfg := config.ModelConfig{
Threads: &threads,
Backend: "localai-proxy",
Proxy: config.ProxyConfig{
UpstreamURL: "http://127.0.0.1:8081",
Mode: config.ProxyModePassthrough,
},
}
opts := grpcModelOpts(cfg, "/tmp/models")
Expect(opts.Proxy).NotTo(BeNil())
Expect(opts.Proxy.UpstreamUrl).To(Equal("http://127.0.0.1:8081"))
Expect(opts.Proxy.Mode).To(Equal(config.ProxyModePassthrough))
})
It("sends the config name as the localai-proxy upstream model when none is set", func() {
threads := 1
cfg := config.ModelConfig{
Name: "argus-whisper",
Threads: &threads,
Backend: "localai-proxy",
Proxy: config.ProxyConfig{UpstreamURL: "http://127.0.0.1:8081"},
}
cfg.Model = "some-file"
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("argus-whisper"))
cfg.Proxy.UpstreamModel = "whisper-large"
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(Equal("whisper-large"))
})
It("leaves the cloud-proxy upstream model unset so translate mode keeps its fallback", func() {
threads := 1
cfg := config.ModelConfig{
Name: "claude-strict",
Threads: &threads,
Backend: "cloud-proxy",
Proxy: config.ProxyConfig{UpstreamURL: "https://api.example.com", Mode: config.ProxyModeTranslate},
}
Expect(grpcModelOpts(cfg, "/tmp/models").Proxy.UpstreamModel).To(BeEmpty())
})
It("leaves Proxy nil for a backend that is not a proxy", func() {
threads := 1
opts := grpcModelOpts(config.ModelConfig{Threads: &threads, Backend: "llama-cpp"}, "/tmp/models")
Expect(opts.Proxy).To(BeNil())
})
})
var _ = Describe("grpcModelOpts Diffusers options", func() {
It("forwards original_config_file without rewriting it", func() {
threads := 1
+6 -6
View File
@@ -66,12 +66,12 @@ func pipelineStages(cl *config.ModelConfigLoader, p *config.Pipeline, modelPath
}
var stages []PreloadStage
for _, s := range []struct{ role, name string }{
{"vad", p.VAD},
{"transcription", p.Transcription},
{"llm", p.LLM},
{"tts", p.TTS},
{"sound_detection", p.SoundDetection},
{"voice_recognition", voiceRec},
{config.PipelineStageVAD, p.VAD},
{config.PipelineStageTranscription, p.Transcription},
{config.PipelineStageLLM, p.LLM},
{config.PipelineStageTTS, p.TTS},
{config.PipelineStageSoundDetection, p.SoundDetection},
{config.PipelineStageVoiceRecognition, voiceRec},
} {
if s.name == "" {
continue
+1
View File
@@ -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(
+7
View File
@@ -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(),
+30
View File
@@ -0,0 +1,30 @@
package config
import "github.com/mudler/LocalAI/pkg/model"
func init() {
RegisterBackendHook(model.LocalAIProxyBackend, localAIProxyDefaults)
}
// localAIProxyDefaults makes chat requests reach the upstream as structured
// messages. Without the tokenizer template core renders the prompt itself
// with no template, and the proxy can only send that text to
// /v1/completions, bypassing the upstream model's chat template, tool
// handling and reasoning parsing.
//
// Only configs that declare the chat usecase get it: usecase guessing reads
// the tokenizer template as "this model chats", so setting it on a
// transcription or embedding proxy would offer that model to chat pickers
// and default-model selection. A config that brings its own templates keeps
// them: the operator chose local templating.
func localAIProxyDefaults(cfg *ModelConfig, _ string) {
t := cfg.TemplateConfig
if t.UseTokenizerTemplate || t.Chat != "" || t.ChatMessage != "" || t.Completion != "" || t.Edit != "" {
return
}
declared := GetUsecasesFromYAML(cfg.KnownUsecaseStrings)
if declared == nil || *declared&FLAG_CHAT != FLAG_CHAT {
return
}
cfg.TemplateConfig.UseTokenizerTemplate = true
}
+32
View File
@@ -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{
+33
View File
@@ -419,6 +419,39 @@ func DefaultRegistry() map[string]FieldMetaOverride {
Order: 0,
},
// --- Failover ---
"failover.targets": {
Section: "failover",
Label: "Failover targets",
Description: "Ordered list of models that serve this chain. The first healthy target serves each request; later targets take over when it fails. Mark a local target warm to keep it loaded.",
Component: "failover-targets",
Order: 0,
},
"failover.probe.interval": {
Section: "failover", Label: "Probe interval", Component: "input", Order: 1, Advanced: true,
Description: "How often an idle target is checked, as a duration (default 15s).", Placeholder: "15s",
},
"failover.probe.timeout": {
Section: "failover", Label: "Probe timeout", Component: "input", Order: 2, Advanced: true,
Description: "How long one probe may take (default 5s).", Placeholder: "5s",
},
"failover.trip.errors": {
Section: "failover", Label: "Errors to trip", Component: "number", Order: 3, Advanced: true,
Description: "Failures within the trip window that mark a target down (default 1).",
},
"failover.trip.window": {
Section: "failover", Label: "Trip window", Component: "input", Order: 4, Advanced: true,
Description: "Window in which failures are counted (default 30s).", Placeholder: "30s",
},
"failover.recovery.probes": {
Section: "failover", Label: "Recovery probes", Component: "number", Order: 5, Advanced: true,
Description: "Consecutive real test requests a target must pass before it is used again (default 3).",
},
"failover.recovery.min_dwell": {
Section: "failover", Label: "Minimum time on fallback", Component: "input", Order: 6, Advanced: true,
Description: "Minimum time on a lower target before traffic moves back to a recovered higher one (default 60s).", Placeholder: "60s",
},
// --- Pipeline ---
"pipeline.llm": {
Section: "pipeline",
+12
View File
@@ -28,6 +28,18 @@ var _ = Describe("alias field metadata", func() {
}
Expect(found).To(BeTrue(), "DefaultSections should include an alias section")
})
It("registers the failover section", func() {
reg := meta.DefaultRegistry()
Expect(reg).To(HaveKey("failover.targets"))
Expect(reg["failover.targets"].Section).To(Equal("failover"))
Expect(reg["failover.targets"].Component).To(Equal("failover-targets"))
var ids []string
for _, s := range meta.DefaultSections() {
ids = append(ids, s.ID)
}
Expect(ids).To(ContainElement("failover"))
})
})
var _ = Describe("MCP field metadata", func() {
+1
View File
@@ -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},
+52 -3
View File
@@ -17,6 +17,7 @@ import (
"github.com/mudler/LocalAI/core/services/routing/piipattern"
"github.com/mudler/LocalAI/pkg/downloader"
"github.com/mudler/LocalAI/pkg/functions"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/modelartifacts"
"github.com/mudler/LocalAI/pkg/reasoning"
"github.com/mudler/cogito"
@@ -75,6 +76,10 @@ type ModelConfig struct {
// at create/swap time). See docs/content for Model Aliases.
Alias string `yaml:"alias,omitempty" json:"alias,omitempty"`
// Failover makes this config a failover chain over other models. Like an
// alias it has no backend of its own.
Failover *FailoverConfig `yaml:"failover,omitempty" json:"failover,omitempty"`
F16 *bool `yaml:"f16,omitempty" json:"f16,omitempty"`
Threads *int `yaml:"threads,omitempty" json:"threads,omitempty"`
Debug *bool `yaml:"debug,omitempty" json:"debug,omitempty"`
@@ -298,12 +303,37 @@ const (
ProxyProviderAnthropic = "anthropic"
)
// ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or
// "" when neither is set. Mirrored (not imported, to keep backends independent
// of core's package layout) by resolveAPIKey in backend/go/cloud-proxy/proxy.go
// — keep the two in sync, empty-value handling included.
func (p ProxyConfig) ResolveAPIKey(envLookup func(string) string) (string, error) {
switch {
case p.APIKeyEnv != "":
var v string
if envLookup != nil {
v = envLookup(p.APIKeyEnv)
}
if v == "" {
return "", fmt.Errorf("proxy api_key_env %q is unset", p.APIKeyEnv)
}
return v, nil
case p.APIKeyFile != "":
b, err := os.ReadFile(p.APIKeyFile)
if err != nil {
return "", fmt.Errorf("proxy api_key_file %q: %w", p.APIKeyFile, err)
}
return strings.TrimSpace(string(b)), nil
}
return "", nil
}
// IsCloudProxyBackendPassthrough reports whether this model uses the
// cloud-proxy gRPC backend in passthrough mode. Empty Mode counts as
// passthrough (SetDefaults normalises it, but Validate accepts empty
// too — handlers should not rely on a particular call order).
func (c *ModelConfig) IsCloudProxyBackendPassthrough() bool {
if c.Backend != "cloud-proxy" {
if c.Backend != model.CloudProxyBackend {
return false
}
return c.Proxy.Mode == "" || c.Proxy.Mode == ProxyModePassthrough
@@ -641,7 +671,7 @@ func (c *ModelConfig) PIIIsEnabled() bool {
if c.PII.Enabled != nil {
return *c.PII.Enabled
}
return c.Backend == "cloud-proxy"
return c.Backend == model.CloudProxyBackend
}
// PIIDetectors returns the names of the token-classification models that
@@ -672,7 +702,7 @@ var piiCoverableUsecases = []ModelConfigUsecase{FLAG_CHAT, FLAG_COMPLETION, FLAG
// false naturally: HasUsecases short-circuits to false for any usecase a
// declared score/token_classify model did not itself declare.
func (c *ModelConfig) PIIFilterApplies() bool {
if c.Backend == "cloud-proxy" {
if c.Backend == model.CloudProxyBackend {
return true
}
return slices.ContainsFunc(piiCoverableUsecases, c.HasUsecases)
@@ -770,6 +800,18 @@ type MCPSTDIOServer struct {
Command string `json:"command,omitempty"`
}
// Pipeline stage names. They match the Pipeline yaml keys and are the stage
// identifiers the realtime endpoint routes by (failover chains per stage,
// model_failover events, preload roles), so every user shares one spelling.
const (
PipelineStageVAD = "vad"
PipelineStageTranscription = "transcription"
PipelineStageLLM = "llm"
PipelineStageTTS = "tts"
PipelineStageSoundDetection = "sound_detection"
PipelineStageVoiceRecognition = "voice_recognition"
)
// @Description Pipeline defines other models to use for audio-to-audio
type Pipeline struct {
TTS string `yaml:"tts,omitempty" json:"tts,omitempty"`
@@ -1644,6 +1686,13 @@ func (c *ModelConfig) Validate() (bool, error) {
return false, fmt.Errorf("a config with artifacts must declare exactly one %q target, found %d", modelartifacts.TargetModel, primaries)
}
if c.IsFailover() {
if err := c.validateFailover(); err != nil {
return false, err
}
return true, nil
}
// An alias is a pure redirect: validate only its own shape here. Target
// existence and the no-chain rule need the full config set, so the loader
// (load-time) and the create/swap endpoints enforce those.
+163
View File
@@ -0,0 +1,163 @@
package config
import (
"fmt"
"time"
"github.com/mudler/LocalAI/pkg/model"
)
// FailoverConfig turns a model config into a failover chain: requests for the
// chain name are served by its highest-priority healthy target. See
// core/services/failover for the runtime side.
type FailoverConfig struct {
Targets []FailoverTarget `yaml:"targets" json:"targets"`
Probe FailoverProbe `yaml:"probe,omitempty" json:"probe,omitempty"`
Trip FailoverTrip `yaml:"trip,omitempty" json:"trip,omitempty"`
Recovery FailoverRecovery `yaml:"recovery,omitempty" json:"recovery,omitempty"`
}
type FailoverTarget struct {
Model string `yaml:"model" json:"model"`
// Warm keeps a local target loaded and exempt from eviction, so a switch
// does not wait for a cold load.
Warm bool `yaml:"warm,omitempty" json:"warm,omitempty"`
}
type FailoverProbe struct {
Interval string `yaml:"interval,omitempty" json:"interval,omitempty"`
Timeout string `yaml:"timeout,omitempty" json:"timeout,omitempty"`
}
type FailoverTrip struct {
Errors int `yaml:"errors,omitempty" json:"errors,omitempty"`
Window string `yaml:"window,omitempty" json:"window,omitempty"`
}
type FailoverRecovery struct {
Probes int `yaml:"probes,omitempty" json:"probes,omitempty"`
MinDwell string `yaml:"min_dwell,omitempty" json:"min_dwell,omitempty"`
}
const (
DefaultFailoverProbeInterval = 15 * time.Second
DefaultFailoverProbeTimeout = 5 * time.Second
DefaultFailoverTripErrors = 1
DefaultFailoverTripWindow = 30 * time.Second
DefaultFailoverRecoveryProbes = 3
DefaultFailoverMinDwell = 60 * time.Second
)
// IsFailover reports whether this config is a failover chain.
func (c ModelConfig) IsFailover() bool { return c.Failover != nil }
// IsRemoteProxy reports whether the model is served by a remote upstream
// through a proxy backend. Failover probes such a target over HTTP, and
// `warm` has no effect on it.
func (c ModelConfig) IsRemoteProxy() bool {
switch c.Backend {
case model.CloudProxyBackend, model.LocalAIProxyBackend:
return true
}
return false
}
func (f FailoverConfig) ProbeInterval() time.Duration {
return durationOr(f.Probe.Interval, DefaultFailoverProbeInterval)
}
func (f FailoverConfig) ProbeTimeout() time.Duration {
return durationOr(f.Probe.Timeout, DefaultFailoverProbeTimeout)
}
func (f FailoverConfig) TripWindow() time.Duration {
return durationOr(f.Trip.Window, DefaultFailoverTripWindow)
}
func (f FailoverConfig) MinDwell() time.Duration {
return durationOr(f.Recovery.MinDwell, DefaultFailoverMinDwell)
}
func (f FailoverConfig) TripErrors() int {
if f.Trip.Errors <= 0 {
return DefaultFailoverTripErrors
}
return f.Trip.Errors
}
func (f FailoverConfig) RecoveryProbes() int {
if f.Recovery.Probes <= 0 {
return DefaultFailoverRecoveryProbes
}
return f.Recovery.Probes
}
// WarmFailoverTargets returns the targets marked warm, in chain order.
func (c ModelConfig) WarmFailoverTargets() []string {
if c.Failover == nil {
return nil
}
var out []string
for _, t := range c.Failover.Targets {
if t.Warm {
out = append(out, t.Model)
}
}
return out
}
func durationOr(s string, def time.Duration) time.Duration {
if s == "" {
return def
}
d, err := time.ParseDuration(s)
if err != nil || d <= 0 {
return def
}
return d
}
// validateFailover checks what a chain can check without other configs.
// Target existence is checked by ModelConfigLoader.ValidateFailoverTargets.
func (c *ModelConfig) validateFailover() error {
if c.Name == "" {
return fmt.Errorf("failover config requires a name")
}
if c.IsAlias() {
return fmt.Errorf("model %q cannot set both alias and failover", c.Name)
}
if c.Backend != "" || c.Model != "" {
return fmt.Errorf("failover config %q must not set backend or parameters.model: a chain is a pure redirect", c.Name)
}
f := c.Failover
if len(f.Targets) < 2 {
return fmt.Errorf("failover chain %q needs at least 2 targets", c.Name)
}
seen := map[string]bool{}
for _, t := range f.Targets {
switch {
case t.Model == "":
return fmt.Errorf("failover chain %q has a target with no model", c.Name)
case t.Model == c.Name:
return fmt.Errorf("failover chain %q cannot list itself", c.Name)
case seen[t.Model]:
return fmt.Errorf("failover chain %q lists %q twice", c.Name, t.Model)
}
seen[t.Model] = true
}
for key, v := range map[string]string{
"probe.interval": f.Probe.Interval,
"probe.timeout": f.Probe.Timeout,
"trip.window": f.Trip.Window,
"recovery.min_dwell": f.Recovery.MinDwell,
} {
if v == "" {
continue
}
if d, err := time.ParseDuration(v); err != nil || d <= 0 {
return fmt.Errorf("failover chain %q: invalid %s %q", c.Name, key, v)
}
}
if f.Trip.Errors < 0 {
return fmt.Errorf("failover chain %q: trip.errors must not be negative", c.Name)
}
if f.Recovery.Probes < 0 {
return fmt.Errorf("failover chain %q: recovery.probes must not be negative", c.Name)
}
return nil
}
+101
View File
@@ -0,0 +1,101 @@
package config
import (
"os"
"path/filepath"
"time"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"gopkg.in/yaml.v3"
)
var _ = Describe("ModelConfig failover", func() {
chain := func(targets ...string) ModelConfig {
c := ModelConfig{Name: "chain", Failover: &FailoverConfig{}}
for _, t := range targets {
c.Failover.Targets = append(c.Failover.Targets, FailoverTarget{Model: t})
}
return c
}
It("parses the YAML block and applies defaults", func() {
var c ModelConfig
Expect(yaml.Unmarshal([]byte(`
name: assistant-llm
failover:
targets:
- model: argus-llm
- model: gemma-local
warm: true
recovery:
probes: 5
`), &c)).To(Succeed())
Expect(c.IsFailover()).To(BeTrue())
Expect(c.Failover.Targets).To(Equal([]FailoverTarget{{Model: "argus-llm"}, {Model: "gemma-local", Warm: true}}))
Expect(c.Failover.ProbeInterval()).To(Equal(15 * time.Second))
Expect(c.Failover.ProbeTimeout()).To(Equal(5 * time.Second))
Expect(c.Failover.TripErrors()).To(Equal(1))
Expect(c.Failover.TripWindow()).To(Equal(30 * time.Second))
Expect(c.Failover.RecoveryProbes()).To(Equal(5))
Expect(c.Failover.MinDwell()).To(Equal(60 * time.Second))
Expect(c.WarmFailoverTargets()).To(Equal([]string{"gemma-local"}))
})
It("accepts a valid chain", func() {
c := chain("a", "b")
ok, err := c.Validate()
Expect(err).ToNot(HaveOccurred())
Expect(ok).To(BeTrue())
})
DescribeTable("rejects invalid chains",
func(mutate func(*ModelConfig), want string) {
c := chain("a", "b")
mutate(&c)
ok, err := c.Validate()
Expect(ok).To(BeFalse())
Expect(err).To(MatchError(ContainSubstring(want)))
},
Entry("alias and failover", func(c *ModelConfig) { c.Alias = "x" }, "both alias and failover"),
Entry("backend set", func(c *ModelConfig) { c.Backend = "llama-cpp" }, "must not set backend"),
Entry("one target", func(c *ModelConfig) { c.Failover.Targets = c.Failover.Targets[:1] }, "at least 2 targets"),
Entry("empty target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "" }, "no model"),
Entry("self target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "chain" }, "cannot list itself"),
Entry("duplicate target", func(c *ModelConfig) { c.Failover.Targets[1].Model = "a" }, "twice"),
Entry("bad duration", func(c *ModelConfig) { c.Failover.Probe.Interval = "soon" }, "invalid probe.interval"),
Entry("negative errors", func(c *ModelConfig) { c.Failover.Trip.Errors = -1 }, "trip.errors"),
Entry("no name", func(c *ModelConfig) { c.Name = "" }, "requires a name"),
)
})
var _ = Describe("ProxyConfig.ResolveAPIKey", func() {
It("reads the key through the application lookup", func() {
app := NewApplicationConfig(WithProxyAPIKeyEnvLookup(func(name string) string {
Expect(name).To(Equal("FAILOVER_TEST_KEY"))
return "k1"
}))
Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(app.ProxyAPIKeyEnvLookup)).To(Equal("k1"))
})
It("fails without an environment lookup", func() {
_, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(nil)
Expect(err).To(HaveOccurred())
})
It("fails on an unset env var", func() {
_, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey(os.Getenv)
Expect(err).To(HaveOccurred())
})
It("fails on a set-but-empty env var", func() {
GinkgoT().Setenv("FAILOVER_TEST_EMPTY_KEY", "")
_, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_EMPTY_KEY"}.ResolveAPIKey(os.Getenv)
Expect(err).To(HaveOccurred())
})
It("reads and trims the key file", func() {
f := filepath.Join(GinkgoT().TempDir(), "key")
Expect(os.WriteFile(f, []byte(" k2\n"), 0o600)).To(Succeed())
Expect(ProxyConfig{APIKeyFile: f}.ResolveAPIKey(os.Getenv)).To(Equal("k2"))
})
It("returns empty when nothing is set", func() {
Expect(ProxyConfig{}.ResolveAPIKey(os.Getenv)).To(Equal(""))
})
})
+138
View File
@@ -17,6 +17,7 @@ import (
"github.com/charmbracelet/glamour"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/pkg/downloader"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/modelartifacts"
"github.com/mudler/LocalAI/pkg/safefile"
"github.com/mudler/LocalAI/pkg/utils"
@@ -501,6 +502,102 @@ func (bcl *ModelConfigLoader) ValidateAliasTarget(cfg *ModelConfig) error {
return nil
}
// failoverUsecases are the single usecases a chain can share. Checking one
// flag at a time avoids treating "chat+tts" and "tts" as unrelated.
var failoverUsecases = []ModelConfigUsecase{
FLAG_CHAT, FLAG_COMPLETION, FLAG_EMBEDDINGS, FLAG_RERANK, FLAG_IMAGE,
FLAG_TRANSCRIPT, FLAG_TTS, FLAG_SOUND_GENERATION, FLAG_VAD, FLAG_VIDEO,
FLAG_SOUND_CLASSIFICATION,
}
// ValidateFailoverTargets checks that every target of a chain exists and is
// not itself a chain. Alias targets are allowed and resolve one hop.
func (bcl *ModelConfigLoader) ValidateFailoverTargets(cfg *ModelConfig) error {
return validateFailoverTargets(cfg, bcl.GetModelConfig)
}
// FailoverTargetsShareUsecase reports whether all targets of a chain have at
// least one usecase in common. A false result is only a warning: usecases are
// often inferred.
func (bcl *ModelConfigLoader) FailoverTargetsShareUsecase(cfg *ModelConfig) bool {
return failoverTargetsShareUsecase(cfg, bcl.GetModelConfig)
}
func validateFailoverTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) error {
if cfg == nil || !cfg.IsFailover() {
return nil
}
for _, t := range cfg.Failover.Targets {
target, ok := lookup(t.Model)
if !ok {
return fmt.Errorf("failover chain %q: target %q does not exist", cfg.Name, t.Model)
}
if target.IsAlias() {
if resolved, ok := lookup(target.Alias); ok {
target = resolved
}
}
if target.IsFailover() {
return fmt.Errorf("failover chain %q: target %q is a chain (chains do not nest)", cfg.Name, t.Model)
}
}
return nil
}
// failoverWarmRemoteTargets lists the chain's targets marked warm that are
// remote. Warm only keeps a local model loaded, so the flag does nothing there.
func failoverWarmRemoteTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) []string {
if cfg == nil || !cfg.IsFailover() {
return nil
}
var out []string
for _, t := range cfg.Failover.Targets {
if !t.Warm {
continue
}
target, ok := lookup(t.Model)
if ok && target.IsAlias() {
target, ok = lookup(target.Alias)
}
if ok && target.IsRemoteProxy() {
out = append(out, t.Model)
}
}
return out
}
func failoverTargetsShareUsecase(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) bool {
if cfg == nil || !cfg.IsFailover() {
return true
}
var targets []ModelConfig
for _, t := range cfg.Failover.Targets {
target, ok := lookup(t.Model)
if !ok {
return true // missing targets are reported by validateFailoverTargets
}
if target.IsAlias() {
if resolved, ok := lookup(target.Alias); ok {
target = resolved
}
}
targets = append(targets, target)
}
for _, u := range failoverUsecases {
all := true
for i := range targets {
if !targets[i].HasUsecases(u) {
all = false
break
}
}
if all {
return true
}
}
return false
}
type preloadWork struct {
key string
config ModelConfig
@@ -837,6 +934,47 @@ func (bcl *ModelConfigLoader) loadModelConfigsFromPath(path string, strict bool,
}
}
// Reject failover chains whose targets are missing or are themselves
// chains. bcl.Lock() is held here, so look up configs directly rather
// than through GetModelConfig, which would deadlock on the same mutex.
lookup := func(n string) (ModelConfig, bool) { c, ok := bcl.configs[n]; return c, ok }
for name, cfg := range bcl.configs {
if !cfg.IsFailover() {
continue
}
c := cfg
if err := validateFailoverTargets(&c, lookup); err != nil {
if strict {
return fmt.Errorf("invalid model config %q: %w", name, err)
}
xlog.Error("skipping invalid failover chain", "model", name, "error", err)
delete(bcl.configs, name)
continue
}
if !failoverTargetsShareUsecase(&c, lookup) {
xlog.Warn("failover chain targets share no known usecase", "model", name)
}
if remote := failoverWarmRemoteTargets(&c, lookup); len(remote) > 0 {
xlog.Warn("failover chain: warm has no effect on remote targets", "model", name, "targets", remote)
}
}
// localai-proxy proxies straight through to another LocalAI instance, so
// it never translates the wire protocol the way cloud-proxy does; warn
// when a config carries settings that only make sense there, or lacks
// the usecases failover's own usecase-sharing check depends on.
for name, cfg := range bcl.configs {
if cfg.Backend != model.LocalAIProxyBackend {
continue
}
if cfg.Proxy.Mode == ProxyModeTranslate || cfg.Proxy.Provider != "" {
xlog.Warn("localai-proxy backend proxies to another LocalAI instance and ignores proxy.mode/proxy.provider", "model", name)
}
if len(cfg.KnownUsecaseStrings) == 0 {
xlog.Warn("localai-proxy config has no known_usecases; failover usecase matching will skip it", "model", name)
}
}
return nil
}
+127
View File
@@ -1,7 +1,9 @@
package config
import (
"bytes"
"context"
"log/slog"
"os"
"path/filepath"
@@ -9,6 +11,7 @@ import (
. "github.com/onsi/gomega"
"github.com/mudler/LocalAI/pkg/modelartifacts"
"github.com/mudler/xlog"
)
type preloadArtifactMaterializer struct {
@@ -402,3 +405,127 @@ var _ = Describe("ModelConfigLoader ResolveAliasName", func() {
Expect(target).To(BeEmpty())
})
})
var _ = Describe("ModelConfigLoader failover validation", func() {
var loader *ModelConfigLoader
chain := func(targets ...string) *ModelConfig {
c := &ModelConfig{Name: "chain", Failover: &FailoverConfig{}}
for _, t := range targets {
c.Failover.Targets = append(c.Failover.Targets, FailoverTarget{Model: t})
}
return c
}
BeforeEach(func() {
loader = NewModelConfigLoader("")
loader.configs["a"] = ModelConfig{Name: "a", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}}
loader.configs["b"] = ModelConfig{Name: "b", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}}
loader.configs["tts"] = ModelConfig{Name: "tts", Backend: "piper", KnownUsecaseStrings: []string{"tts"}}
loader.configs["alias-b"] = ModelConfig{Name: "alias-b", Alias: "b"}
loader.configs["other-chain"] = *chain("a", "b")
loader.configs["alias-chain"] = ModelConfig{Name: "alias-chain", Alias: "other-chain"}
for k, c := range loader.configs {
c.KnownUsecases = GetUsecasesFromYAML(c.KnownUsecaseStrings)
loader.configs[k] = c
}
})
It("accepts existing targets and alias targets", func() {
Expect(loader.ValidateFailoverTargets(chain("a", "alias-b"))).To(Succeed())
})
It("rejects a missing target", func() {
Expect(loader.ValidateFailoverTargets(chain("a", "nope"))).To(MatchError(ContainSubstring("does not exist")))
})
It("rejects a nested chain, directly or through an alias", func() {
Expect(loader.ValidateFailoverTargets(chain("a", "other-chain"))).To(MatchError(ContainSubstring("chains do not nest")))
Expect(loader.ValidateFailoverTargets(chain("a", "alias-chain"))).To(MatchError(ContainSubstring("chains do not nest")))
})
It("reports whether targets share a usecase", func() {
Expect(loader.FailoverTargetsShareUsecase(chain("a", "b"))).To(BeTrue())
Expect(loader.FailoverTargetsShareUsecase(chain("a", "tts"))).To(BeFalse())
})
It("finds warm targets that are remote, where warm has no effect", func() {
loader.configs["remote"] = ModelConfig{Name: "remote", Backend: "cloud-proxy"}
loader.configs["alias-remote"] = ModelConfig{Name: "alias-remote", Alias: "remote"}
c := chain("remote", "alias-remote", "a")
for i := range c.Failover.Targets {
c.Failover.Targets[i].Warm = true
}
Expect(failoverWarmRemoteTargets(c, loader.GetModelConfig)).To(Equal([]string{"remote", "alias-remote"}))
Expect(failoverWarmRemoteTargets(chain("remote", "a"), loader.GetModelConfig)).To(BeEmpty())
})
})
var _ = Describe("ModelConfigLoader localai-proxy load-time warnings", func() {
var captured *bytes.Buffer
BeforeEach(func() {
captured = &bytes.Buffer{}
handler := slog.NewTextHandler(captured, &slog.HandlerOptions{Level: slog.LevelWarn})
xlog.SetLogger(xlog.NewLoggerWithHandler(handler, xlog.LogLevelWarn))
})
AfterEach(func() {
// xlog exposes no getter for the package logger, so restore the same
// default the suite entrypoint installs rather than the prior value.
xlog.SetLogger(xlog.NewLogger(xlog.LogLevel("info"), "text"))
})
It("warns and still loads when proxy.mode/proxy.provider are set (ignored by localai-proxy)", func() {
modelsPath := GinkgoT().TempDir()
cfgYAML := `
name: proxied
backend: localai-proxy
known_usecases: [chat]
proxy:
mode: translate
provider: openai
upstream_url: http://127.0.0.1:8081
`
Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed())
loader := NewModelConfigLoader(modelsPath)
Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed())
_, ok := loader.GetModelConfig("proxied")
Expect(ok).To(BeTrue())
Expect(captured.String()).To(ContainSubstring("proxy.mode/proxy.provider"))
Expect(captured.String()).To(ContainSubstring("proxied"))
})
It("warns and still loads when known_usecases is empty", func() {
modelsPath := GinkgoT().TempDir()
cfgYAML := `
name: proxied-no-usecase
backend: localai-proxy
proxy:
upstream_url: http://127.0.0.1:8081
`
Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed())
loader := NewModelConfigLoader(modelsPath)
Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed())
_, ok := loader.GetModelConfig("proxied-no-usecase")
Expect(ok).To(BeTrue())
Expect(captured.String()).To(ContainSubstring("known_usecases"))
Expect(captured.String()).To(ContainSubstring("proxied-no-usecase"))
})
It("does not warn when localai-proxy uses only passthrough and known_usecases", func() {
modelsPath := GinkgoT().TempDir()
cfgYAML := `
name: proxied-clean
backend: localai-proxy
known_usecases: [chat]
proxy:
upstream_url: http://127.0.0.1:8081
`
Expect(os.WriteFile(filepath.Join(modelsPath, "proxied.yaml"), []byte(cfgYAML), 0o600)).To(Succeed())
loader := NewModelConfigLoader(modelsPath)
Expect(loader.LoadModelConfigsFromPath(modelsPath)).To(Succeed())
Expect(captured.String()).To(BeEmpty())
})
})
+1
View File
@@ -475,6 +475,7 @@ func API(application *application.Application) (*echo.Echo, error) {
mcpJobsMw := auth.RequireFeature(application.AuthDB(), auth.FeatureMCPJobs)
requestExtractor := httpMiddleware.NewRequestExtractor(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig())
requestExtractor.SetFailoverManager(application.FailoverManager())
// Register auth routes (login, callback, API keys, user management)
routes.RegisterAuthRoutes(e, application)
+7
View File
@@ -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)
+35
View File
@@ -223,6 +223,17 @@ var _ = Describe("Auth Middleware", func() {
Expect(rec.Code).To(Equal(http.StatusOK))
})
It("allows requests to the failover status endpoint with a valid session", func() {
sessionID := createTestSession(db, user.ID)
rec := doRequest(app, http.MethodGet, "/api/failover", withSessionCookie(sessionID))
Expect(rec.Code).To(Equal(http.StatusOK))
})
It("returns 401 for the failover status endpoint without credentials", func() {
rec := doRequest(app, http.MethodGet, "/api/failover")
Expect(rec.Code).To(Equal(http.StatusUnauthorized))
})
It("allows authenticated users to call moderation by default", func() {
sessionID := createTestSession(db, user.ID)
rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID))
@@ -526,6 +537,30 @@ var _ = Describe("Auth Middleware", func() {
Expect(rec.Code).To(Equal(http.StatusForbidden))
})
It("allows admin to pin and unpin failover chains", func() {
admin := createTestUser(db, "admin5@example.com", auth.RoleAdmin, auth.ProviderGitHub)
sessionID := createTestSession(db, admin.ID)
app := newAdminTestApp(db, appConfig)
rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID))
Expect(rec.Code).To(Equal(http.StatusOK))
rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID))
Expect(rec.Code).To(Equal(http.StatusOK))
})
It("blocks non-admin from pinning or unpinning failover chains", func() {
user := createTestUser(db, "user5@example.com", auth.RoleUser, auth.ProviderGitHub)
sessionID := createTestSession(db, user.ID)
app := newAdminTestApp(db, appConfig)
rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID))
Expect(rec.Code).To(Equal(http.StatusForbidden))
rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID))
Expect(rec.Code).To(Equal(http.StatusForbidden))
})
It("allows user to access regular inference endpoints", func() {
user := createTestUser(db, "user@example.com", auth.RoleUser, auth.ProviderGitHub)
sessionID := createTestSession(db, user.ID)
@@ -129,6 +129,12 @@ var instructionDefs = []instructionDef{
Tags: []string{"middleware", "pii", "router"},
Intro: "GET /api/middleware/status is the single round-trip the /app/middleware admin page reads to render the current state: every model's resolved PII enabled state and the NER detector models it references, recent event count, and the active routing models with their classifier configurations. Admin-only (the synthetic local user is admin in no-auth mode). PII detection policy is edited on each detector model's `pii_detection:` block via the model-config tools/UI — there is no global pattern set to mutate. GET /api/router/decisions returns the routing decision log filtered by correlation_id / user_id / router_model. The same surface is exposed as MCP tools (`get_middleware_status`, `get_pii_events`, `get_router_decisions`) for agent-driven inspection.",
},
{
Name: "failover",
Description: "Model failover chains: target health, pinning and switch events",
Tags: []string{"failover"},
Intro: "A failover chain is a model config with a failover block. Requests for the chain name are served by its highest-priority healthy target; the X-LocalAI-Served-Model response header names it. Subscribe to GET /api/failover/events (SSE) to follow switches.",
},
{
Name: "intelligent-routing",
Description: "Per-model `router:` configuration that classifies requests and rewrites the served model",
@@ -39,7 +39,7 @@ var _ = Describe("API Instructions Endpoints", func() {
instructions, ok := resp["instructions"].([]any)
Expect(ok).To(BeTrue())
Expect(instructions).To(HaveLen(19))
Expect(instructions).To(HaveLen(20))
// Verify each instruction has required fields and correct URL format
for _, s := range instructions {
@@ -81,6 +81,7 @@ var _ = Describe("API Instructions Endpoints", func() {
"intelligent-routing",
"voice-library",
"3d",
"failover",
))
})
})
+169
View File
@@ -0,0 +1,169 @@
package localai
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"time"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
)
type FailoverChainsResponse struct {
Chains []failover.ChainStatus `json:"chains"`
}
type FailoverPinRequest struct {
Target string `json:"target"`
}
func failoverError(c echo.Context, code int, msg string) error {
return c.JSON(code, schema.ErrorResponse{Error: &schema.APIError{Message: msg, Code: code, Type: "failover_error"}})
}
// ListFailoverChainsEndpoint lists failover chains and the health of their targets
//
// @Summary List failover chains and the health of their targets
// @Tags failover
// @Produce json
// @Success 200 {object} FailoverChainsResponse
// @Router /api/failover [get]
func ListFailoverChainsEndpoint(fm *failover.Manager) echo.HandlerFunc {
return func(c echo.Context) error {
return c.JSON(http.StatusOK, FailoverChainsResponse{Chains: fm.Status()})
}
}
// GetFailoverChainEndpoint returns one failover chain
//
// @Summary Get one failover chain
// @Tags failover
// @Produce json
// @Param chain path string true "Chain name"
// @Success 200 {object} failover.ChainStatus
// @Failure 404 {object} schema.ErrorResponse
// @Router /api/failover/{chain} [get]
func GetFailoverChainEndpoint(fm *failover.Manager) echo.HandlerFunc {
return func(c echo.Context) error {
st, ok := fm.ChainStatus(c.Param("chain"))
if !ok {
return failoverError(c, http.StatusNotFound, fmt.Sprintf("failover chain %q not found", c.Param("chain")))
}
return c.JSON(http.StatusOK, st)
}
}
// PinFailoverTargetEndpoint forces a chain to one target
//
// @Summary Pin a failover chain to one target
// @Tags failover
// @Accept json
// @Produce json
// @Param chain path string true "Chain name"
// @Param request body FailoverPinRequest true "Target to pin"
// @Success 200 {object} failover.ChainStatus
// @Failure 400 {object} schema.ErrorResponse
// @Failure 404 {object} schema.ErrorResponse
// @Router /api/failover/{chain}/pin [post]
func PinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc {
return func(c echo.Context) error {
var req FailoverPinRequest
if err := c.Bind(&req); err != nil || req.Target == "" {
return failoverError(c, http.StatusBadRequest, "request body must set \"target\"")
}
chain := c.Param("chain")
if err := fm.Pin(chain, req.Target); err != nil {
return pinError(c, err)
}
st, _ := fm.ChainStatus(chain)
return c.JSON(http.StatusOK, st)
}
}
// UnpinFailoverTargetEndpoint removes a pin
//
// @Summary Remove the pin from a failover chain
// @Tags failover
// @Produce json
// @Param chain path string true "Chain name"
// @Success 200 {object} failover.ChainStatus
// @Failure 404 {object} schema.ErrorResponse
// @Router /api/failover/{chain}/pin [delete]
func UnpinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc {
return func(c echo.Context) error {
chain := c.Param("chain")
if err := fm.Unpin(chain); err != nil {
return pinError(c, err)
}
st, _ := fm.ChainStatus(chain)
return c.JSON(http.StatusOK, st)
}
}
func pinError(c echo.Context, err error) error {
switch {
case errors.Is(err, failover.ErrChainNotFound):
return failoverError(c, http.StatusNotFound, err.Error())
case errors.Is(err, failover.ErrTargetNotInChain):
return failoverError(c, http.StatusBadRequest, err.Error())
}
return failoverError(c, http.StatusInternalServerError, err.Error())
}
// FailoverEventsEndpoint streams failover events
//
// @Summary Stream failover events (server-sent events)
// @Description The first event is "snapshot" with the full state, then "chain.switched" and "target.state" events.
// @Tags failover
// @Produce text/event-stream
// @Success 200
// @Router /api/failover/events [get]
func FailoverEventsEndpoint(fm *failover.Manager) echo.HandlerFunc {
return func(c echo.Context) error {
// Subscribe before the snapshot so no event falls between the two.
events, cancel := fm.Subscribe(64)
defer cancel()
w := c.Response()
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
w.WriteHeader(http.StatusOK)
send := func(name string, v any) error {
data, err := json.Marshal(v)
if err != nil {
return err
}
if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", name, data); err != nil {
return err
}
w.Flush()
return nil
}
if err := send("snapshot", FailoverChainsResponse{Chains: fm.Status()}); err != nil {
return nil
}
keepalive := time.NewTicker(15 * time.Second)
defer keepalive.Stop()
for {
select {
case <-c.Request().Context().Done():
return nil
case <-keepalive.C:
if _, err := fmt.Fprint(w, ": keepalive\n\n"); err != nil {
return nil
}
w.Flush()
case ev, ok := <-events:
if !ok {
return nil
}
if err := send(string(ev.Type), ev); err != nil {
return nil
}
}
}
}
}
@@ -0,0 +1,70 @@
package localai
import (
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/http/middleware"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/system"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"google.golang.org/grpc/codes"
grpcstatus "google.golang.org/grpc/status"
)
// The non-OpenAI endpoints (depth, detection, face_*, voice_*, images, video,
// 3d) return mapBackendError's echo 501 for a backend without the method, not
// the raw gRPC status. The failover retry must still see a capability gap.
var _ = Describe("mapBackendError under a failover chain", func() {
It("spills a backend's Unimplemented to the next target without tripping it", func() {
dir := GinkgoT().TempDir()
write := func(name, body string) {
Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed())
}
write("a", "name: a\nbackend: fake-a\n")
write("b", "name: b\nbackend: fake-b\n")
write("chain", "name: chain\nfailover:\n targets:\n - model: a\n - model: b\n")
ss := &system.SystemState{Model: system.Model{ModelsPath: dir}}
appConfig := config.NewApplicationConfig()
appConfig.SystemState = ss
mcl := config.NewModelConfigLoader(dir)
Expect(mcl.LoadModelConfigsFromPath(dir)).To(Succeed())
re := middleware.NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig)
fm := failover.New(mcl)
fm.Sync()
re.SetFailoverManager(fm)
var calls []string
handler := func(c echo.Context) error {
cfg := c.Get(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
calls = append(calls, cfg.Name)
if cfg.Name == "a" {
return mapBackendError(grpcstatus.Error(codes.Unimplemented, "unimplemented: Detect"))
}
return c.JSON(http.StatusOK, map[string]string{"served": cfg.Name})
}
app := echo.New()
app.POST("/v1/detection", handler,
re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.DetectionRequest) }))
req := httptest.NewRequest(http.MethodPost, "/v1/detection", strings.NewReader(`{"model":"chain","image":"x"}`))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
app.ServeHTTP(rec, req)
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(calls).To(Equal([]string{"a", "b"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
})
@@ -0,0 +1,110 @@
package localai
import (
"bufio"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"time"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/services/failover"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
type mapSource map[string]config.ModelConfig
func (s mapSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok }
func (s mapSource) GetAllModelsConfigs() []config.ModelConfig {
var out []config.ModelConfig
for _, c := range s {
out = append(out, c)
}
return out
}
var _ = Describe("failover endpoints", func() {
var (
e *echo.Echo
fm *failover.Manager
)
BeforeEach(func() {
src := mapSource{
"a": {Name: "a", Backend: "cloud-proxy"},
"b": {Name: "b", Backend: "llama-cpp"},
"chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}},
}
fm = failover.New(src)
e = echo.New()
e.GET("/api/failover", ListFailoverChainsEndpoint(fm))
e.GET("/api/failover/events", FailoverEventsEndpoint(fm))
e.GET("/api/failover/:chain", GetFailoverChainEndpoint(fm))
e.POST("/api/failover/:chain/pin", PinFailoverTargetEndpoint(fm))
e.DELETE("/api/failover/:chain/pin", UnpinFailoverTargetEndpoint(fm))
})
do := func(method, path, body string) *httptest.ResponseRecorder {
req := httptest.NewRequest(method, path, strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
e.ServeHTTP(rec, req)
return rec
}
It("lists chains", func() {
rec := do(http.MethodGet, "/api/failover", "")
Expect(rec.Code).To(Equal(http.StatusOK))
var out FailoverChainsResponse
Expect(json.Unmarshal(rec.Body.Bytes(), &out)).To(Succeed())
Expect(out.Chains).To(HaveLen(1))
Expect(out.Chains[0].Active).To(Equal("a"))
})
It("gets one chain or 404", func() {
Expect(do(http.MethodGet, "/api/failover/chain", "").Code).To(Equal(http.StatusOK))
Expect(do(http.MethodGet, "/api/failover/nope", "").Code).To(Equal(http.StatusNotFound))
})
It("pins and unpins", func() {
rec := do(http.MethodPost, "/api/failover/chain/pin", `{"target":"b"}`)
Expect(rec.Code).To(Equal(http.StatusOK))
st, _ := fm.ChainStatus("chain")
Expect(st.Active).To(Equal("b"))
Expect(do(http.MethodPost, "/api/failover/chain/pin", `{"target":"zzz"}`).Code).To(Equal(http.StatusBadRequest))
Expect(do(http.MethodPost, "/api/failover/chain/pin", `{}`).Code).To(Equal(http.StatusBadRequest))
Expect(do(http.MethodPost, "/api/failover/nope/pin", `{"target":"a"}`).Code).To(Equal(http.StatusNotFound))
Expect(do(http.MethodDelete, "/api/failover/chain/pin", "").Code).To(Equal(http.StatusOK))
st, _ = fm.ChainStatus("chain")
Expect(st.Pinned).To(BeNil())
})
It("streams a snapshot, then switch events", func() {
srv := httptest.NewServer(e)
defer srv.Close()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil)
resp, err := http.DefaultClient.Do(req)
Expect(err).ToNot(HaveOccurred())
defer func() { _ = resp.Body.Close() }()
Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream"))
r := bufio.NewReader(resp.Body)
next := func() string {
for {
line, err := r.ReadString('\n')
Expect(err).ToNot(HaveOccurred())
if strings.HasPrefix(line, "event: ") {
return strings.TrimSpace(strings.TrimPrefix(line, "event: "))
}
}
}
Expect(next()).To(Equal("snapshot"))
Expect(fm.Pin("chain", "b")).To(Succeed())
Expect(next()).To(Equal("chain.switched"))
})
})
@@ -187,6 +187,12 @@ func ImportModelEndpoint(cl *config.ModelConfigLoader, gs *galleryop.GalleryServ
return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()})
}
// Reject failover chains whose targets are missing or are themselves
// chains, for the same reason.
if err := cl.ValidateFailoverTargets(&modelConfig); err != nil {
return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()})
}
// Create the configuration file
configPath := filepath.Join(appConfig.SystemState.Model.ModelsPath, modelConfig.Name+".yaml")
if err := utils.VerifyPath(modelConfig.Name+".yaml", appConfig.SystemState.Model.ModelsPath); err != nil {
@@ -214,3 +214,11 @@ func (stubClient) SeedRouterCorpus(_ context.Context, req localaitools.RouterCor
func (stubClient) ClearRouterCorpus(_ context.Context, routerModel string) (*localaitools.RouterCorpusClearResult, error) {
return &localaitools.RouterCorpusClearResult{Router: routerModel}, nil
}
func (stubClient) ListFailoverChains(_ context.Context) ([]localaitools.FailoverChainInfo, error) {
return []localaitools.FailoverChainInfo{}, nil
}
func (stubClient) PinFailoverTarget(_ context.Context, _, _ string) error { return nil }
func (stubClient) UnpinFailoverTarget(_ context.Context, _ string) error { return nil }
@@ -0,0 +1,78 @@
package openai
import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/http/middleware"
"github.com/mudler/LocalAI/core/schema"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
// A broken or missing upload is the caller's fault: the audio endpoints must
// answer 400, not surface the multipart parser error as a 500.
var _ = Describe("audio upload endpoints reject bad uploads as client errors", func() {
type endpointCase struct {
path string
handler func() echo.HandlerFunc
}
cases := map[string]endpointCase{
"transcription": {"/v1/audio/transcriptions", func() echo.HandlerFunc {
return TranscriptEndpoint(nil, nil, config.NewApplicationConfig())
}},
"diarization": {"/v1/audio/diarization", func() echo.HandlerFunc {
return DiarizationEndpoint(nil, nil, config.NewApplicationConfig())
}},
"sound classification": {"/v1/audio/classifications", func() echo.HandlerFunc {
return SoundClassificationEndpoint(nil, nil, config.NewApplicationConfig())
}},
}
run := func(ec endpointCase, contentType, body string) error {
req := httptest.NewRequest(http.MethodPost, ec.path, strings.NewReader(body))
req.Header.Set("Content-Type", contentType)
c := echo.New().NewContext(req, httptest.NewRecorder())
c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.OpenAIRequest{PredictionOptions: schema.PredictionOptions{BasicModelRequest: schema.BasicModelRequest{Model: "m"}}})
c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{})
return ec.handler()(c)
}
expectBadRequest := func(err error) {
Expect(err).To(HaveOccurred())
var he *echo.HTTPError
Expect(errors.As(err, &he)).To(BeTrue(), "expected *echo.HTTPError, got %T: %v", err, err)
Expect(he.Code).To(Equal(http.StatusBadRequest))
}
// An unknown response_format is the caller's fault too, and must be
// rejected before the backend runs: a failover chain would otherwise
// transcribe on every target and count the error against each of them.
for name, ec := range map[string]endpointCase{
"transcription": cases["transcription"],
"diarization": cases["diarization"],
} {
It(name+": unknown response_format", func() {
body := "--xyz\r\nContent-Disposition: form-data; name=\"response_format\"\r\n\r\nbogus\r\n" +
"--xyz\r\nContent-Disposition: form-data; name=\"file\"; filename=\"a.wav\"\r\n\r\nRIFF\r\n--xyz--\r\n"
var err error
Expect(func() { err = run(ec, "multipart/form-data; boundary=xyz", body) }).NotTo(Panic(), "the backend must not be reached")
expectBadRequest(err)
})
}
for name, ec := range cases {
It(name+": multipart content type without a boundary", func() {
expectBadRequest(run(ec, "multipart/form-data", ""))
})
It(name+": multipart body without the file field", func() {
body := "--xyz\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\nm\r\n--xyz--\r\n"
expectBadRequest(run(ec, "multipart/form-data; boundary=xyz", body))
})
}
})
+9 -3
View File
@@ -1,7 +1,6 @@
package openai
import (
"errors"
"fmt"
"io"
"net/http"
@@ -75,8 +74,15 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap
if responseFormat == "" {
responseFormat = schema.DiarizationResponseFormatJson
}
switch responseFormat {
case schema.DiarizationResponseFormatJson, schema.DiarizationResponseFormatJsonVerbose, schema.DiarizationResponseFormatRTTM:
default:
// Checked before the backend runs, for the same reason as in
// TranscriptEndpoint.
return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format (expected: json, verbose_json, rttm)")
}
file, err := c.FormFile("file")
file, err := uploadedFile(c, "file")
if err != nil {
return err
}
@@ -126,7 +132,7 @@ func DiarizationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, ap
case schema.DiarizationResponseFormatJsonVerbose:
return c.JSON(http.StatusOK, result)
default:
return errors.New("invalid response_format (expected: json, verbose_json, rttm)")
return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format (expected: json, verbose_json, rttm)")
}
}
}
+18 -3
View File
@@ -32,6 +32,7 @@ import (
"github.com/mudler/LocalAI/core/http/endpoints/openai/turncoord"
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/routing/router"
"github.com/mudler/LocalAI/core/services/voiceprofile"
"github.com/mudler/LocalAI/core/templates"
@@ -638,6 +639,7 @@ func runRealtimeSession(application *application.Application, t Transport, model
application.ModelConfigLoader(),
application.ModelLoader(),
application.ApplicationConfig(),
application.FailoverManager(),
)
} else {
m, err = newModel(
@@ -709,7 +711,7 @@ func runRealtimeSession(application *application.Application, t Transport, model
var gateErr error
if session.voiceGate != nil {
_, gateErr = backend.PreloadStages(context.Background(), application.ModelLoader(), application.ApplicationConfig(), []backend.PreloadStage{
{Role: "voice_recognition", Cfg: session.voiceGate.recCfg},
{Role: config.PipelineStageVoiceRecognition, Cfg: session.voiceGate.recCfg},
})
}
if err := errors.Join(<-warmErr, gateErr); err != nil {
@@ -754,6 +756,13 @@ func runRealtimeSession(application *application.Application, t Transport, model
Session: session.ToServer(),
})
// Sent after session.created, which clients expect as the first event.
// This function runs until the connection closes, so the defer stops the
// events at session end. A transcription session.update that swaps the
// model restarts them for the new model's chains.
stopFailoverEvents := startModelFailoverEvents(t, m)
defer func() { stopFailoverEvents() }()
var (
msg []byte
wg sync.WaitGroup
@@ -824,12 +833,14 @@ func runRealtimeSession(application *application.Application, t Transport, model
// Handle transcription session update
if e.Session.Transcription != nil {
prevModel := session.ModelInterface
if err := updateTransSession(
session,
&e.Session,
application.ModelConfigLoader(),
application.ModelLoader(),
application.ApplicationConfig(),
application.FailoverManager(),
); err != nil {
xlog.Error("failed to update session", "error", err)
// The cause is validation feedback on the client's own
@@ -846,6 +857,10 @@ func runRealtimeSession(application *application.Application, t Transport, model
},
Session: session.ToServer(),
})
if session.ModelInterface != prevModel {
stopFailoverEvents()
stopFailoverEvents = startModelFailoverEvents(t, session.ModelInterface)
}
}
// Handle realtime session update
@@ -1143,7 +1158,7 @@ func sendTestTone(t Transport) {
}
}
func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) error {
func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) error {
sessionLock.Lock()
defer sessionLock.Unlock()
@@ -1166,7 +1181,7 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config
return fmt.Errorf("model is not a valid pipeline model: %s", trUpd.Model)
}
m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig)
m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig, fm)
if err != nil {
return err
}
@@ -45,7 +45,7 @@ var classifierTestHistory = schema.Messages{
func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent {
var out []types.ClassifierResultEvent
for _, e := range t.events {
for _, e := range t.events() {
if ev, ok := e.(types.ClassifierResultEvent); ok {
out = append(out, ev)
}
@@ -57,7 +57,7 @@ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent {
// item — what a classifier response actually "spoke".
func replyTexts(t *fakeTransport) []string {
var out []string
for _, e := range t.events {
for _, e := range t.events() {
if ev, ok := e.(types.ResponseOutputTextDoneEvent); ok {
out = append(out, ev.Text)
}
@@ -277,7 +277,7 @@ var _ = Describe("classifierRespond", func() {
Expect(t.countEvents(types.ServerEventTypeResponseOutputTextDone)).To(Equal(1))
Expect(t.countEvents(types.ServerEventTypeResponseFunctionCallArgumentsDone)).To(Equal(1))
var fcArgs string
for _, e := range t.events {
for _, e := range t.events() {
if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok {
fcArgs = done.Arguments
}
@@ -656,7 +656,7 @@ var _ = Describe("classifierRespond slot filling", func() {
Expect(results[0].Arguments).To(MatchJSON(`{"direction":"up","distance":3,"units":"meters"}`))
var fcArgs string
for _, e := range t.events {
for _, e := range t.events() {
if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok {
fcArgs = done.Arguments
}
@@ -695,7 +695,7 @@ var _ = Describe("classifierRespond slot filling", func() {
Expect(handled).To(BeTrue())
var fcArgs string
for _, e := range t.events {
for _, e := range t.events() {
if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok {
fcArgs = done.Arguments
}
@@ -17,8 +17,10 @@ import (
// so streaming behaviour can be asserted without a real WebSocket/WebRTC peer.
// It is not a *WebRTCTransport, so handler code takes the WebSocket path.
type fakeTransport struct {
events []types.ServerEvent
audio []fakeAudioChunk
// mu guards sent: some specs send from a background goroutine.
mu sync.Mutex
sent []types.ServerEvent
audio []fakeAudioChunk
}
type fakeAudioChunk struct {
@@ -27,10 +29,19 @@ type fakeAudioChunk struct {
}
func (f *fakeTransport) SendEvent(e types.ServerEvent) error {
f.events = append(f.events, e)
f.mu.Lock()
defer f.mu.Unlock()
f.sent = append(f.sent, e)
return nil
}
// events returns a copy of the server events sent so far.
func (f *fakeTransport) events() []types.ServerEvent {
f.mu.Lock()
defer f.mu.Unlock()
return append([]types.ServerEvent(nil), f.sent...)
}
func (f *fakeTransport) ReadEvent() ([]byte, error) { return nil, nil }
func (f *fakeTransport) SendAudio(_ context.Context, pcm []byte, sampleRate int) error {
@@ -43,7 +54,7 @@ func (f *fakeTransport) Close() error { return nil }
// countEvents returns how many recorded events have the given type.
func (f *fakeTransport) countEvents(et types.ServerEventType) int {
n := 0
for _, e := range f.events {
for _, e := range f.events() {
if e.ServerEventType() == et {
n++
}
@@ -55,7 +66,7 @@ func (f *fakeTransport) countEvents(et types.ServerEventType) int {
// delta event — i.e. the text streamed to the client as it is generated.
func (f *fakeTransport) transcriptDeltaText() string {
var b strings.Builder
for _, e := range f.events {
for _, e := range f.events() {
if d, ok := e.(types.ResponseOutputAudioTranscriptDeltaEvent); ok {
b.WriteString(d.Delta)
}
@@ -0,0 +1,180 @@
package openai
import (
"context"
"errors"
"fmt"
"sort"
"sync"
"github.com/mudler/LocalAI/core/backend"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/pkg/model"
)
// stageRouter routes realtime pipeline stages that name a failover chain.
// Every realtime model kind (full pipeline, transcription-only, sound-only)
// embeds one, so a chain resolves the same way whatever the session does.
type stageRouter struct {
// failover and stageChains route pipeline stages that name a failover
// chain; stageChains maps a stage (config.PipelineStage*) to its chain.
// The model's *Config fields then hold the target that was active at
// session start, for the checks that run once (voice, templates).
failover *failover.Manager
stageChains map[string]string
stageTargetConfig func(name string) (*config.ModelConfig, error)
appTracing bool
}
func newStageRouter(fm *failover.Manager, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) stageRouter {
return stageRouter{
failover: fm,
stageChains: map[string]string{},
stageTargetConfig: func(name string) (*config.ModelConfig, error) {
cfg, err := cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err != nil {
return nil, err
}
failover.PrepareTarget(cfg)
return cfg, nil
},
appTracing: appConfig.EnableTracing,
}
}
// resolveStage records a stage that names a chain and returns the chain's
// active target, so everything that inspects stage configs at session start
// sees a real model. Any other config is returned as is. A chain config
// reaching the model loader would have no backend and trigger backend
// auto-detection.
func (r *stageRouter) resolveStage(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) {
if cfg == nil || !cfg.IsFailover() {
return cfg, nil
}
if r.failover == nil {
return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name)
}
st, ok := r.failover.ChainStatus(cfg.Name)
if !ok {
return nil, fmt.Errorf("failover chain %q not found", cfg.Name)
}
r.stageChains[stage] = cfg.Name
return r.stageTargetConfig(st.Active)
}
// isChainStage reports whether stage names a failover chain.
func (r *stageRouter) isChainStage(stage string) bool {
_, ok := r.stageChains[stage]
return ok && r.failover != nil
}
// hasChainStages reports whether any stage names a chain.
func (r *stageRouter) hasChainStages() bool {
return r.failover != nil && len(r.stageChains) > 0
}
func (r *stageRouter) router() *stageRouter { return r }
// stageRouted is implemented by every realtime model that embeds a
// stageRouter; the session uses it to send failover events.
type stageRouted interface{ router() *stageRouter }
// stageCall runs fn with the config that should serve stage now. A plain
// stage uses base. A chain stage goes through the failover plan on every
// call, so a switch takes effect on the next call without rebuilding the
// session, and fn is retried on the next target until it calls commit.
func (r *stageRouter) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error {
if !r.isChainStage(stage) {
return fn(base, func() {})
}
chain := r.stageChains[stage]
return r.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error {
cfg, err := r.stageTargetConfig(target)
if err != nil {
return err
}
err = fn(cfg, commit)
if err != nil {
failover.RecordAttemptTrace(r.appTracing, chain, target, err)
}
return err
})
}
// warmStages preloads the stages. A chain stage warms through its failover
// plan: a target that fails to load moves the stage to the next one instead
// of failing the session.
func (r *stageRouter) warmStages(ctx context.Context, ml *model.ModelLoader, appConfig *config.ApplicationConfig, stages []backend.PreloadStage) error {
var (
plain []backend.PreloadStage
wg sync.WaitGroup
mu sync.Mutex
errs []error
)
for _, s := range stages {
if !r.isChainStage(s.Role) {
plain = append(plain, s)
continue
}
wg.Go(func() {
err := r.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error {
_, err := backend.PreloadStages(ctx, ml, appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}})
return err
})
mu.Lock()
errs = append(errs, err)
mu.Unlock()
})
}
_, err := backend.PreloadStages(ctx, ml, appConfig, plain)
wg.Wait()
return errors.Join(append(errs, err)...)
}
// startModelFailoverEvents starts failover events for m when it has chain
// stages. The returned func stops them and is never nil.
func startModelFailoverEvents(t Transport, m Model) func() {
sr, ok := m.(stageRouted)
if !ok {
return func() {}
}
r := sr.router()
if !r.hasChainStages() {
return func() {}
}
return startFailoverEvents(t, r.failover, r.stageChains)
}
// startFailoverEvents tells the client which target serves each chain stage
// now, and again whenever a chain switches. The returned func stops it.
func startFailoverEvents(t Transport, fm *failover.Manager, stageChains map[string]string) func() {
// Subscribe before reading the status, so a switch that lands in
// between is still delivered.
events, cancel := fm.Subscribe(16)
stages := make([]string, 0, len(stageChains))
for s := range stageChains {
stages = append(stages, s)
}
sort.Strings(stages)
for _, stage := range stages {
chain := stageChains[stage]
if st, ok := fm.ChainStatus(chain); ok {
sendEvent(t, types.ModelFailoverEvent{Chain: chain, Stage: stage, To: st.Active, State: string(st.State), Reason: string(failover.ReasonInitial)})
}
}
go func() {
for ev := range events {
if ev.Type != failover.EventChainSwitched {
continue
}
for _, stage := range stages {
if stageChains[stage] == ev.Chain {
sendEvent(t, types.ModelFailoverEvent{Chain: ev.Chain, Stage: stage, From: ev.From, To: ev.To, State: ev.State, Reason: string(ev.Reason)})
}
}
}
}()
return cancel
}
@@ -0,0 +1,206 @@
package openai
import (
"context"
"errors"
"os"
"path/filepath"
"time"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/system"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
type rtSource map[string]config.ModelConfig
func (s rtSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok }
func (s rtSource) GetAllModelsConfigs() []config.ModelConfig {
var out []config.ModelConfig
for _, c := range s {
out = append(out, c)
}
return out
}
var _ = Describe("realtime failover", func() {
var fm *failover.Manager
BeforeEach(func() {
fm = failover.New(rtSource{
"a": {Name: "a", Backend: "cloud-proxy"},
"b": {Name: "b", Backend: "llama-cpp"},
"chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}},
})
})
chainModel := func() *wrappedModel {
return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{config.PipelineStageTTS: "chain"},
stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }}}
}
It("routes a chain stage through the plan and retries before commit", func() {
m := chainModel()
var tried []string
err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, _ func()) error {
tried = append(tried, cfg.Name)
if cfg.Name == "a" {
return errors.New("dial tcp: refused")
}
return nil
})
Expect(err).ToNot(HaveOccurred())
Expect(tried).To(Equal([]string{"a", "b"}))
})
It("does not retry a chain stage once output was committed", func() {
m := chainModel()
var tried []string
err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, commit func()) error {
tried = append(tried, cfg.Name)
commit()
return errors.New("dial tcp: refused")
})
Expect(err).To(HaveOccurred())
Expect(tried).To(Equal([]string{"a"}))
})
It("calls a plain stage once with its own config", func() {
m := &wrappedModel{}
base := &config.ModelConfig{Name: "plain"}
calls := 0
err := m.stageCall(context.Background(), config.PipelineStageTTS, base, func(cfg *config.ModelConfig, _ func()) error {
calls++
Expect(cfg).To(BeIdenticalTo(base))
return nil
})
Expect(err).ToNot(HaveOccurred())
Expect(calls).To(Equal(1))
})
It("sends initial events, then switch events, and stops on cancel", func() {
t := &fakeTransport{}
failoverEvents := func() []types.ModelFailoverEvent {
var out []types.ModelFailoverEvent
for _, e := range t.events() {
if fe, ok := e.(types.ModelFailoverEvent); ok {
out = append(out, fe)
}
}
return out
}
stop := startFailoverEvents(t, fm, map[string]string{config.PipelineStageLLM: "chain"})
Eventually(failoverEvents).Should(ContainElement(And(
HaveField("Reason", "initial"), HaveField("To", "a"), HaveField("Stage", "llm"))))
fm.ReportFailure("a", errors.New("dial tcp: refused"))
Eventually(failoverEvents).Should(ContainElement(And(
HaveField("Reason", "trip"), HaveField("From", "a"), HaveField("To", "b"), HaveField("Chain", "chain"))))
stop()
})
})
var _ = Describe("realtime failover in transcription-only and sound-only sessions", func() {
var (
cl *config.ModelConfigLoader
ml *model.ModelLoader
appConfig *config.ApplicationConfig
fm *failover.Manager
)
BeforeEach(func() {
dir := GinkgoT().TempDir()
write := func(name, body string) {
Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed())
}
write("vad", "name: vad\nbackend: silero-vad\n")
write("stt-a", "name: stt-a\nbackend: fake-stt-a\n")
write("stt-b", "name: stt-b\nbackend: fake-stt-b\n")
write("stt-chain", "name: stt-chain\nfailover:\n targets:\n - model: stt-a\n - model: stt-b\n")
write("sound-a", "name: sound-a\nbackend: fake-sound-a\n")
write("sound-b", "name: sound-b\nbackend: fake-sound-b\n")
write("sound-chain", "name: sound-chain\nfailover:\n targets:\n - model: sound-a\n - model: sound-b\n")
ss := &system.SystemState{Model: system.Model{ModelsPath: dir}}
appConfig = config.NewApplicationConfig()
appConfig.SystemState = ss
cl = config.NewModelConfigLoader(dir)
Expect(cl.LoadModelConfigsFromPath(dir)).To(Succeed())
ml = model.NewModelLoader(ss)
fm = failover.New(cl)
})
// failoverEvents collects the failover events a transport received.
failoverEvents := func(t *fakeTransport) func() []types.ModelFailoverEvent {
return func() []types.ModelFailoverEvent {
var out []types.ModelFailoverEvent
for _, e := range t.events() {
if fe, ok := e.(types.ModelFailoverEvent); ok {
out = append(out, fe)
}
}
return out
}
}
It("resolves a sound_detection chain and routes each call through it", func() {
m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, fm)
Expect(err).ToNot(HaveOccurred())
tm := m.(*transcriptOnlyModel)
Expect(tm.SoundDetectionConfig.Name).To(Equal("sound-a"))
Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageSoundDetection: "sound-chain"}))
var tried []string
tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) {
tried = append(tried, name)
return nil, errors.New("dial tcp: refused")
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, err = tm.SoundDetection(ctx, "a.wav", 3, 0.1)
Expect(err).To(HaveOccurred())
Expect(tried).To(Equal([]string{"sound-a", "sound-b"}))
t := &fakeTransport{}
stop := startModelFailoverEvents(t, m)
defer stop()
Eventually(failoverEvents(t)).Should(ContainElement(And(
HaveField("Stage", "sound_detection"), HaveField("Chain", "sound-chain"), HaveField("Reason", "initial"))))
})
It("resolves a transcription chain in a transcription-only session", func() {
m, cfg, err := newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, fm)
Expect(err).ToNot(HaveOccurred())
Expect(cfg.Name).To(Equal("stt-a"))
tm := m.(*transcriptOnlyModel)
Expect(tm.VADConfig.Name).To(Equal("vad"))
Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageTranscription: "stt-chain"}))
var tried []string
tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) {
tried = append(tried, name)
return nil, errors.New("dial tcp: refused")
}
_, err = tm.Transcribe(context.Background(), "a.wav", "", false, false, "")
Expect(err).To(HaveOccurred())
Expect(tried).To(Equal([]string{"stt-a", "stt-b"}))
})
It("fails before touching a backend when failover is not running", func() {
_, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, nil)
Expect(err).To(MatchError(ContainSubstring("failover is not running")))
_, _, err = newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, nil)
Expect(err).To(MatchError(ContainSubstring("failover is not running")))
Expect(ml.ListLoadedModels()).To(BeEmpty())
})
It("sends no failover events for a session without chains", func() {
m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-a"}, cl, ml, appConfig, fm)
Expect(err).ToNot(HaveOccurred())
t := &fakeTransport{}
startModelFailoverEvents(t, m)()
Consistently(failoverEvents(t), 200*time.Millisecond).Should(BeEmpty())
})
})
+226 -29
View File
@@ -19,6 +19,7 @@ import (
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
"github.com/mudler/LocalAI/core/http/middleware"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/routing/router"
"github.com/mudler/LocalAI/core/services/voiceprofile"
"github.com/mudler/LocalAI/core/templates"
@@ -84,6 +85,11 @@ type wrappedModel struct {
routerStore router.DecisionStore
routerSessionID string
routerUserID string
stageRouter
// tuneLLM applies the pipeline's LLM overrides (reasoning effort,
// disable_thinking) to a chain target loaded per call.
tuneLLM func(cfg *config.ModelConfig)
}
// anyToAnyModel represent a model which supports Any-to-Any operations
@@ -106,18 +112,38 @@ type transcriptOnlyModel struct {
appConfig *config.ApplicationConfig
modelLoader *model.ModelLoader
confLoader *config.ModelConfigLoader
stageRouter
}
func (m *transcriptOnlyModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) {
return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig)
var res *schema.VADResponse
err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg)
return err
})
return res, err
}
func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) {
return backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig)
var res *schema.TranscriptionResult
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *cfg, m.appConfig)
return err
})
return res, err
}
func (m *transcriptOnlyModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) {
return modelSoundDetection(ctx, m.modelLoader, m.appConfig, m.SoundDetectionConfig, audio, topK, threshold)
var res *schema.SoundClassificationResult
err := m.stageCall(ctx, config.PipelineStageSoundDetection, m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = modelSoundDetection(ctx, m.modelLoader, m.appConfig, cfg, audio, topK, threshold)
return err
})
return res, err
}
func (m *transcriptOnlyModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
@@ -144,11 +170,27 @@ func (m *transcriptOnlyModel) TTSStream(ctx context.Context, text, voice, langua
}
func (m *transcriptOnlyModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) {
return transcribeStream(ctx, m.modelLoader, *m.TranscriptionConfig, m.appConfig, audio, language, translate, diarize, prompt, onDelta)
var res *schema.TranscriptionResult
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error {
var err error
res, err = transcribeStream(ctx, m.modelLoader, *cfg, m.appConfig, audio, language, translate, diarize, prompt, func(s string) {
commit()
onDelta(s)
})
return err
})
return res, err
}
func (m *transcriptOnlyModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) {
return backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *m.TranscriptionConfig, m.appConfig, onEvent)
var live backend.LiveTranscriptionSession
// Only opening the live session can move to the next target.
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent)
return err
})
return live, err
}
func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig {
@@ -156,24 +198,41 @@ func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig {
}
func (m *transcriptOnlyModel) Warmup(ctx context.Context) error {
_, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{
{Role: "vad", Cfg: m.VADConfig},
{Role: "transcription", Cfg: m.TranscriptionConfig},
{Role: "sound_detection", Cfg: m.SoundDetectionConfig},
return m.warmStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{
{Role: config.PipelineStageVAD, Cfg: m.VADConfig},
{Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig},
{Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig},
})
return err
}
func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) {
return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig)
var res *schema.VADResponse
err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg)
return err
})
return res, err
}
func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) {
return backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig)
var res *schema.TranscriptionResult
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *cfg, m.appConfig)
return err
})
return res, err
}
func (m *wrappedModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) {
return modelSoundDetection(ctx, m.modelLoader, m.appConfig, m.SoundDetectionConfig, audio, topK, threshold)
var res *schema.SoundClassificationResult
err := m.stageCall(ctx, config.PipelineStageSoundDetection, m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
res, err = modelSoundDetection(ctx, m.modelLoader, m.appConfig, cfg, audio, topK, threshold)
return err
})
return res, err
}
func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
@@ -181,11 +240,22 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im
Messages: messages,
}
toolsJSON, toolChoiceJSON := realtimeToolsJSON(tools, toolChoice)
// infer renders the prompt for cfg and starts inference on it. Everything
// that reads the LLM config lives here, so a chain stage can run it again
// against the next target.
infer := func(cfg *config.ModelConfig, cb func(string, backend.TokenUsage) bool) (func() (backend.LLMResponse, error), error) {
predInput := m.renderPredictPrompt(input, cfg, tools, toolChoice)
return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, cfg, m.confLoader, m.appConfig, cb, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil)
}
// Per-turn routing: when the session's LLMConfig is a router, swap
// to the candidate the classifier picks for this turn's prompt.
// LLMConfig itself is held by value (we never mutate it) — turnCfg
// is the config we dispatch against.
turnCfg := m.LLMConfig
routed := false
if m.LLMConfig.HasRouter() && m.routerDeps != nil {
chosen, err := m.routeTurn(ctx, &input)
if err != nil {
@@ -193,9 +263,47 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im
"router_model", m.LLMConfig.Name, "error", err)
} else if chosen != nil {
turnCfg = chosen
routed = true
}
}
// A routed turn dispatches to the router's pick: chains as router
// candidates are not resolved here.
if routed || !m.isChainStage(config.PipelineStageLLM) {
return infer(turnCfg, tokenCallback)
}
return func() (backend.LLMResponse, error) {
var resp backend.LLMResponse
err := m.stageCall(ctx, config.PipelineStageLLM, turnCfg, func(cfg *config.ModelConfig, commit func()) error {
if m.tuneLLM != nil {
m.tuneLLM(cfg)
}
// Without a callback nothing reaches the client before the
// reply is complete, so every failure can still be retried.
var cb func(string, backend.TokenUsage) bool
if tokenCallback != nil {
cb = func(s string, u backend.TokenUsage) bool {
commit()
return tokenCallback(s, u)
}
}
predict, err := infer(cfg, cb)
if err != nil {
return err
}
resp, err = predict()
return err
})
return resp, err
}, nil
}
// renderPredictPrompt templates the turn's prompt for cfg. It also applies
// the turn's tool choice and function-calling grammar to cfg, which the
// backend reads when inference starts. The prompt is empty for models that
// use the tokenizer's template.
func (m *wrappedModel) renderPredictPrompt(input schema.OpenAIRequest, turnCfg *config.ModelConfig, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) string {
// Surface the resolved reasoning effort to the Go-side template path too
// (jinja models get it via backend metadata in gRPCPredictOpts; Go-templated
// models like gpt-oss read it from the template's .ReasoningEffort).
@@ -303,6 +411,12 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im
}
}
return predInput
}
// realtimeToolsJSON serializes the turn's tools and tool choice the way the
// backends expect them. Neither depends on the LLM config.
func realtimeToolsJSON(tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) (string, string) {
var toolsJSON string
if len(tools) > 0 {
// Convert tools to OpenAI Chat Completions format (nested)
@@ -348,7 +462,7 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im
toolChoiceJSON = string(b)
}
return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, turnCfg, m.confLoader, m.appConfig, tokenCallback, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil)
return toolsJSON, toolChoiceJSON
}
// routeTurn classifies this turn's prompt against the session's router
@@ -395,7 +509,16 @@ func newRealtimeDecisionID() string {
}
func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) {
return backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *m.TTSConfig)
var (
out string
res *proto.Result
)
err := m.stageCall(ctx, config.PipelineStageTTS, m.TTSConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
out, res, err = backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *cfg)
return err
})
return out, res, err
}
func (m *wrappedModel) setTTSParams(params map[string]string) {
@@ -403,7 +526,14 @@ func (m *wrappedModel) setTTSParams(params map[string]string) {
}
func (m *wrappedModel) TTSStream(ctx context.Context, text, voice, language string, onAudio func(pcm []byte, sampleRate int) error) error {
return ttsStream(ctx, m.modelLoader, m.appConfig, *m.TTSConfig, text, voice, language, maps.Clone(m.ttsParams), onAudio)
return m.stageCall(ctx, config.PipelineStageTTS, m.TTSConfig, func(cfg *config.ModelConfig, commit func()) error {
// Audio that reached the client cannot be taken back, so the first
// chunk ends the retries.
return ttsStream(ctx, m.modelLoader, m.appConfig, *cfg, text, voice, language, maps.Clone(m.ttsParams), func(pcm []byte, sr int) error {
commit()
return onAudio(pcm, sr)
})
})
}
func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig *config.ModelConfig, profiles *voiceprofile.Store) (string, map[string]string, func(), error) {
@@ -431,11 +561,28 @@ func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig
}
func (m *wrappedModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) {
return transcribeStream(ctx, m.modelLoader, *m.TranscriptionConfig, m.appConfig, audio, language, translate, diarize, prompt, onDelta)
var res *schema.TranscriptionResult
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error {
var err error
res, err = transcribeStream(ctx, m.modelLoader, *cfg, m.appConfig, audio, language, translate, diarize, prompt, func(s string) {
commit()
onDelta(s)
})
return err
})
return res, err
}
func (m *wrappedModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) {
return backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *m.TranscriptionConfig, m.appConfig, onEvent)
var live backend.LiveTranscriptionSession
// Only opening the live session can move to the next target: once it is
// open, events flow to the client for the rest of the utterance.
err := m.stageCall(ctx, config.PipelineStageTranscription, m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error {
var err error
live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent)
return err
})
return live, err
}
func (m *wrappedModel) PredictConfig() *config.ModelConfig {
@@ -683,18 +830,17 @@ func (m *wrappedModel) FillToolArguments(ctx context.Context, messages schema.Me
func (m *wrappedModel) Warmup(ctx context.Context) error {
stages := []backend.PreloadStage{
{Role: "vad", Cfg: m.VADConfig},
{Role: "transcription", Cfg: m.TranscriptionConfig},
{Role: "llm", Cfg: m.LLMConfig},
{Role: "tts", Cfg: m.TTSConfig},
{Role: "sound_detection", Cfg: m.SoundDetectionConfig},
{Role: config.PipelineStageVAD, Cfg: m.VADConfig},
{Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig},
{Role: config.PipelineStageLLM, Cfg: m.LLMConfig},
{Role: config.PipelineStageTTS, Cfg: m.TTSConfig},
{Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig},
}
// The scoring model is a separate stage only when it isn't the LLM.
if m.ScoreConfig != nil && m.ScoreConfig != m.LLMConfig {
stages = append(stages, backend.PreloadStage{Role: "classifier", Cfg: m.ScoreConfig})
}
_, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, stages)
return err
return m.warmStages(ctx, m.modelLoader, m.appConfig, stages)
}
// wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream
@@ -787,8 +933,12 @@ func loadSoundDetectionConfig(pipeline *config.Pipeline, cl *config.ModelConfigL
return cfg, nil
}
func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, *config.ModelConfig, error) {
func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, *config.ModelConfig, error) {
sr := newStageRouter(fm, cl, ml, appConfig)
cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgVAD, err = sr.resolveStage(config.PipelineStageVAD, cfgVAD)
}
if err != nil {
return nil, nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -799,6 +949,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig
}
cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgSST, err = sr.resolveStage(config.PipelineStageTranscription, cfgSST)
}
if err != nil {
return nil, nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -809,6 +962,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig
}
cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig)
if err == nil {
cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound)
}
if err != nil {
return nil, nil, err
}
@@ -821,6 +977,7 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig
confLoader: cl,
modelLoader: ml,
appConfig: appConfig,
stageRouter: sr,
}, cfgSST, nil
}
@@ -829,8 +986,12 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig
// a sound-detection-only realtime session, which activates on sounds (not
// speech) and is driven by client-side windowing (turn_detection none +
// input_audio_buffer.commit) rather than the voice VAD loop.
func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, error) {
func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, error) {
sr := newStageRouter(fm, cl, ml, appConfig)
cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig)
if err == nil {
cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound)
}
if err != nil {
return nil, err
}
@@ -842,6 +1003,7 @@ func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfi
confLoader: cl,
modelLoader: ml,
appConfig: appConfig,
stageRouter: sr,
}, nil
}
@@ -854,6 +1016,8 @@ type RealtimeRoutingContext struct {
Store router.DecisionStore
SessionID string
UserID string
// Failover resolves pipeline stages that name a failover chain.
Failover *failover.Manager
}
// buildRealtimeRoutingContext assembles the routing dependencies the
@@ -875,6 +1039,7 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) *
Store: a.RouterDecisions(),
SessionID: sessionID,
UserID: userID,
Failover: a.FailoverManager(),
}
}
@@ -882,7 +1047,20 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) *
func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, evaluator *templates.Evaluator, routing *RealtimeRoutingContext) (Model, error) {
xlog.Debug("Creating new model pipeline model", "pipeline", pipeline)
// A stage that names a failover chain is resolved on every call. Here it
// takes the chain's active target, so everything that inspects stage
// configs at session start (voice, reasoning, templates) sees a real model.
var fm *failover.Manager
if routing != nil {
fm = routing.Failover
}
sr := newStageRouter(fm, cl, ml, appConfig)
resolveStage := sr.resolveStage
cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgVAD, err = resolveStage(config.PipelineStageVAD, cfgVAD)
}
if err != nil {
return nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -894,6 +1072,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model
// TODO: Do we always need a transcription model? It can be disabled. Note that any-to-any instruction following models don't transcribe as such, so if transcription is required it is a separate process
cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgSST, err = resolveStage(config.PipelineStageTranscription, cfgSST)
}
if err != nil {
return nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -926,6 +1107,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model
// Otherwise we want to return a wrapped model, which is a "virtual" model that re-uses other models to perform operations
cfgLLM, err := cl.LoadResolvedModelConfig(pipeline.LLM, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgLLM, err = resolveStage(config.PipelineStageLLM, cfgLLM)
}
if err != nil {
return nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -937,10 +1121,17 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model
// Let the pipeline set the LLM's reasoning effort and force thinking off
// (cfgLLM is a per-session copy). disable_thinking applies after the effort.
applyPipelineReasoning(cfgLLM, *pipeline)
applyPipelineThinking(cfgLLM, *pipeline)
pipelineCopy := *pipeline
tuneLLM := func(cfg *config.ModelConfig) {
applyPipelineReasoning(cfg, pipelineCopy)
applyPipelineThinking(cfg, pipelineCopy)
}
tuneLLM(cfgLLM)
cfgTTS, err := cl.LoadResolvedModelConfig(pipeline.TTS, ml.ModelPath, appConfig.ToConfigLoaderOptions()...)
if err == nil {
cfgTTS, err = resolveStage(config.PipelineStageTTS, cfgTTS)
}
if err != nil {
return nil, fmt.Errorf("failed to load backend config: %w", err)
@@ -951,6 +1142,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model
}
cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig)
if err == nil {
cfgSound, err = resolveStage(config.PipelineStageSoundDetection, cfgSound)
}
if err != nil {
return nil, err
}
@@ -1000,6 +1194,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model
modelLoader: ml,
appConfig: appConfig,
evaluator: evaluator,
stageRouter: sr,
tuneLLM: tuneLLM,
}
if routing != nil {
wm.routerDeps = routing.Deps
@@ -269,7 +269,7 @@ var _ = Describe("liveTurnState", func() {
lts.drainEvents(1.0)
var got []types.ConversationItemInputAudioTranscriptionDeltaEvent
for _, e := range ftr.events {
for _, e := range ftr.events() {
if d, ok := e.(types.ConversationItemInputAudioTranscriptionDeltaEvent); ok {
got = append(got, d)
}
@@ -335,7 +335,7 @@ var _ = Describe("commitUtteranceWithTranscript", func() {
Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1))
var completed types.ConversationItemInputAudioTranscriptionCompletedEvent
for _, e := range tr.events {
for _, e := range tr.events() {
if c, ok := e.(types.ConversationItemInputAudioTranscriptionCompletedEvent); ok {
completed = c
}
@@ -407,7 +407,7 @@ var _ = Describe("emitPrecomputedTranscription", func() {
Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionDelta)).To(Equal(2), "empty deltas skipped")
Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1))
for _, e := range tr.events {
for _, e := range tr.events() {
switch ev := e.(type) {
case types.ConversationItemInputAudioTranscriptionDeltaEvent:
Expect(ev.ItemID).To(Equal("item42"))
@@ -38,7 +38,7 @@ var _ = Describe("emitSoundDetection", func() {
Expect(err).ToNot(HaveOccurred())
Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1))
ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent)
ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent)
Expect(ok).To(BeTrue())
Expect(ev.ItemID).To(Equal("item1"))
Expect(ev.ContentIndex).To(Equal(0))
@@ -62,7 +62,7 @@ var _ = Describe("emitSoundDetection", func() {
Expect(err).ToNot(HaveOccurred())
Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1))
ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent)
ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent)
Expect(ok).To(BeTrue())
Expect(ev.Detections).To(BeEmpty())
})
@@ -250,8 +250,8 @@ var _ = Describe("triggerResponse", func() {
// The single terminal carries the produced output item and the usage —
// both empty in the legacy code.
var done *types.ResponseDoneEvent
for i := range t.events {
if d, ok := t.events[i].(types.ResponseDoneEvent); ok {
for i := range t.events() {
if d, ok := t.events()[i].(types.ResponseDoneEvent); ok {
done = &d
}
}
@@ -287,8 +287,8 @@ var _ = Describe("triggerResponse", func() {
var created *types.ResponseCreatedEvent
var done *types.ResponseDoneEvent
for i := range t.events {
switch e := t.events[i].(type) {
for i := range t.events() {
switch e := t.events()[i].(type) {
case types.ResponseCreatedEvent:
created = &e
case types.ResponseDoneEvent:
@@ -317,8 +317,8 @@ var _ = Describe("triggerResponse", func() {
triggerResponse(context.Background(), session, &Conversation{}, t, nil)
for i := range t.events {
if d, ok := t.events[i].(types.ResponseDoneEvent); ok {
for i := range t.events() {
if d, ok := t.events()[i].(types.ResponseDoneEvent); ok {
Expect(d.Response.Metadata).To(BeEmpty())
}
}
@@ -67,7 +67,7 @@ func itSession(gate *voiceGate) (*Session, *fakeModel) {
// hasSpeakerNotAuthorized reports whether a speaker_not_authorized error event
// was emitted to the client.
func hasSpeakerNotAuthorized(tr *fakeTransport) bool {
for _, e := range tr.events {
for _, e := range tr.events() {
if ev, ok := e.(types.ErrorEvent); ok && ev.Error.Code == "speaker_not_authorized" {
return true
}
@@ -48,7 +48,7 @@ func SoundClassificationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLo
Threshold: float32(parseFormFloat(c, "threshold", 0)),
}
file, err := c.FormFile("file")
file, err := uploadedFile(c, "file")
if err != nil {
return err
}
+32 -3
View File
@@ -2,9 +2,9 @@ package openai
import (
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"os"
"path"
@@ -56,6 +56,18 @@ func resolveTranscriptionTranslate(formTranslate string, configTranslate bool) b
return configTranslate
}
// uploadedFile reads a required multipart file field. Any failure here (no
// multipart boundary, malformed body, missing field) is caused by the request,
// so it maps to 400 instead of leaking the parser error as a 500.
func uploadedFile(c echo.Context, field string) (*multipart.FileHeader, error) {
file, err := c.FormFile(field)
if err != nil {
return nil, echo.NewHTTPError(http.StatusBadRequest,
fmt.Sprintf("missing or invalid %q file upload: %v", field, err))
}
return file, nil
}
func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) echo.HandlerFunc {
return func(c echo.Context) error {
input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest)
@@ -104,8 +116,15 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app
}
}
// Reject an unknown format before the backend runs: the backend work
// would be wasted, and a failover chain would count the error
// against every target.
if !stream && !validTranscriptionResponseFormat(responseFormat) {
return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format")
}
// retrieve the file data from the request
file, err := c.FormFile("file")
file, err := uploadedFile(c, "file")
if err != nil {
return err
}
@@ -217,11 +236,21 @@ func TranscriptEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, app
}
return c.JSON(http.StatusOK, trs)
default:
return errors.New("invalid response_format")
return echo.NewHTTPError(http.StatusBadRequest, "invalid response_format")
}
}
}
func validTranscriptionResponseFormat(f schema.TranscriptionResponseFormatType) bool {
switch f {
case "", schema.TranscriptionResponseFormatLrc, schema.TranscriptionResponseFormatText,
schema.TranscriptionResponseFormatSrt, schema.TranscriptionResponseFormatVtt,
schema.TranscriptionResponseFormatJson, schema.TranscriptionResponseFormatJsonVerbose:
return true
}
return false
}
// streamTranscription emits OpenAI-format SSE events for a transcription
// request: one `transcript.text.delta` per backend chunk, a final
// `transcript.text.done` with the assembled text, and `[DONE]`. Backends that
@@ -0,0 +1,46 @@
package types
import "encoding/json"
// ModelFailoverEvent is a LocalAI extension server event
// (localai.model.failover). It tells a client which target serves a
// pipeline stage that names a failover chain: once per chain stage at
// session start (reason "initial"), then on every switch of that chain.
type ModelFailoverEvent struct {
ServerEventBase
// The failover chain the stage names.
Chain string `json:"chain"`
// The pipeline stage: vad, transcription, llm, tts or sound_detection.
Stage string `json:"stage"`
// The target that served the stage before the switch; "" at session start.
From string `json:"from"`
// The target that serves the stage now.
To string `json:"to"`
// The chain state: primary, fallback or degraded.
State string `json:"state"`
// Why the chain switched, or "initial" at session start.
Reason string `json:"reason"`
}
func (m ModelFailoverEvent) ServerEventType() ServerEventType {
return ServerEventTypeModelFailover
}
func (m ModelFailoverEvent) MarshalJSON() ([]byte, error) {
type typeAlias ModelFailoverEvent
type typeWrapper struct {
typeAlias
Type ServerEventType `json:"type"`
}
shadow := typeWrapper{
typeAlias: typeAlias(m),
Type: m.ServerEventType(),
}
return json.Marshal(shadow)
}
@@ -27,7 +27,11 @@ const (
// ServerEventTypeClassifierResult is a LocalAI extension: it carries the
// classifier-mode score distribution and decision for a response. OpenAI
// clients ignore it.
ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result"
ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result"
// ServerEventTypeModelFailover is a LocalAI extension: it names the target
// that serves a pipeline stage backed by a failover chain, at session start
// and on every chain switch. OpenAI clients ignore it.
ServerEventTypeModelFailover ServerEventType = "localai.model.failover"
ServerEventTypeInputAudioBufferCommitted ServerEventType = "input_audio_buffer.committed"
ServerEventTypeInputAudioBufferCleared ServerEventType = "input_audio_buffer.cleared"
ServerEventTypeInputAudioBufferSpeechStarted ServerEventType = "input_audio_buffer.speech_started"
+1
View File
@@ -38,6 +38,7 @@ func AdmissionControl(limiter *admission.Limiter, events pii.EventStore) echo.Mi
if !ok {
retryAfter := admission.RetryAfter(cfg.Limits.RetryAfterSeconds)
recordAdmissionRejection(events, cfg.Name, retryAfter)
c.Set(ContextKeyAdmissionRejected, true)
c.Response().Header().Set("Retry-After", strconv.Itoa(int(retryAfter.Seconds())))
return c.JSON(http.StatusTooManyRequests, map[string]any{
"error": map[string]any{
+9
View File
@@ -47,4 +47,13 @@ const (
// router nor the body-parse path has produced one. Distinct from
// ContextKeyServedModel, which is the router's resolved choice.
ContextKeyResponseModel = "routing.response_model"
// ContextKeyFailoverAttempt holds the *failoverState of a request whose
// model is a failover chain.
ContextKeyFailoverAttempt = "failover.attempt"
// ContextKeyAdmissionRejected is set to true by AdmissionControl when it
// turns a request away because the model is at capacity. The failover
// retry sends such a request to the next target without tripping this one.
ContextKeyAdmissionRejected = "admission.rejected"
)
+335
View File
@@ -0,0 +1,335 @@
package middleware
import (
"bufio"
"bytes"
"fmt"
"io"
"net"
"net/http"
"slices"
"strings"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/services/failover"
)
const (
HeaderServedModel = "X-LocalAI-Served-Model"
HeaderFailover = "X-LocalAI-Failover"
)
// MaxFailoverReplayBody caps the request body kept for a retry. A larger body
// is still served, by one target only.
const MaxFailoverReplayBody = 32 << 20
type failoverState struct {
attempt *failover.Attempt
}
// SetFailoverManager enables failover chains. Without it, a request for a
// chain fails with 503.
func (re *RequestExtractor) SetFailoverManager(m *failover.Manager) { re.failover = m }
// resolveFailover returns the config of the target that should serve this
// attempt. The first attempt plans the chain; retries reuse the plan.
func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, chain *config.ModelConfig) (*config.ModelConfig, error) {
st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState)
if st == nil || st.attempt.Chain() != chain.Name {
if re.failover == nil {
return nil, fmt.Errorf("model %q is a failover chain, but failover is not running", chain.Name)
}
att, err := re.failover.Plan(chain.Name)
if err != nil {
return nil, err
}
st = &failoverState{attempt: att}
c.Set(ContextKeyFailoverAttempt, st)
}
for {
cfg, err := re.loadFailoverTarget(st.attempt.Target())
// A target that became a chain after its chain was saved has no
// backend of its own; like a disabled one it is skipped, not tripped.
if err == nil && (cfg.IsDisabled() || cfg.IsFailover()) {
// Disabled on purpose, not broken: move on without a trip.
if st.attempt.Skip() {
continue
}
c.Set(ContextKeyFailoverAttempt, nil)
return nil, fmt.Errorf("failover chain %q: target %q is disabled or is itself a chain", chain.Name, cfg.Name)
}
if err == nil {
failover.PrepareTarget(cfg) // cfg is a copy
c.Set(ContextKeyRequestedModel, requested)
c.Set(ContextKeyServedModel, cfg.Name)
setFailoverHeaders(c.Response().Header(), st.attempt)
return cfg, nil
}
if !st.attempt.Fail(err) {
// Clear the state so the retry wrapper sends this 503 as is.
c.Set(ContextKeyFailoverAttempt, nil)
return nil, err
}
}
}
func (re *RequestExtractor) loadFailoverTarget(name string) (*config.ModelConfig, error) {
cfg, err := re.modelConfigLoader.LoadModelConfigFileByNameDefaultOptions(name, re.applicationConfig)
if err != nil {
return nil, err
}
resolved, _, err := re.modelConfigLoader.ResolveAlias(cfg)
return resolved, err
}
func setFailoverHeaders(h http.Header, att *failover.Attempt) {
h.Set(HeaderServedModel, att.Target())
switch {
case att.Degraded():
h.Set(HeaderFailover, "degraded")
case att.Target() != att.Primary():
h.Set(HeaderFailover, "fallback")
default:
h.Del(HeaderFailover)
}
}
// failoverRetry runs h again on the next target while the response is not
// committed. h is SetModelAndConfig's body plus the rest of the chain, so
// every attempt binds the request again from the replayed body.
func (re *RequestExtractor) failoverRetry(h echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
// Without chains there is nothing to retry, so plain installations
// pay neither for the body copy nor for the writer.
if !re.failover.HasChains() {
return h(c)
}
req := c.Request()
src := req.Body
if src == nil {
src = http.NoBody
}
rec := &replayBody{src: src, limit: MaxFailoverReplayBody}
req.Body = rec
// A default-model middleware may have parsed a multipart form before
// this point, draining the body. That form stays valid for every
// attempt; a form parsed during an attempt is dropped and parsed again
// from the replayed body.
entryMultipart, entryForm, entryPostForm := req.MultipartForm, req.Form, req.PostForm
resp := c.Response()
orig := resp.Writer
baseHeader := resp.Header().Clone()
defer func() { resp.Writer = orig }()
active := func() bool {
st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState)
return st != nil
}
tracing := re.applicationConfig != nil && re.applicationConfig.EnableTracing
for {
w := &failoverWriter{ResponseWriter: orig, active: active}
resp.Writer = w
c.Set(ContextKeyAdmissionRejected, nil)
err := h(c)
st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState)
if st == nil {
w.release()
return err
}
att := st.attempt
status := w.held
if err == nil && status == 0 {
// A 4xx says nothing about the target's health.
if w.status < http.StatusBadRequest {
att.Succeed()
}
return nil
}
rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool)
// A handler that wrote its 501 itself instead of returning it
// reports the same gap.
gap := (failover.IsCapabilityGap(err) || status == http.StatusNotImplemented) && !w.committed
if (rejected && !w.committed) || gap {
// The target is at capacity, or cannot serve this kind of
// request at all: spill to the next target without counting
// a failure. Neither says anything about the target's health.
if !rec.replayable() || !att.Skip() {
w.release()
return err
}
} else {
cause := attemptError(err, status, w.body.Bytes())
retryable := req.Context().Err() == nil && failover.IsRetryable(err, status)
if !retryable || w.committed || !rec.replayable() {
if retryable {
att.Report(cause)
}
w.release()
return err
}
failover.RecordAttemptTrace(tracing, att.Chain(), att.Target(), cause)
if !att.Fail(cause) {
w.release()
return err
}
}
// Temp files of a form parsed during this attempt would otherwise
// outlive the request: the server only cleans up the last form.
if mf := c.Request().MultipartForm; mf != nil && mf != entryMultipart {
_ = mf.RemoveAll()
}
req.Body = rec.replay()
req.MultipartForm, req.Form, req.PostForm = entryMultipart, entryForm, entryPostForm
// Middleware after this one may have replaced the request; the
// next attempt starts again from the request as it arrived here.
c.SetRequest(req)
resetResponse(resp, baseHeader)
}
}
}
// stopFailoverRecording releases the recorded body of a request whose model
// turned out not to be a chain: it will never be replayed.
func stopFailoverRecording(c echo.Context) {
if rb, ok := c.Request().Body.(*replayBody); ok {
rb.stop()
}
}
func attemptError(err error, status int, body []byte) error {
if err != nil {
return err
}
msg := strings.TrimSpace(string(body))
if len(msg) > 200 {
msg = msg[:200]
}
return fmt.Errorf("HTTP %d: %s", status, msg)
}
func resetResponse(resp *echo.Response, base http.Header) {
h := resp.Header()
for k := range h {
delete(h, k)
}
for k, v := range base {
h[k] = slices.Clone(v)
}
resp.Committed = false
resp.Status = http.StatusOK
resp.Size = 0
}
// failoverWriter holds back an error response (status >= 500) of a chain
// request until the handler returns, so the retry can drop it.
type failoverWriter struct {
http.ResponseWriter
active func() bool
held int
body bytes.Buffer
committed bool
// status is the code sent to the client, 0 until one is sent.
status int
}
func (w *failoverWriter) WriteHeader(code int) {
if w.held != 0 {
return
}
if !w.committed && code >= 500 && w.active() {
w.held = code
return
}
w.committed = true
w.status = code
w.ResponseWriter.WriteHeader(code)
}
func (w *failoverWriter) Write(b []byte) (int, error) {
if w.held != 0 {
return w.body.Write(b)
}
w.committed = true
return w.ResponseWriter.Write(b)
}
// FlushError and Hijack go through a ResponseController so the capabilities of
// wrapped writers further down stay reachable, as they were before this writer
// was inserted.
func (w *failoverWriter) FlushError() error {
w.release()
w.committed = true
return http.NewResponseController(w.ResponseWriter).Flush()
}
func (w *failoverWriter) Flush() { _ = w.FlushError() }
func (w *failoverWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
w.committed = true
return http.NewResponseController(w.ResponseWriter).Hijack()
}
func (w *failoverWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter }
// release sends a held error response to the client.
func (w *failoverWriter) release() {
if w.held == 0 {
return
}
code := w.held
w.held = 0
w.committed = true
w.ResponseWriter.WriteHeader(code)
_, _ = w.ResponseWriter.Write(w.body.Bytes())
w.body.Reset()
}
// replayBody records what the handler reads, up to limit, so the body can be
// sent again to the next target. Decoders often stop at the end of the value
// without reading to EOF, so a replay is the recorded bytes followed by
// whatever the previous attempt left unread.
type replayBody struct {
src io.ReadCloser
buf bytes.Buffer
limit int
overflow bool
stopped bool
}
func (r *replayBody) Read(p []byte) (int, error) {
n, err := r.src.Read(p)
if n > 0 && !r.overflow && !r.stopped {
if r.buf.Len()+n > r.limit {
r.overflow = true
// A new buffer, not Reset: Reset keeps the memory.
r.buf = bytes.Buffer{}
} else {
r.buf.Write(p[:n])
}
}
return n, err
}
func (r *replayBody) Close() error { return r.src.Close() }
// replayable reports whether everything read so far was kept.
func (r *replayBody) replayable() bool { return !r.overflow && !r.stopped }
// stop ends recording and frees what was kept.
func (r *replayBody) stop() {
r.stopped = true
r.buf = bytes.Buffer{}
}
// replay rewinds to the start of the body and keeps recording, so a third
// attempt can replay too.
func (r *replayBody) replay() io.ReadCloser {
data := bytes.Clone(r.buf.Bytes())
rest := r.src
r.src = struct {
io.Reader
io.Closer
}{io.MultiReader(bytes.NewReader(data), rest), rest}
r.buf.Reset()
return r
}
+381
View File
@@ -0,0 +1,381 @@
package middleware
import (
"bytes"
"context"
"errors"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing/iotest"
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/routing/admission"
"github.com/mudler/LocalAI/pkg/model"
"github.com/mudler/LocalAI/pkg/system"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
"google.golang.org/grpc/codes"
grpcstatus "google.golang.org/grpc/status"
)
var _ = Describe("failover chains in the request pipeline", func() {
var (
app *echo.Echo
fm *failover.Manager
re *RequestExtractor
limiter *admission.Limiter
mu sync.Mutex
calls []string
behavior map[string]func(c echo.Context) error
)
served := func(c echo.Context) error {
cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
return c.JSON(http.StatusOK, map[string]string{"served": cfg.Name})
}
handler := func(c echo.Context) error {
cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
mu.Lock()
calls = append(calls, cfg.Name)
b := behavior[cfg.Name]
mu.Unlock()
if b == nil {
return served(c)
}
return b(c)
}
post := func(path, body string) *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
app.ServeHTTP(rec, req)
return rec
}
chat := func(model string) *httptest.ResponseRecorder {
return post("/v1/chat/completions", `{"model":"`+model+`","messages":[{"role":"user","content":"hi"}]}`)
}
BeforeEach(func() {
calls = nil
behavior = map[string]func(c echo.Context) error{}
dir := GinkgoT().TempDir()
write := func(name, body string) {
Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed())
}
write("a", "name: a\nbackend: fake-a\n")
write("b", "name: b\nbackend: fake-b\n")
write("plain", "name: plain\nbackend: fake-p\n")
write("chain", "name: chain\nfailover:\n targets:\n - model: a\n - model: b\n")
write("capped", "name: capped\nbackend: fake-c\nlimits:\n max_concurrent: 1\n")
write("off", "name: off\nbackend: fake-o\ndisabled: true\n")
write("chain-capped", "name: chain-capped\nfailover:\n targets:\n - model: capped\n - model: b\n")
write("chain-off", "name: chain-off\nfailover:\n targets:\n - model: off\n - model: b\n")
write("remote", "name: remote\nbackend: cloud-proxy\nproxy:\n mode: passthrough\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n")
write("remote-mapped", "name: remote-mapped\nbackend: cloud-proxy\nproxy:\n mode: translate\n provider: openai\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n upstream_model: big-llm\n")
write("chain-remote", "name: chain-remote\nfailover:\n targets:\n - model: remote\n - model: remote-mapped\n")
ss := &system.SystemState{Model: system.Model{ModelsPath: dir}}
appConfig := config.NewApplicationConfig()
appConfig.SystemState = ss
mcl := config.NewModelConfigLoader(dir)
Expect(mcl.LoadModelConfigsFromPath(dir)).To(Succeed())
re = NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig)
fm = failover.New(mcl)
// The scheduler's first tick syncs in the application; HasChains
// answers from the last sync.
fm.Sync()
re.SetFailoverManager(fm)
app = echo.New()
// echo's default handler hides internal error messages; the specs
// below check which target's error reached the client.
app.HTTPErrorHandler = func(err error, c echo.Context) {
code := http.StatusInternalServerError
var he *echo.HTTPError
if errors.As(err, &he) {
code = he.Code
}
_ = c.JSON(code, map[string]string{"error": err.Error()})
}
app.POST("/v1/chat/completions", handler,
re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }))
app.POST("/v1/audio/transcriptions", handler,
re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }))
limiter = admission.New()
app.POST("/v1/chat/admitted", handler,
re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }),
AdmissionControl(limiter, nil))
// The real transcription route resolves a default model first, which
// parses the multipart form before SetModelAndConfig runs.
app.POST("/v1/audio/transcriptions-default", handler,
re.BuildFilteredFirstAvailableDefaultModel(config.NoFilterFn),
re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }))
})
It("serves from the next target when the first fails before responding", func() {
behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: connection refused") }
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusOK))
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(rec.Header().Get(HeaderServedModel)).To(Equal("b"))
Expect(rec.Header().Get(HeaderFailover)).To(Equal("fallback"))
Expect(calls).To(Equal([]string{"a", "b"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Active).To(Equal("b"))
})
It("sends a remote target its own upstream model, not the chain name", func() {
upstream := map[string]string{}
record := func(c echo.Context) error {
cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
mu.Lock()
upstream[cfg.Name] = cfg.Proxy.UpstreamModel
mu.Unlock()
return nil
}
behavior["remote"] = func(c echo.Context) error {
_ = record(c)
return errors.New("dial tcp: connection refused")
}
behavior["remote-mapped"] = func(c echo.Context) error { _ = record(c); return served(c) }
Expect(chat("chain-remote").Code).To(Equal(http.StatusOK))
// The probe checks the same name, so a request cannot fail on a model
// the liveness probe just found.
Expect(upstream).To(Equal(map[string]string{"remote": "remote", "remote-mapped": "big-llm"}))
for _, name := range []string{"remote", "remote-mapped"} {
cfg, ok := re.modelConfigLoader.GetModelConfig(name)
Expect(ok).To(BeTrue())
Expect(upstream[name]).To(Equal(failover.UpstreamModel(cfg)))
}
stored, _ := re.modelConfigLoader.GetModelConfig("remote")
Expect(stored.Proxy.UpstreamModel).To(BeEmpty(), "the shared config must not change")
})
It("serves the primary without the failover header", func() {
rec := chat("chain")
Expect(rec.Header().Get(HeaderServedModel)).To(Equal("a"))
Expect(rec.Header().Get(HeaderFailover)).To(BeEmpty())
})
It("drops a buffered 5xx response and retries", func() {
behavior["a"] = func(c echo.Context) error {
return c.JSON(http.StatusServiceUnavailable, map[string]string{"error": "no healthy nodes"})
}
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusOK))
Expect(rec.Body.String()).ToNot(ContainSubstring("no healthy nodes"))
})
It("does not retry after streaming started, and trips the target", func() {
behavior["a"] = func(c echo.Context) error {
c.Response().Header().Set("Content-Type", "text/event-stream")
_, _ = c.Response().Write([]byte("data: x\n\n"))
c.Response().Flush()
return errors.New("connection reset by peer")
}
rec := chat("chain")
Expect(rec.Body.String()).To(HavePrefix("data: x"))
Expect(calls).To(Equal([]string{"a"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateDown))
})
It("does not retry or trip on 4xx", func() {
behavior["a"] = func(echo.Context) error { return echo.NewHTTPError(http.StatusBadRequest, "bad") }
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusBadRequest))
Expect(calls).To(Equal([]string{"a"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
It("does not retry when the client cancelled", func() {
ctx, cancel := context.WithCancel(context.Background())
behavior["a"] = func(echo.Context) error { cancel(); return context.Canceled }
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions",
strings.NewReader(`{"model":"chain","messages":[{"role":"user","content":"hi"}]}`)).WithContext(ctx)
req.Header.Set("Content-Type", "application/json")
app.ServeHTTP(httptest.NewRecorder(), req)
Expect(calls).To(Equal([]string{"a"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
It("gives each attempt a fresh request", func() {
behavior["a"] = func(c echo.Context) error {
in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest)
in.Messages = nil
return errors.New("dial tcp: connection refused")
}
behavior["b"] = func(c echo.Context) error {
in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest)
Expect(in.Messages).To(HaveLen(1))
return served(c)
}
Expect(chat("chain").Code).To(Equal(http.StatusOK))
})
It("degraded: tries every target in priority order and returns the last error", func() {
behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: a down") }
behavior["b"] = func(echo.Context) error { return errors.New("dial tcp: b down") }
chat("chain") // trips both
calls = nil
rec := chat("chain")
Expect(calls).To(Equal([]string{"a", "b"}))
Expect(rec.Code).To(Equal(http.StatusInternalServerError))
Expect(rec.Body.String()).To(ContainSubstring("b down"))
Expect(rec.Header().Get(HeaderFailover)).To(Equal("degraded"))
})
DescribeTable("replays a multipart body for the next target", func(path string) {
var body bytes.Buffer
mw := multipart.NewWriter(&body)
_ = mw.WriteField("model", "chain")
fw, _ := mw.CreateFormFile("file", "a.wav")
_, _ = fw.Write(bytes.Repeat([]byte{7}, 4096))
Expect(mw.Close()).To(Succeed())
size := func(c echo.Context) int64 {
fh, err := c.FormFile("file")
Expect(err).ToNot(HaveOccurred())
f, _ := fh.Open()
n, _ := io.Copy(io.Discard, f)
return n
}
behavior["a"] = func(c echo.Context) error { size(c); return errors.New("dial tcp: refused") }
behavior["b"] = func(c echo.Context) error {
Expect(size(c)).To(Equal(int64(4096)))
return served(c)
}
req := httptest.NewRequest(http.MethodPost, path, &body)
req.Header.Set("Content-Type", mw.FormDataContentType())
rec := httptest.NewRecorder()
app.ServeHTTP(rec, req)
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(calls).To(Equal([]string{"a", "b"}))
},
Entry("read by the handler", "/v1/audio/transcriptions"),
Entry("parsed before SetModelAndConfig", "/v1/audio/transcriptions-default"),
)
It("spills an admission rejection to the next target without tripping it", func() {
release, ok := limiter.Acquire("capped", 1)
Expect(ok).To(BeTrue())
defer release()
rec := post("/v1/chat/admitted", `{"model":"chain-capped","messages":[{"role":"user","content":"hi"}]}`)
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(rec.Header().Get("Retry-After")).To(BeEmpty())
st, _ := fm.ChainStatus("chain-capped")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
Expect(st.Active).To(Equal("capped"))
})
It("spills a capability gap (gRPC Unimplemented) to the next target without tripping it", func() {
behavior["a"] = func(echo.Context) error {
return grpcstatus.Error(codes.Unimplemented, "localai-proxy: Rerank has no upstream counterpart")
}
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(calls).To(Equal([]string{"a", "b"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
It("spills a written 501 to the next target without tripping it", func() {
behavior["a"] = func(c echo.Context) error {
return c.JSON(http.StatusNotImplemented, map[string]string{"error": "not supported"})
}
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(calls).To(Equal([]string{"a", "b"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
It("fails a rate-limited target (gRPC ResourceExhausted) over to the next target and trips it", func() {
behavior["a"] = func(echo.Context) error {
return grpcstatus.Error(codes.ResourceExhausted, "localai-proxy: upstream /v1/chat/completions returned 429: slow down")
}
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(calls).To(Equal([]string{"a", "b"}))
st, _ := fm.ChainStatus("chain")
Expect(st.Targets[0].State).To(Equal(failover.StateDown))
})
It("skips a disabled target without tripping it", func() {
rec := chat("chain-off")
Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String())
Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`))
Expect(calls).To(Equal([]string{"b"}))
st, _ := fm.ChainStatus("chain-off")
Expect(st.Targets[0].State).To(Equal(failover.StateHealthy))
})
It("counts a 4xx response as neither success nor failure", func() {
behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: a down") }
behavior["b"] = func(echo.Context) error { return errors.New("dial tcp: b down") }
chat("chain") // trips both
behavior["a"] = func(c echo.Context) error {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "bad"})
}
rec := chat("chain")
Expect(rec.Code).To(Equal(http.StatusBadRequest))
st, _ := fm.ChainStatus("chain")
// a is a cold local target: a success would have recovered it.
Expect(st.Targets[0].State).To(Equal(failover.StateDown))
})
It("stops recording the body once the model is known not to be a chain", func() {
behavior["plain"] = func(c echo.Context) error {
rb, ok := c.Request().Body.(*replayBody)
Expect(ok).To(BeTrue())
Expect(rb.replayable()).To(BeFalse())
Expect(rb.buf.Cap()).To(BeZero())
return served(c)
}
Expect(chat("plain").Code).To(Equal(http.StatusOK))
})
It("does not record the body without a failover manager", func() {
re.SetFailoverManager(nil)
behavior["plain"] = func(c echo.Context) error {
_, ok := c.Request().Body.(*replayBody)
Expect(ok).To(BeFalse())
return served(c)
}
Expect(chat("plain").Code).To(Equal(http.StatusOK))
})
It("releases the recorded bytes on overflow", func() {
// One byte per read, so some bytes are recorded before the limit is hit.
rb := &replayBody{src: io.NopCloser(iotest.OneByteReader(bytes.NewReader(make([]byte, 64)))), limit: 16}
_, _ = io.ReadAll(rb)
Expect(rb.replayable()).To(BeFalse())
Expect(rb.buf.Cap()).To(BeZero())
})
It("leaves plain models untouched", func() {
behavior["plain"] = func(c echo.Context) error {
return c.JSON(http.StatusInternalServerError, map[string]string{"error": "boom"})
}
rec := chat("plain")
Expect(rec.Code).To(Equal(http.StatusInternalServerError))
Expect(rec.Body.String()).To(ContainSubstring("boom"))
Expect(rec.Header().Get(HeaderServedModel)).To(BeEmpty())
})
})
+22 -2
View File
@@ -12,6 +12,7 @@ import (
"github.com/labstack/echo/v4"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/schema"
"github.com/mudler/LocalAI/core/services/failover"
"github.com/mudler/LocalAI/core/services/galleryop"
"github.com/mudler/LocalAI/core/templates"
"github.com/mudler/LocalAI/pkg/distributedhdr"
@@ -29,6 +30,7 @@ type RequestExtractor struct {
modelConfigLoader *config.ModelConfigLoader
modelLoader *model.ModelLoader
applicationConfig *config.ApplicationConfig
failover *failover.Manager
}
func NewRequestExtractor(modelConfigLoader *config.ModelConfigLoader, modelLoader *model.ModelLoader, applicationConfig *config.ApplicationConfig) *RequestExtractor {
@@ -122,7 +124,7 @@ func (re *RequestExtractor) BuildFilteredFirstAvailableDefaultModel(filterFn con
// Otherwise, it's in its own method below for now
func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIRequest) echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
return re.failoverRetry(func(c echo.Context) error {
input := initializer()
if input == nil {
return echo.NewHTTPError(http.StatusBadRequest, "unable to initialize body")
@@ -194,6 +196,24 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR
cfg = resolved
}
// A failover chain resolves to one of its targets, like an alias.
// failoverRetry re-runs this middleware for the next target.
if cfg != nil && cfg.IsFailover() {
resolved, fErr := re.resolveFailover(c, modelName, cfg)
if fErr != nil {
return c.JSON(http.StatusServiceUnavailable, schema.ErrorResponse{
Error: &schema.APIError{
Message: fErr.Error(),
Code: http.StatusServiceUnavailable,
Type: "failover_unavailable",
},
})
}
cfg = resolved
} else {
stopFailoverRecording(c)
}
// Check if the model is disabled
if cfg != nil && cfg.IsDisabled() {
return c.JSON(http.StatusForbidden, schema.ErrorResponse{
@@ -209,7 +229,7 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR
c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg)
return next(c)
}
})
}
}
@@ -0,0 +1,175 @@
import { test, expect } from './coverage-fixtures.js'
// Failover Chain template + FailoverTargetsEditor regression tests.
//
// A failover chain is a model config whose `failover.targets` field lists,
// in order, the downstream models that answer for one name — the first
// healthy target serves each request, the rest take over when it fails.
// This covers:
// - the create-flow template gallery exposes a "Failover Chain" card that
// seeds a minimal name + two empty targets
// - the dedicated FailoverTargetsEditor renders {model, warm} rows with
// add/remove/move controls
// - inline warnings for a single-target chain, a duplicated target model,
// and a target that names the chain itself
// - saving an edited chain sends the updated failover.targets array in
// the PATCH body
const FAILOVER_METADATA = {
sections: [
{ id: 'general', label: 'General', icon: 'settings', order: 0 },
{ id: 'failover', label: 'Failover', icon: 'shuffle', order: 90 },
],
fields: [
{
path: 'name', yaml_key: 'name', go_type: 'string', ui_type: 'string',
section: 'general', label: 'Model Name', component: 'input', order: 0,
},
{
path: 'failover.targets', yaml_key: 'targets', go_type: '[]FailoverTarget', ui_type: 'object',
section: 'failover', label: 'Failover targets', component: 'failover-targets',
description: 'Ordered list of models that serve this chain.',
order: 1,
},
],
}
async function mockCommon(page) {
await page.route('**/api/auth/status', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ authEnabled: false, staticApiKeyRequired: false, providers: [] }) }))
await page.route('**/api/models/config-metadata*', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(FAILOVER_METADATA) }))
await page.route('**/api/models/config-metadata/autocomplete/**', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ values: [] }) }))
page.on('pageerror', (err) => {
throw new Error(`uncaught page error: ${err.message}`)
})
}
test.describe('Failover Chain template - create flow', () => {
test.beforeEach(async ({ page }) => {
await mockCommon(page)
})
test('template gallery exposes the Failover Chain card', async ({ page }) => {
await page.goto('/app/model-editor')
await expect(page.getByRole('button', { name: /Failover Chain/i })).toBeVisible({ timeout: 10_000 })
})
test('failover template loads the editor with two target rows', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
await expect(page.getByText(/Unexpected Application Error/i)).toHaveCount(0)
await expect(page.locator('h1.page-title')).toBeVisible({ timeout: 10_000 })
await expect(page.getByText('Failover targets').first()).toBeVisible()
// Two empty {model: ''} rows seeded by the template.
await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(2)
await expect(page.getByRole('button', { name: /Add target/i }).first()).toBeVisible()
})
test('Add target adds a row', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
await page.getByRole('button', { name: /Add target/i }).first().click()
await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(3)
})
test('removing a row removes it', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
await page.locator('button[title="Remove target"]').first().click()
await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(1)
})
test('move down/up reorders the target rows', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
const rows = page.locator('input[placeholder="target model..."]')
await rows.nth(0).fill('model-a')
await rows.nth(1).fill('model-b')
// Move the first row down — it should now be the second row's value.
await page.locator('button[title="Move down"]').first().click()
await expect(rows.nth(0)).toHaveValue('model-b')
await expect(rows.nth(1)).toHaveValue('model-a')
// Move it back up with the second row's "Move up" control.
await page.locator('button[title="Move up (tried earlier)"]').nth(1).click()
await expect(rows.nth(0)).toHaveValue('model-a')
await expect(rows.nth(1)).toHaveValue('model-b')
})
test('a single remaining target shows a too-few warning', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
await page.locator('button[title="Remove target"]').first().click()
await expect(page.locator('input[placeholder="target model..."]')).toHaveCount(1)
await expect(page.getByText(/nothing to fail over to/i)).toBeVisible()
})
test('duplicate target models flag both rows', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
const rows = page.locator('input[placeholder="target model..."]')
await rows.nth(0).fill('same-model')
await rows.nth(1).fill('same-model')
await expect(page.getByText(/Duplicate target/i)).toHaveCount(2)
})
test('a target naming the chain itself is flagged', async ({ page }) => {
await page.goto('/app/model-editor?template=failover')
// Create mode renders the model name through a dedicated input (not the
// generic field renderer), whose placeholder comes from the modelEditor
// i18n namespace rather than the field registry.
await page.locator('input[placeholder="my-model-name"]').fill('my-chain')
const rows = page.locator('input[placeholder="target model..."]')
await rows.nth(0).fill('my-chain')
await expect(page.getByText(/chain's own name/i)).toBeVisible()
})
})
test.describe('Failover Chain - saving an edited chain', () => {
const MOCK_YAML = 'name: my-chain\nfailover:\n targets:\n - model: model-a\n warm: true\n - model: model-b\n warm: false\n'
test.beforeEach(async ({ page }) => {
await mockCommon(page)
await page.route('**/api/models/edit/my-chain', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: MOCK_YAML, name: 'my-chain' }) }))
// ModelFailoverStatus mounts unconditionally in edit mode; keep its
// fetch + SSE subscription harmless for a chain the test doesn't care
// about the live status of.
await page.route('**/api/failover', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [] }) }))
await page.route('**/api/failover/events', (route) =>
route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: '' }))
})
test('saving sends the updated failover.targets array in the PATCH body', async ({ page }) => {
let patchBody = null
await page.route('**/api/models/config-json/my-chain', (route) => {
if (route.request().method() === 'PATCH') {
patchBody = route.request().postDataJSON()
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ success: true, message: "Model 'my-chain' updated successfully" }) })
} else {
route.fulfill({ contentType: 'application/json', body: '{}' })
}
})
await page.goto('/app/model-editor/my-chain')
await expect(page.locator('h1', { hasText: 'Model Editor' })).toBeVisible({ timeout: 10_000 })
// Existing targets loaded from YAML.
const rows = page.locator('input[placeholder="target model..."]')
await expect(rows).toHaveCount(2)
// Add a third target, then save.
await page.getByRole('button', { name: /Add target/i }).first().click()
await rows.nth(2).fill('model-c')
await page.locator('button', { hasText: 'Save Changes' }).click()
await expect(page.locator('text=Configuration saved')).toBeVisible({ timeout: 5_000 })
expect(patchBody).toBeTruthy()
expect(patchBody.failover.targets).toEqual([
{ model: 'model-a', warm: true },
{ model: 'model-b', warm: false },
{ model: 'model-c', warm: false },
])
})
})
@@ -0,0 +1,124 @@
import { test, expect } from './coverage-fixtures.js'
// Live failover chain health in the Model Editor. A model that is a failover
// chain shows a strip under the editor header with the chain state, the
// active target, and a per-target health table fed by GET /api/failover plus
// the /api/failover/events SSE stream. Pin controls are admin-only.
const MOCK_METADATA = {
sections: [{ id: 'general', label: 'General', icon: 'settings', order: 0 }],
fields: [
{ path: 'name', yaml_key: 'name', go_type: 'string', ui_type: 'string', section: 'general', label: 'Model Name', description: 'id', component: 'input', order: 0 },
],
}
const MOCK_YAML = 'name: chain\nfailover:\n targets: [a, b]\n'
const CHAIN = {
name: 'chain',
state: 'primary',
active: 'a',
active_since: '2026-09-26T09:00:00Z',
pinned: null,
targets: [
{ model: 'a', kind: 'local', warm: true, state: 'healthy', consecutive_ok: 5, last_probe: '2026-09-26T09:59:00Z' },
{ model: 'b', kind: 'remote', warm: false, state: 'healthy', consecutive_ok: 3, last_error: 'dial tcp: connection refused while probing upstream' },
],
}
// The stream replays a snapshot (primary on a) then a switch to b. The body
// ends after the two frames; EventSource reconnects and replays them, so the
// settled state stays fallback on b.
const SSE_BODY =
`event: snapshot\ndata: ${JSON.stringify({ chains: [CHAIN] })}\n\n` +
'event: chain.switched\ndata: {"type":"chain.switched","chain":"chain","from":"a","to":"b","state":"fallback","reason":"trip","at":"2026-09-26T10:00:00Z"}\n\n'
async function mockEditor(page, authStatus) {
await page.route('**/api/auth/status', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(authStatus) }))
await page.route('**/api/models/config-metadata*', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK_METADATA) }))
await page.route('**/api/models/edit/**', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: MOCK_YAML, name: 'chain' }) }))
await page.route('**/api/models/config-json/**', (route) =>
route.fulfill({ contentType: 'application/json', body: '{}' }))
await page.route('**/api/failover', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [CHAIN] }) }))
await page.route('**/api/failover/chain', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(CHAIN) }))
await page.route('**/api/failover/events', (route) =>
route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: SSE_BODY }))
}
const NO_AUTH = { authEnabled: false, staticApiKeyRequired: false, providers: [] }
const NON_ADMIN = {
authEnabled: true,
staticApiKeyRequired: false,
providers: ['local'],
user: { id: 'user-uuid', name: 'User', role: 'user', provider: 'local' },
}
test.describe('Model Editor — failover chain health', () => {
test('shows the chain strip and follows chain.switched to fallback', async ({ page }) => {
await mockEditor(page, NO_AUTH)
await page.goto('/app/model-editor/chain')
const strip = page.locator('.failover-status')
await expect(strip).toBeVisible({ timeout: 10_000 })
const chainPill = strip.locator('.failover-status__state .status-pill')
await expect(chainPill).toHaveText(/fallback/i)
await expect(chainPill).toHaveClass(/status-pill--warning/)
await expect(strip.locator('.failover-status__active')).toHaveText('b')
// Both targets are listed with their health.
const rows = strip.locator('tbody tr')
await expect(rows).toHaveCount(2)
await expect(rows.nth(0)).toContainText('a')
await expect(rows.nth(1).locator('.status-pill')).toHaveClass(/status-pill--success/)
// Long errors are truncated but kept whole in the title.
await expect(rows.nth(1).locator('.failover-status__error')).toHaveAttribute('title', /connection refused/)
})
test('admin can pin a target after confirming', async ({ page }) => {
await mockEditor(page, NO_AUTH)
let pinBody = null
await page.route('**/api/failover/chain/pin', async (route) => {
if (route.request().method() === 'POST') pinBody = route.request().postDataJSON()
await route.fulfill({ contentType: 'application/json', body: JSON.stringify({ ...CHAIN, pinned: 'b' }) })
})
await page.goto('/app/model-editor/chain')
const strip = page.locator('.failover-status')
await expect(strip).toBeVisible({ timeout: 10_000 })
const pinB = strip.locator('tbody tr').nth(1).getByRole('button', { name: /pin/i })
await expect(pinB).toBeVisible()
await pinB.click()
const dialog = page.getByRole('alertdialog')
await expect(dialog).toBeVisible()
await dialog.getByRole('button', { name: /^pin$/i }).click()
await expect.poll(() => pinBody).toEqual({ target: 'b' })
})
// The editor route itself is admin-only (RequireAdmin), so a non-admin never
// reaches the strip or its pin controls; the component additionally gates
// pinning on isAdmin for surfaces that are not admin-gated.
test('non-admin users get no pin controls', async ({ page }) => {
await mockEditor(page, NON_ADMIN)
await page.route('**/api/auth/me', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(NON_ADMIN.user) }))
await page.goto('/app/model-editor/chain')
await page.waitForURL(/\/app(?!\/model-editor)/, { timeout: 5000 })
await expect(page.locator('.failover-status')).toHaveCount(0)
await expect(page.getByRole('button', { name: /^pin/i })).toHaveCount(0)
})
test('models that are not failover chains show no strip', async ({ page }) => {
await mockEditor(page, NO_AUTH)
await page.route('**/api/models/edit/**', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ config: 'name: plain\n', name: 'plain' }) }))
await page.goto('/app/model-editor/plain')
await expect(page.locator('h1.page-title')).toBeVisible({ timeout: 10_000 })
await expect(page.locator('.failover-status')).toHaveCount(0)
})
})
@@ -0,0 +1,134 @@
import { test, expect } from './coverage-fixtures.js'
// Failover overview page (Operate -> Runtime -> Failover) and the matching
// "chain" badge on an installed model that is itself a failover chain.
//
// The overview is a dense table over GET /api/failover, kept live by the
// /api/failover/events SSE stream (same hook the Model Editor's chain strip
// uses), plus an empty state that points at the Failover Chain template.
const NO_AUTH = { authEnabled: false, staticApiKeyRequired: false, providers: [] }
const CHAIN_A = {
name: 'chain-a',
state: 'primary',
active: 'a',
active_since: '2026-09-26T09:00:00Z',
pinned: null,
targets: [
{ model: 'a', kind: 'local', warm: true, state: 'healthy', last_probe: '2026-09-26T09:59:00Z' },
{ model: 'b', kind: 'remote', warm: false, state: 'healthy' },
],
}
const CHAIN_B = {
name: 'chain-b',
state: 'degraded',
active: 'x',
active_since: '2026-09-26T08:00:00Z',
pinned: null,
targets: [
{ model: 'x', kind: 'local', warm: true, state: 'down' },
{ model: 'y', kind: 'local', warm: false, state: 'healthy' },
],
}
async function mockAuth(page, authStatus = NO_AUTH) {
await page.route('**/api/auth/status', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify(authStatus) }))
}
async function mockFailover(page, chains) {
await page.route('**/api/failover', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains }) }))
await page.route('**/api/failover/events', (route) =>
route.fulfill({
status: 200,
headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' },
body: `event: snapshot\ndata: ${JSON.stringify({ chains })}\n\n`,
}))
}
test.describe('Failover overview', () => {
test.beforeEach(async ({ page }) => {
page.on('pageerror', (err) => {
throw new Error(`uncaught page error: ${err.message}`)
})
})
test('admin sees the Failover entry in the Operate rail', async ({ page }) => {
await mockAuth(page)
await mockFailover(page, [CHAIN_A])
await page.goto('/app/operate')
const rail = page.locator('.console-layout > .console-rail')
const link = rail.locator('a.nav-item[href="/app/failover"]')
await expect(link).toBeVisible({ timeout: 10_000 })
await expect(link).toContainText('Failover')
})
test('lists chains from GET /api/failover with their live health', async ({ page }) => {
await mockAuth(page)
await mockFailover(page, [CHAIN_A, CHAIN_B])
await page.goto('/app/failover')
const rows = page.locator('[data-testid="failover-overview"] tbody tr')
await expect(rows).toHaveCount(2)
const first = rows.nth(0)
await expect(first.getByRole('link', { name: 'chain-a' })).toHaveAttribute('href', '/app/model-editor/chain-a')
await expect(first.locator('.status-pill').first()).toHaveText(/primary/i)
await expect(first).toContainText('a')
await expect(first.locator('.failover-overview__targets .status-pill')).toHaveCount(2)
const second = rows.nth(1)
await expect(second.getByRole('link', { name: 'chain-b' })).toHaveAttribute('href', '/app/model-editor/chain-b')
await expect(second.locator('.status-pill').first()).toHaveText(/degraded/i)
})
test('follows target.state over SSE without a page reload', async ({ page }) => {
await mockAuth(page)
await page.route('**/api/failover', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ chains: [CHAIN_A] }) }))
const sseBody =
`event: snapshot\ndata: ${JSON.stringify({ chains: [CHAIN_A] })}\n\n` +
'event: target.state\ndata: {"type":"target.state","chain":"chain-a","target":"b","from":"healthy","to":"down","error":"dial tcp: refused"}\n\n'
await page.route('**/api/failover/events', (route) =>
route.fulfill({ status: 200, headers: { 'Content-Type': 'text/event-stream', 'Cache-Control': 'no-cache' }, body: sseBody }))
await page.goto('/app/failover')
const row = page.locator('[data-testid="failover-overview"] tbody tr').first()
await expect(row).toBeVisible({ timeout: 10_000 })
const targetB = row.locator('.failover-overview__targets .status-pill', { hasText: 'b' })
await expect(targetB).toHaveClass(/status-pill--error/)
})
test('shows an empty state pointing at the Failover Chain template', async ({ page }) => {
await mockAuth(page)
await mockFailover(page, [])
await page.goto('/app/failover')
await expect(page.locator('.empty-state')).toBeVisible({ timeout: 10_000 })
const cta = page.getByRole('link', { name: /create a failover chain/i })
await expect(cta).toHaveAttribute('href', '/app/model-editor?template=failover')
})
})
test.describe('Installed Models - chain badge', () => {
test.beforeEach(async ({ page }) => {
await mockAuth(page)
await page.route('**/api/models/capabilities', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify({ data: [
{ id: 'chain-a', capabilities: ['chat'], backend: 'llama-cpp' },
] }) }))
await page.route('**/api/aliases', (route) =>
route.fulfill({ contentType: 'application/json', body: JSON.stringify([]) }))
await mockFailover(page, [CHAIN_A])
})
test('renders a read-only chain -> target badge on a chain model', async ({ page }) => {
await page.goto('/app/models?view=installed')
await page.locator('[data-entity="chain-a"]').click()
await expect(page.getByText('chain → a')).toBeVisible({ timeout: 10_000 })
})
})
@@ -47,7 +47,7 @@ test.describe('Nodes fleet dashboard', () => {
const rail = page.locator('.console-layout > .console-rail')
await expect(rail).toBeVisible()
await expect(rail.locator('a.nav-item')).toHaveCount(13)
await expect(rail.locator('a.nav-item')).toHaveCount(14)
await expect(rail.locator('a[href="/app/nodes"]')).toHaveClass(/active/)
await expect(rail.locator('a[href$="/swagger/index.html"]')).toHaveAttribute('target', '_blank')
})
@@ -67,7 +67,7 @@ test.describe('Nodes fleet dashboard', () => {
await expect(rail.locator('.console-rail-groups')).toBeHidden()
await rail.getByRole('button', { name: 'Expand Operate navigation' }).click()
await expect(rail.locator('.console-rail-groups')).toBeVisible()
await expect(rail.locator('a.nav-item')).toHaveCount(13)
await expect(rail.locator('a.nav-item')).toHaveCount(14)
})
test('shows aggregate health, capacity, attention filtering, search, sorting, and grouping', async ({ page }) => {
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -224,5 +225,67 @@
"pickVision": "Bilder und Dokumente lesen",
"pickAudio": "Sprache rein, Sprache raus",
"pickVisual": "Bilder und Video erzeugen"
},
"failover": {
"title": "Failover-Kette",
"servedBy": "Bedient von",
"changed": "geändert {{time}}",
"pinnedTo": "Fixiert auf {{target}}",
"warm": "Vorgeladen",
"never": "Nie",
"states": {
"primary": "Primär",
"fallback": "Fallback",
"degraded": "Beeinträchtigt",
"healthy": "Gesund",
"down": "Ausgefallen",
"recovering": "Erholt sich",
"missing": "Fehlt"
},
"kinds": {
"local": "Lokal",
"remote": "Remote"
},
"columns": {
"target": "Ziel",
"kind": "Art",
"warm": "Vorgeladen",
"health": "Zustand",
"lastProbe": "Letzte Prüfung",
"lastError": "Letzter Fehler",
"actions": "Aktionen"
},
"actions": {
"pin": "Fixieren",
"pinning": "Wird fixiert…",
"unpin": "Lösen",
"unpinning": "Wird gelöst…"
},
"confirm": {
"pinTitle": "{{target}} fixieren?",
"pinMessage": "Nur {{target}} bedient {{chain}}, bis Sie die Fixierung lösen. Das automatische Failover ist für diese Kette angehalten.",
"unpinTitle": "{{chain}} lösen?",
"unpinMessage": "{{chain}} kehrt zum automatischen Failover zurück."
},
"errors": {
"pin": "Fixieren fehlgeschlagen: {{error}}",
"unpin": "Lösen fehlgeschlagen: {{error}}"
},
"overview": {
"title": "Failover-Ketten",
"subtitle": "Live-Zustand und aktives Ziel für jede Failover-Kette.",
"columns": {
"chain": "Kette",
"state": "Status",
"active": "Aktives Ziel",
"targets": "Ziele",
"changed": "Geändert"
},
"empty": {
"title": "Noch keine Failover-Ketten",
"text": "Fügen Sie eine Failover-Kette hinzu, um fehlerhafte Ziele automatisch zu umgehen.",
"cta": "Failover-Kette erstellen"
}
}
}
}
@@ -60,7 +60,8 @@
"middleware": "Middleware",
"activity": "Aktivität",
"overview": "Übersicht",
"thisMachine": "Dieser Rechner"
"thisMachine": "Dieser Rechner",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -241,5 +242,67 @@
"pickVision": "Read images and documents",
"pickAudio": "Speech in and speech out",
"pickVisual": "Generate images and video"
},
"failover": {
"title": "Failover chain",
"servedBy": "Served by",
"changed": "changed {{time}}",
"pinnedTo": "Pinned to {{target}}",
"warm": "Warm",
"never": "Never",
"states": {
"primary": "Primary",
"fallback": "Fallback",
"degraded": "Degraded",
"healthy": "Healthy",
"down": "Down",
"recovering": "Recovering",
"missing": "Missing"
},
"kinds": {
"local": "Local",
"remote": "Remote"
},
"columns": {
"target": "Target",
"kind": "Kind",
"warm": "Warm",
"health": "Health",
"lastProbe": "Last probe",
"lastError": "Last error",
"actions": "Actions"
},
"actions": {
"pin": "Pin",
"pinning": "Pinning…",
"unpin": "Unpin",
"unpinning": "Unpinning…"
},
"confirm": {
"pinTitle": "Pin {{target}}?",
"pinMessage": "Only {{target}} serves {{chain}} until you unpin it. Automatic failover stops for this chain.",
"unpinTitle": "Unpin {{chain}}?",
"unpinMessage": "{{chain}} returns to automatic failover."
},
"errors": {
"pin": "Pin failed: {{error}}",
"unpin": "Unpin failed: {{error}}"
},
"overview": {
"title": "Failover chains",
"subtitle": "Live health and the active target for every failover chain.",
"columns": {
"chain": "Chain",
"state": "State",
"active": "Active target",
"targets": "Targets",
"changed": "Changed"
},
"empty": {
"title": "No failover chains yet",
"text": "Add a failover chain to route around unhealthy targets automatically.",
"cta": "Create a failover chain"
}
}
}
}
@@ -61,7 +61,8 @@
"api": "API",
"activity": "Activity",
"overview": "Overview",
"thisMachine": "This machine"
"thisMachine": "This machine",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -224,5 +225,67 @@
"pickVision": "Leer imágenes y documentos",
"pickAudio": "Voz de entrada y de salida",
"pickVisual": "Generar imágenes y vídeo"
},
"failover": {
"title": "Cadena de failover",
"servedBy": "Servido por",
"changed": "cambió {{time}}",
"pinnedTo": "Fijado en {{target}}",
"warm": "Precargado",
"never": "Nunca",
"states": {
"primary": "Primario",
"fallback": "Respaldo",
"degraded": "Degradado",
"healthy": "Correcto",
"down": "Caído",
"recovering": "Recuperándose",
"missing": "Ausente"
},
"kinds": {
"local": "Local",
"remote": "Remoto"
},
"columns": {
"target": "Destino",
"kind": "Tipo",
"warm": "Precargado",
"health": "Estado",
"lastProbe": "Última comprobación",
"lastError": "Último error",
"actions": "Acciones"
},
"actions": {
"pin": "Fijar",
"pinning": "Fijando…",
"unpin": "Liberar",
"unpinning": "Liberando…"
},
"confirm": {
"pinTitle": "¿Fijar {{target}}?",
"pinMessage": "Solo {{target}} atiende {{chain}} hasta que lo liberes. El failover automático se detiene para esta cadena.",
"unpinTitle": "¿Liberar {{chain}}?",
"unpinMessage": "{{chain}} vuelve al failover automático."
},
"errors": {
"pin": "No se pudo fijar: {{error}}",
"unpin": "No se pudo liberar: {{error}}"
},
"overview": {
"title": "Cadenas de failover",
"subtitle": "Estado en vivo y destino activo de cada cadena de failover.",
"columns": {
"chain": "Cadena",
"state": "Estado",
"active": "Destino activo",
"targets": "Destinos",
"changed": "Cambiado"
},
"empty": {
"title": "Aún no hay cadenas de failover",
"text": "Agrega una cadena de failover para evitar automáticamente los destinos en mal estado.",
"cta": "Crear una cadena de failover"
}
}
}
}
@@ -60,7 +60,8 @@
"middleware": "Middleware",
"activity": "Actividad",
"overview": "Resumen",
"thisMachine": "Esta máquina"
"thisMachine": "Esta máquina",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -237,5 +238,67 @@
"pickVision": "Membaca gambar dan dokumen",
"pickAudio": "Suara masuk dan keluar",
"pickVisual": "Membuat gambar dan video"
},
"failover": {
"title": "Rantai failover",
"servedBy": "Dilayani oleh",
"changed": "berubah {{time}}",
"pinnedTo": "Disematkan ke {{target}}",
"warm": "Siap",
"never": "Belum pernah",
"states": {
"primary": "Utama",
"fallback": "Cadangan",
"degraded": "Menurun",
"healthy": "Sehat",
"down": "Mati",
"recovering": "Memulihkan",
"missing": "Tidak ada"
},
"kinds": {
"local": "Lokal",
"remote": "Jarak jauh"
},
"columns": {
"target": "Target",
"kind": "Jenis",
"warm": "Siap",
"health": "Kesehatan",
"lastProbe": "Pemeriksaan terakhir",
"lastError": "Kesalahan terakhir",
"actions": "Aksi"
},
"actions": {
"pin": "Sematkan",
"pinning": "Menyematkan…",
"unpin": "Lepas",
"unpinning": "Melepas…"
},
"confirm": {
"pinTitle": "Sematkan {{target}}?",
"pinMessage": "Hanya {{target}} yang melayani {{chain}} sampai Anda melepasnya. Failover otomatis berhenti untuk rantai ini.",
"unpinTitle": "Lepas {{chain}}?",
"unpinMessage": "{{chain}} kembali ke failover otomatis."
},
"errors": {
"pin": "Gagal menyematkan: {{error}}",
"unpin": "Gagal melepas: {{error}}"
},
"overview": {
"title": "Rantai failover",
"subtitle": "Kesehatan langsung dan target aktif untuk setiap rantai failover.",
"columns": {
"chain": "Rantai",
"state": "Status",
"active": "Target aktif",
"targets": "Target",
"changed": "Berubah"
},
"empty": {
"title": "Belum ada rantai failover",
"text": "Tambahkan rantai failover untuk melewati target yang tidak sehat secara otomatis.",
"cta": "Buat rantai failover"
}
}
}
}
@@ -61,7 +61,8 @@
"api": "API",
"activity": "Aktivitas",
"overview": "Ikhtisar",
"thisMachine": "Mesin ini"
"thisMachine": "Mesin ini",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -224,5 +225,67 @@
"pickVision": "Leggere immagini e documenti",
"pickAudio": "Voce in ingresso e in uscita",
"pickVisual": "Generare immagini e video"
},
"failover": {
"title": "Catena di failover",
"servedBy": "Servito da",
"changed": "cambiato {{time}}",
"pinnedTo": "Fissato su {{target}}",
"warm": "Pronto",
"never": "Mai",
"states": {
"primary": "Primario",
"fallback": "Fallback",
"degraded": "Degradato",
"healthy": "Integro",
"down": "Non disponibile",
"recovering": "In ripristino",
"missing": "Mancante"
},
"kinds": {
"local": "Locale",
"remote": "Remoto"
},
"columns": {
"target": "Destinazione",
"kind": "Tipo",
"warm": "Pronto",
"health": "Stato",
"lastProbe": "Ultimo controllo",
"lastError": "Ultimo errore",
"actions": "Azioni"
},
"actions": {
"pin": "Fissa",
"pinning": "Fissaggio…",
"unpin": "Sblocca",
"unpinning": "Sblocco…"
},
"confirm": {
"pinTitle": "Fissare {{target}}?",
"pinMessage": "Solo {{target}} serve {{chain}} finché non lo sblocchi. Il failover automatico si ferma per questa catena.",
"unpinTitle": "Sbloccare {{chain}}?",
"unpinMessage": "{{chain}} torna al failover automatico."
},
"errors": {
"pin": "Fissaggio non riuscito: {{error}}",
"unpin": "Sblocco non riuscito: {{error}}"
},
"overview": {
"title": "Catene di failover",
"subtitle": "Stato in tempo reale e destinazione attiva per ogni catena di failover.",
"columns": {
"chain": "Catena",
"state": "Stato",
"active": "Destinazione attiva",
"targets": "Destinazioni",
"changed": "Cambiato"
},
"empty": {
"title": "Nessuna catena di failover",
"text": "Aggiungi una catena di failover per aggirare automaticamente le destinazioni non integre.",
"cta": "Crea una catena di failover"
}
}
}
}
@@ -60,7 +60,8 @@
"middleware": "Middleware",
"activity": "Attività",
"overview": "Panoramica",
"thisMachine": "Questa macchina"
"thisMachine": "Questa macchina",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -208,5 +209,67 @@
"pickVision": "이미지와 문서 읽기",
"pickAudio": "음성 입력과 출력",
"pickVisual": "이미지와 영상 생성"
},
"failover": {
"title": "페일오버 체인",
"servedBy": "처리 대상",
"changed": "{{time}} 변경됨",
"pinnedTo": "{{target}}에 고정됨",
"warm": "예열됨",
"never": "없음",
"states": {
"primary": "기본",
"fallback": "대체",
"degraded": "성능 저하",
"healthy": "정상",
"down": "중단",
"recovering": "복구 중",
"missing": "없음"
},
"kinds": {
"local": "로컬",
"remote": "원격"
},
"columns": {
"target": "대상",
"kind": "종류",
"warm": "예열",
"health": "상태",
"lastProbe": "마지막 확인",
"lastError": "마지막 오류",
"actions": "작업"
},
"actions": {
"pin": "고정",
"pinning": "고정 중…",
"unpin": "고정 해제",
"unpinning": "고정 해제 중…"
},
"confirm": {
"pinTitle": "{{target}}을(를) 고정할까요?",
"pinMessage": "고정을 해제할 때까지 {{target}}만 {{chain}}을(를) 처리합니다. 이 체인의 자동 페일오버가 중지됩니다.",
"unpinTitle": "{{chain}} 고정을 해제할까요?",
"unpinMessage": "{{chain}}이(가) 자동 페일오버로 돌아갑니다."
},
"errors": {
"pin": "고정 실패: {{error}}",
"unpin": "고정 해제 실패: {{error}}"
},
"overview": {
"title": "페일오버 체인",
"subtitle": "모든 페일오버 체인의 실시간 상태와 활성 대상입니다.",
"columns": {
"chain": "체인",
"state": "상태",
"active": "활성 대상",
"targets": "대상",
"changed": "변경됨"
},
"empty": {
"title": "아직 페일오버 체인이 없습니다",
"text": "상태가 나빌 대상을 자동으로 우회하려면 페일오버 체인을 추가하세요.",
"cta": "페일오버 체인 만들기"
}
}
}
}
@@ -60,7 +60,8 @@
"api": "API",
"activity": "활동",
"overview": "개요",
"thisMachine": "이 머신"
"thisMachine": "이 머신",
"failover": "페일오버"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -240,5 +241,67 @@
"pickVision": "Leia imagens e documentos",
"pickAudio": "Fala de entrada e saída",
"pickVisual": "Gere imagens e vídeos"
},
"failover": {
"title": "Cadeia de failover",
"servedBy": "Atendido por",
"changed": "alterado {{time}}",
"pinnedTo": "Fixado em {{target}}",
"warm": "Pré-carregado",
"never": "Nunca",
"states": {
"primary": "Primário",
"fallback": "Reserva",
"degraded": "Degradado",
"healthy": "Saudável",
"down": "Fora do ar",
"recovering": "Recuperando",
"missing": "Ausente"
},
"kinds": {
"local": "Local",
"remote": "Remoto"
},
"columns": {
"target": "Destino",
"kind": "Tipo",
"warm": "Pré-carregado",
"health": "Saúde",
"lastProbe": "Última verificação",
"lastError": "Último erro",
"actions": "Ações"
},
"actions": {
"pin": "Fixar",
"pinning": "Fixando…",
"unpin": "Desafixar",
"unpinning": "Desafixando…"
},
"confirm": {
"pinTitle": "Fixar {{target}}?",
"pinMessage": "Somente {{target}} atende {{chain}} até você desafixar. O failover automático para nesta cadeia.",
"unpinTitle": "Desafixar {{chain}}?",
"unpinMessage": "{{chain}} volta ao failover automático."
},
"errors": {
"pin": "Falha ao fixar: {{error}}",
"unpin": "Falha ao desafixar: {{error}}"
},
"overview": {
"title": "Cadeias de failover",
"subtitle": "Saúde em tempo real e destino ativo de cada cadeia de failover.",
"columns": {
"chain": "Cadeia",
"state": "Estado",
"active": "Destino ativo",
"targets": "Destinos",
"changed": "Alterado"
},
"empty": {
"title": "Ainda não há cadeias de failover",
"text": "Adicione uma cadeia de failover para contornar automaticamente destinos com problemas.",
"cta": "Criar uma cadeia de failover"
}
}
}
}
@@ -61,7 +61,8 @@
"api": "API",
"activity": "Atividade",
"overview": "Visão geral",
"thisMachine": "Esta máquina"
"thisMachine": "Esta máquina",
"failover": "Failover"
},
"footer": {
"github": "GitHub",
@@ -40,7 +40,8 @@
"license": "License", "tags": "Tags", "links": "Links", "distributed": "Distributed", "source": "Source",
"files": "Files", "fileCount": "{{count}} file", "fileCount_other": "{{count}} files",
"adopted": "Adopted", "adoptedHint": "Discovered on a worker but not configured locally. Persist the config to make it permanent.",
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}"
"alias": "alias → {{target}}", "aliasTitle": "Alias → {{target}}",
"chain": "chain → {{target}}", "chainTitle": "Failover chain → {{target}}"
},
"open": {
"title": "Open", "chat": "Chat", "completion": "Completion", "image": "Image", "video": "Video", "tts": "TTS",
@@ -224,5 +225,67 @@
"pickVision": "读取图像与文档",
"pickAudio": "语音输入与输出",
"pickVisual": "生成图像与视频"
},
"failover": {
"title": "故障转移链",
"servedBy": "当前服务",
"changed": "{{time}}变更",
"pinnedTo": "已固定到 {{target}}",
"warm": "预热",
"never": "从未",
"states": {
"primary": "主用",
"fallback": "备用",
"degraded": "降级",
"healthy": "健康",
"down": "不可用",
"recovering": "恢复中",
"missing": "缺失"
},
"kinds": {
"local": "本地",
"remote": "远程"
},
"columns": {
"target": "目标",
"kind": "类型",
"warm": "预热",
"health": "健康状态",
"lastProbe": "上次探测",
"lastError": "上次错误",
"actions": "操作"
},
"actions": {
"pin": "固定",
"pinning": "正在固定…",
"unpin": "取消固定",
"unpinning": "正在取消固定…"
},
"confirm": {
"pinTitle": "固定 {{target}}?",
"pinMessage": "在取消固定之前,只有 {{target}} 为 {{chain}} 提供服务。此链的自动故障转移将停止。",
"unpinTitle": "取消固定 {{chain}}?",
"unpinMessage": "{{chain}} 将恢复自动故障转移。"
},
"errors": {
"pin": "固定失败:{{error}}",
"unpin": "取消固定失败:{{error}}"
},
"overview": {
"title": "故障转移链",
"subtitle": "每条故障转移链的实时健康状态和当前目标。",
"columns": {
"chain": "链",
"state": "状态",
"active": "当前目标",
"targets": "目标",
"changed": "变更时间"
},
"empty": {
"title": "尚无故障转移链",
"text": "添加故障转移链以自动绕过不健康的目标。",
"cta": "创建故障转移链"
}
}
}
}
@@ -60,7 +60,8 @@
"middleware": "Middleware",
"activity": "活动",
"overview": "概览",
"thisMachine": "本机"
"thisMachine": "本机",
"failover": "故障转移"
},
"footer": {
"github": "GitHub",
+51
View File
@@ -12275,6 +12275,15 @@ button.collapsible-header:focus-visible {
.cfr-num { width: 120px; font-size: var(--text-sm); }
.cfr-select { width: 220px; font-size: var(--text-sm); }
.cfr-radio { display: flex; align-items: center; gap: var(--spacing-sm); font-size: var(--text-sm); cursor: pointer; }
/* Failover targets editor — chain member list with a warm toggle and
duplicate/self-reference warnings. */
.fte-list { display: flex; flex-direction: column; gap: var(--spacing-sm); width: 100%; }
.fte-empty { font-size: var(--text-sm); color: var(--color-text-muted); padding: var(--spacing-sm) 0; }
.fte-warning { font-size: var(--text-xs); color: var(--color-warning); padding: var(--spacing-sm) 0; }
.fte-row { padding: var(--spacing-sm); display: flex; flex-direction: column; gap: var(--spacing-xs); }
.fte-row--error { border-color: var(--color-error); }
.fte-row__head { font-size: var(--text-xs); color: var(--color-text-muted); }
.pill-tiny { padding: 2px 6px; font-size: var(--text-xs); }
/* ===========================================================================
@@ -12361,6 +12370,48 @@ button.collapsible-header:focus-visible {
justify-content: space-between;
padding: var(--spacing-lg) var(--spacing-lg) var(--spacing-md);
}
/* Failover chain health strip (Model Editor, under the header) */
.failover-status {
padding: 0 var(--spacing-lg) var(--spacing-md);
border-bottom: 1px solid var(--color-border);
}
.failover-status__summary {
display: flex;
flex-wrap: wrap;
align-items: center;
gap: var(--spacing-xs) var(--spacing-md);
padding-bottom: var(--spacing-sm);
font-size: var(--text-sm);
color: var(--color-text-secondary);
}
.failover-status__label {
font-size: 0.75rem;
font-weight: 500;
letter-spacing: 0.04em;
text-transform: uppercase;
color: var(--color-text-muted);
}
.failover-status__active { color: var(--color-text-primary); font-family: var(--font-mono); }
.failover-status__pinned { color: var(--color-warning); }
.failover-status__unpin { margin-left: auto; }
.failover-status__table td, .failover-status__table th { white-space: nowrap; }
.failover-status__table code { font-family: var(--font-mono); font-size: 0.8125rem; }
.failover-status__row--active td:first-child { box-shadow: inset 2px 0 0 var(--color-success); }
.failover-status__error {
display: inline-block;
max-width: 28ch;
overflow: hidden;
text-overflow: ellipsis;
vertical-align: bottom;
color: var(--color-error);
}
.failover-status__action { text-align: right; }
/* Failover overview (Operate -> Runtime): one row per chain, with a compact
strip of per-target pills standing in for the detailed table the Model
Editor shows for a single chain. */
.failover-overview__targets { display: inline-flex; flex-wrap: wrap; gap: 4px; }
.failover-overview__target-pill { font-size: 0.75rem; }
.me-tabs { display: flex; gap: 0; padding: 0 var(--spacing-lg); border-bottom: 1px solid var(--color-border); }
.me-tab {
padding: var(--spacing-sm) var(--spacing-md);
@@ -11,6 +11,7 @@ import PatternListEditor from './PatternListEditor'
import ModelMultiSelect from './ModelMultiSelect'
import RouterCandidatesEditor from './RouterCandidatesEditor'
import RouterPoliciesEditor from './RouterPoliciesEditor'
import FailoverTargetsEditor from './FailoverTargetsEditor'
// Map autocomplete provider to SearchableModelSelect capability
const PROVIDER_TO_CAPABILITY = {
@@ -381,6 +382,23 @@ export default function ConfigFieldRenderer({ field, value, onChange, onRemove,
)
}
// Failover targets — ordered member list of a failover chain. Each row
// is {model, warm}; duplicate/self-reference detection reads the edited
// model's own name from FormContext.
if (component === 'failover-targets') {
return (
<div className="list-row">
<div className="hstack hstack--between mb-xs">
<div>
<div className="text-base fw-medium"><FieldLabel field={field} /></div>
<div className="text-meta mt-xs">{description}</div>
</div>
</div>
<FailoverTargetsEditor value={value} onChange={handleChange} />
</div>
)
}
// PII detectors — a capability-filtered multi-select of token_classify
// models (the consuming model's pii.detectors list).
if (component === 'model-multi-select') {
@@ -0,0 +1,164 @@
import { useState } from 'react'
import { useTranslation } from 'react-i18next'
import StatusPill from './StatusPill'
import ConfirmDialog from './ConfirmDialog'
import useFailoverChains from '../hooks/useFailoverChains'
import { useAuth } from '../context/AuthContext'
import { failoverApi } from '../utils/api'
const UNITS = [
['day', 86_400],
['hour', 3_600],
['minute', 60],
]
// Localized "5 minutes ago" from an RFC 3339 timestamp. Intl keeps the phrase
// in the viewer's language without a translation key per unit. Exported so
// other failover surfaces (the overview table) share the same phrasing.
export function relative(ts, lng) {
const ms = Date.parse(ts)
if (!ms) return null
const seconds = Math.round((ms - Date.now()) / 1000)
const rtf = new Intl.RelativeTimeFormat(lng, { numeric: 'auto' })
for (const [unit, size] of UNITS) {
if (Math.abs(seconds) >= size) return rtf.format(Math.round(seconds / size), unit)
}
return rtf.format(seconds, 'second')
}
// FailoverChainStatus renders the live health of one failover chain: its
// state, the target serving it, and a per-target health table. Pin controls
// appear only when canPin, and every pin change needs a confirmation because
// it overrides automatic failover for all callers.
export default function FailoverChainStatus({ chain, onPin, onUnpin, canPin = false }) {
const { t, i18n } = useTranslation('models')
const [confirm, setConfirm] = useState(null)
const [pending, setPending] = useState(false)
const runConfirmed = async () => {
setPending(true)
try {
if (confirm.kind === 'pin') await onPin?.(confirm.target)
else await onUnpin?.()
} finally {
setPending(false)
setConfirm(null)
}
}
const since = relative(chain.active_since, i18n.language)
return (
<section className="failover-status" aria-label={t('failover.title')}>
<div className="failover-status__summary">
<span className="failover-status__label">{t('failover.title')}</span>
<span className="failover-status__state">
<StatusPill status={chain.state} label={t(`failover.states.${chain.state}`, chain.state)} />
</span>
<span className="failover-status__meta">
{t('failover.servedBy')} <code className="failover-status__active">{chain.active}</code>
</span>
{since && (
<span className="text-muted" title={new Date(chain.active_since).toLocaleString(i18n.language)}>
{t('failover.changed', { time: since })}
</span>
)}
{chain.pinned && (
<span className="failover-status__pinned">
<i className="fas fa-thumbtack" aria-hidden="true" /> {t('failover.pinnedTo', { target: chain.pinned })}
</span>
)}
{canPin && chain.pinned && (
<button type="button" className="btn btn-ghost btn-sm failover-status__unpin" onClick={() => setConfirm({ kind: 'unpin' })}>
{t('failover.actions.unpin')}
</button>
)}
</div>
<div className="table-container">
<table className="table failover-status__table">
<thead>
<tr>
<th>{t('failover.columns.target')}</th>
<th>{t('failover.columns.kind')}</th>
<th>{t('failover.columns.warm')}</th>
<th>{t('failover.columns.health')}</th>
<th>{t('failover.columns.lastProbe')}</th>
<th>{t('failover.columns.lastError')}</th>
{canPin && <th><span className="sr-only">{t('failover.columns.actions')}</span></th>}
</tr>
</thead>
<tbody>
{(chain.targets || []).map(target => (
<tr key={target.model} className={target.model === chain.active ? 'failover-status__row--active' : undefined}>
<td><code>{target.model}</code></td>
<td>{t(`failover.kinds.${target.kind}`, target.kind)}</td>
<td>{target.warm ? t('failover.warm') : <span className="text-muted">-</span>}</td>
<td><StatusPill status={target.state} label={t(`failover.states.${target.state}`, target.state)} /></td>
<td className="text-muted">{relative(target.last_probe, i18n.language) || t('failover.never')}</td>
<td>
{target.last_error
? <span className="failover-status__error" title={target.last_error}>{target.last_error}</span>
: <span className="text-muted">-</span>}
</td>
{canPin && (
<td className="failover-status__action">
{chain.pinned !== target.model && (
<button type="button" className="btn btn-ghost btn-sm" onClick={() => setConfirm({ kind: 'pin', target: target.model })}>
{t('failover.actions.pin')}
</button>
)}
</td>
)}
</tr>
))}
</tbody>
</table>
</div>
<ConfirmDialog
open={!!confirm}
title={confirm?.kind === 'pin'
? t('failover.confirm.pinTitle', { target: confirm.target })
: t('failover.confirm.unpinTitle', { chain: chain.name })}
message={confirm?.kind === 'pin'
? t('failover.confirm.pinMessage', { chain: chain.name, target: confirm.target })
: t('failover.confirm.unpinMessage', { chain: chain.name })}
confirmLabel={confirm?.kind === 'pin' ? t('failover.actions.pin') : t('failover.actions.unpin')}
pendingLabel={confirm?.kind === 'pin' ? t('failover.actions.pinning') : t('failover.actions.unpinning')}
pending={pending}
onConfirm={runConfirmed}
onCancel={() => setConfirm(null)}
/>
</section>
)
}
// ModelFailoverStatus is the Model Editor mount point: it renders the strip
// only when the edited model is a failover chain, and owns the SSE
// subscription so create mode never opens one.
export function ModelFailoverStatus({ name, addToast }) {
const { t } = useTranslation('models')
const { isAdmin } = useAuth()
const { byName, refresh } = useFailoverChains()
const chain = byName[name]
if (!chain) return null
const act = async (call, errorKey) => {
try {
await call()
} catch (err) {
addToast?.(t(errorKey, { error: err.message }), 'error')
}
refresh()
}
return (
<FailoverChainStatus
chain={chain}
canPin={isAdmin}
onPin={(target) => act(() => failoverApi.pin(name, target), 'failover.errors.pin')}
onUnpin={() => act(() => failoverApi.unpin(name), 'failover.errors.unpin')}
/>
)
}
Loaded 100 of 186 files, more files were not shown because too many files have changed in this diff. Show more