mirror of
https://github.com/mudler/LocalAI.git
synced 2026-08-04 12:22:22 -04:00
Compare commits
4 Commits
feat/vllm-
...
bot/issue-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
127a5780f2 | ||
|
|
7916c48f04 | ||
|
|
abb6e2b9d0 | ||
|
|
41d4251cd5 |
@@ -304,9 +304,7 @@ React pages that want to filter the ModelSelector by capability import this symb
|
||||
|
||||
### 4. `docs/content/` (user-facing documentation)
|
||||
|
||||
A new capability deserves its own page under `docs/content/features/`, plus cross-links from related features. See the pattern used by `face-recognition.md` / `object-detection.md`.
|
||||
|
||||
Announcing it is the release's job, not this page's: the capability gets covered in the release blog post under `website/content/blog/`. See [preparing-a-release.md](preparing-a-release.md). `docs/content/whats-new.md` is only a pointer at the blog and GitHub Releases, so there is nothing to add there.
|
||||
A new capability deserves its own page under `docs/content/features/`, plus cross-links from related features and an entry in `docs/content/whats-new.md`. See the pattern used by `face-recognition.md` / `object-detection.md`.
|
||||
|
||||
## Path protection rules
|
||||
|
||||
@@ -336,7 +334,7 @@ When adding a new endpoint:
|
||||
- [ ] Swagger block on the handler: `@Summary`, `@Tags`, `@Param`, `@Success`, `@Router`
|
||||
- [ ] If new capability area (new swagger tag): entry in `instructionDefs` in `core/http/endpoints/localai/api_instructions.go` + test count bumped in `api_instructions_test.go`
|
||||
- [ ] If new `FLAG_*` usecase flag: matching `CAP_*` symbol exported from `core/http/react-ui/src/utils/capabilities.js`
|
||||
- [ ] `docs/content/features/<feature>.md` created; cross-links from related feature pages; capability covered in the release blog post (see [preparing-a-release.md](preparing-a-release.md))
|
||||
- [ ] `docs/content/features/<feature>.md` created; cross-links from related feature pages; entry in `docs/content/whats-new.md`
|
||||
|
||||
**Quality**
|
||||
- [ ] 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`)
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
arch=${1:?target architecture is required}
|
||||
build_type=${2-}
|
||||
|
||||
# SYCL compiles the whole tree with icpx -fsycl, and icpx never finishes
|
||||
# ggml-cpu/arch/x86/repack.cpp at -march=sapphirerapids: the job sits on that one
|
||||
# translation unit until GitHub kills it at 6h. gcc builds the same file in
|
||||
# seconds, so only the SYCL images have to give up the CPU variant matrix.
|
||||
case "$build_type" in
|
||||
sycl*)
|
||||
echo llama-cpp-fallback
|
||||
exit 0
|
||||
;;
|
||||
esac
|
||||
|
||||
# GPU arm64 base images do not consistently provide the gcc-14 toolchain needed
|
||||
# to compile ggml's armv9.2 CPU variants. Keep their portable fallback until the
|
||||
# builder images can supply that compiler.
|
||||
if [ "$arch" = "arm64" ] && [ -n "$build_type" ]; then
|
||||
echo llama-cpp-fallback
|
||||
else
|
||||
echo llama-cpp-cpu-all
|
||||
fi
|
||||
@@ -18,12 +18,10 @@ if [[ -n "${CUDA_DOCKER_ARCH:-}" ]]; then
|
||||
fi
|
||||
|
||||
cd /LocalAI/backend/cpp/llama-cpp
|
||||
BUILD_TARGET=$(/LocalAI/.docker/llama-cpp-build-target.sh "${TARGETARCH}" "${BUILD_TYPE:-}")
|
||||
if [ "$BUILD_TARGET" = "llama-cpp-cpu-all" ]; then
|
||||
# One build with ggml CPU_ALL_VARIANTS replaces the per-microarch binaries (x86:
|
||||
# avx/avx2/avx512/fallback; arm64: armv8.x/armv9.x). BUILD_TYPE remains in the
|
||||
# environment, so GPU builds retain their accelerator backend while ggml dlopens the
|
||||
# best CPU library when work is offloaded to the host.
|
||||
if [ -z "${BUILD_TYPE:-}" ]; then
|
||||
# Pure CPU image (BUILD_TYPE empty): one build with ggml CPU_ALL_VARIANTS replaces the
|
||||
# per-microarch binaries (x86: avx/avx2/avx512/fallback; arm64: armv8.x/armv9.x). ggml
|
||||
# dlopens the best libggml-cpu-*.so at runtime by probing host CPU features.
|
||||
#
|
||||
# arm64: the CPU_ALL_VARIANTS table includes armv9.2 SME variants whose -march=...+sme is
|
||||
# rejected by the Ubuntu 24.04 default gcc-13. gcc-14 accepts it, so build the arm64
|
||||
@@ -37,8 +35,14 @@ if [ "$BUILD_TARGET" = "llama-cpp-cpu-all" ]; then
|
||||
apt-get update -qq && apt-get install -y -qq gcc-14 g++-14
|
||||
export CC=gcc-14 CXX=g++-14
|
||||
fi
|
||||
make llama-cpp-cpu-all
|
||||
else
|
||||
# GPU build (cublas/hipblas/sycl/vulkan/...): the accelerator does the compute, so a
|
||||
# single fallback CPU build is enough - no per-microarch CPU variants needed. (This also
|
||||
# keeps the heavy GPU backend compile from also building the whole CPU variant matrix,
|
||||
# and avoids the gcc-14 apt step on GPU base images such as nvidia l4t.)
|
||||
make llama-cpp-fallback
|
||||
fi
|
||||
make "$BUILD_TARGET"
|
||||
make llama-cpp-grpc
|
||||
make llama-cpp-rpc-server
|
||||
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
arch=${1:?target architecture is required}
|
||||
build_type=${2-}
|
||||
|
||||
# SYCL compiles the whole tree with icpx -fsycl, and icpx never finishes
|
||||
# ggml-cpu/arch/x86/repack.cpp at -march=sapphirerapids: the job sits on that one
|
||||
# translation unit until GitHub kills it at 6h. gcc builds the same file in
|
||||
# seconds, so only the SYCL images have to give up the CPU variant matrix.
|
||||
case "$build_type" in
|
||||
sycl*)
|
||||
echo turboquant-fallback
|
||||
exit 0
|
||||
;;
|
||||
esac
|
||||
|
||||
# GPU arm64 base images do not consistently provide the gcc-14 toolchain needed
|
||||
# to compile ggml's armv9.2 CPU variants. Keep their portable fallback until the
|
||||
# builder images can supply that compiler.
|
||||
if [ "$arch" = "arm64" ] && [ -n "$build_type" ]; then
|
||||
echo turboquant-fallback
|
||||
else
|
||||
echo turboquant-cpu-all
|
||||
fi
|
||||
@@ -19,18 +19,20 @@ fi
|
||||
|
||||
cd /LocalAI/backend/cpp/turboquant
|
||||
|
||||
BUILD_TARGET=$(/LocalAI/.docker/turboquant-build-target.sh "${TARGETARCH}" "${BUILD_TYPE:-}")
|
||||
if [ "$BUILD_TARGET" = "turboquant-cpu-all" ]; then
|
||||
# BUILD_TYPE remains in the environment, so GPU builds retain their accelerator while
|
||||
# ggml selects the best CPU library when model work is offloaded to the host.
|
||||
if [ -z "${BUILD_TYPE:-}" ]; then
|
||||
# Pure CPU image: one ggml CPU_ALL_VARIANTS build replaces the per-microarch binaries.
|
||||
# arm64: the armv9.2 SME variants need gcc-14 (gcc-13 rejects +sme).
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then
|
||||
sh /LocalAI/.docker/apt-mirror.sh || true
|
||||
apt-get update -qq && apt-get install -y -qq gcc-14 g++-14
|
||||
export CC=gcc-14 CXX=g++-14
|
||||
fi
|
||||
make turboquant-cpu-all
|
||||
else
|
||||
# GPU build (cublas/hipblas/sycl/vulkan/...): single fallback CPU build, the accelerator
|
||||
# does the compute. Keeps the GPU compile from also building the CPU variant matrix and
|
||||
# avoids the gcc-14 apt step on GPU base images such as nvidia l4t.
|
||||
make turboquant-fallback
|
||||
fi
|
||||
make "$BUILD_TARGET"
|
||||
make turboquant-grpc
|
||||
make turboquant-rpc-server
|
||||
|
||||
|
||||
12
.github/workflows/bump_deps.yaml
vendored
12
.github/workflows/bump_deps.yaml
vendored
@@ -110,14 +110,10 @@ jobs:
|
||||
variable: "LOCATEANYTHING_VERSION"
|
||||
branch: "master"
|
||||
file: "backend/go/locate-anything-cpp/Makefile"
|
||||
# qwentts.cpp is held, not tracked: upstream master hangs in synthesis
|
||||
# (see the comment on QWEN3TTS_CPP_VERSION in the backend Makefile).
|
||||
# Leaving it here would re-bump the pin back onto the hang every night.
|
||||
# Restore this entry once the upstream fix lands.
|
||||
# - repository: "ServeurpersoCom/qwentts.cpp"
|
||||
# variable: "QWEN3TTS_CPP_VERSION"
|
||||
# branch: "master"
|
||||
# file: "backend/go/qwen3-tts-cpp/Makefile"
|
||||
- repository: "ServeurpersoCom/qwentts.cpp"
|
||||
variable: "QWEN3TTS_CPP_VERSION"
|
||||
branch: "master"
|
||||
file: "backend/go/qwen3-tts-cpp/Makefile"
|
||||
- repository: "ServeurpersoCom/omnivoice.cpp"
|
||||
variable: "OMNIVOICE_VERSION"
|
||||
branch: "master"
|
||||
|
||||
11
.github/workflows/gh-pages.yml
vendored
11
.github/workflows/gh-pages.yml
vendored
@@ -51,16 +51,7 @@ jobs:
|
||||
- name: Setup Go
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
# Track go.mod rather than a literal. Pinned at 1.22 this installed a
|
||||
# toolchain older than the module's `go 1.26.0`, so the `go run` below
|
||||
# downloaded the real one from proxy.golang.org on every run. That
|
||||
# fetch is not always reachable from the runner and the deploy failed
|
||||
# on five of eight consecutive master pushes with:
|
||||
# go: download go1.26.0: ... connect: network is unreachable
|
||||
# ##[error]Command failed: go env GOPATH
|
||||
# Installing the version the module asks for removes the download
|
||||
# instead of depending on it succeeding.
|
||||
go-version-file: go.mod
|
||||
go-version: '1.22'
|
||||
cache: false
|
||||
|
||||
- name: Setup Hugo
|
||||
|
||||
5
.gitignore
vendored
5
.gitignore
vendored
@@ -124,8 +124,3 @@ formal-verification/out/
|
||||
# package directory itself and untrack the source.
|
||||
/apexentries
|
||||
/.github/ci/apexentries/apexentries
|
||||
|
||||
# Runtime state written by `local-ai run` when it is started from the repo
|
||||
# root, which is what a contributor testing a build does. Nothing under here is
|
||||
# source: it is the instance's own models, outputs, traces and identity.
|
||||
/data/
|
||||
|
||||
@@ -161,7 +161,7 @@ local-ai run https://gist.githubusercontent.com/.../phi-2.yaml
|
||||
local-ai run oci://localai/phi-2:latest
|
||||
```
|
||||
|
||||
To work with a running LocalAI server from the terminal, start the built-in agent from another shell. It answers questions, reads your files and runs commands on your machine, asking you to approve anything that changes state. Inside a session, `/models` lists installed models and `/model <name>` switches between them. See the [Terminal agent](https://localai.io/docs/features/terminal-agent/) docs.
|
||||
To test a running LocalAI server from the terminal, open an interactive chat session from another shell. Inside the prompt, `/models` lists installed models and `/model <name>` switches between them.
|
||||
|
||||
```bash
|
||||
# Terminal 1
|
||||
@@ -195,7 +195,7 @@ For more details, see the [Getting Started guide](https://localai.io/basics/gett
|
||||
- **August 2025**: MLX, MLX-VLM, Diffusers, llama.cpp now supported on Apple Silicon
|
||||
- **July 2025**: All backends migrated outside the main binary — [lightweight, modular architecture](https://github.com/mudler/LocalAI/releases/tag/v3.2.0)
|
||||
|
||||
For older news and full release notes, see [GitHub Releases](https://github.com/mudler/LocalAI/releases) and the [blog](https://localai.io/blog/).
|
||||
For older news and full release notes, see [GitHub Releases](https://github.com/mudler/LocalAI/releases) and the [News page](https://localai.io/basics/news/).
|
||||
|
||||
## Features
|
||||
|
||||
@@ -260,7 +260,7 @@ We also maintain [apex-quant](https://github.com/localai-org/apex-quant), a per-
|
||||
- [Kubernetes installation](https://localai.io/basics/getting_started/#run-localai-in-kubernetes)
|
||||
- [Integrations & community projects](https://localai.io/docs/integrations/)
|
||||
- [Installation video walkthrough](https://www.youtube.com/watch?v=cMVNnlqwfw4)
|
||||
- [Blog: release write-ups, benchmarks and engineering notes](https://localai.io/blog/)
|
||||
- [Media & blog posts](https://localai.io/basics/news/#media-blogs-social)
|
||||
- [Examples](https://github.com/mudler/LocalAI-examples) — including the [realtime voice assistant demo](https://github.com/localai-org/localai-realtime-demo) (Go client for the Realtime API with tool calling)
|
||||
|
||||
## Team
|
||||
|
||||
@@ -15,7 +15,6 @@ service Backend {
|
||||
rpc PredictStream(PredictOptions) returns (stream Reply) {}
|
||||
rpc Embedding(PredictOptions) returns (EmbeddingResult) {}
|
||||
rpc GenerateImage(GenerateImageRequest) returns (Result) {}
|
||||
rpc UpscaleImage(UpscaleImageRequest) returns (Result) {}
|
||||
rpc GenerateVideo(GenerateVideoRequest) returns (Result) {}
|
||||
rpc Generate3D(Generate3DRequest) returns (Result) {}
|
||||
rpc AudioTranscription(TranscriptRequest) returns (TranscriptResult) {}
|
||||
@@ -638,12 +637,6 @@ message GenerateImageRequest {
|
||||
string ModelIdentity = 13;
|
||||
}
|
||||
|
||||
message UpscaleImageRequest {
|
||||
string src = 1; // input image path
|
||||
string dst = 2; // output image path
|
||||
int32 scale = 3; // upscale factor (e.g. 2 or 4)
|
||||
}
|
||||
|
||||
message GenerateVideoRequest {
|
||||
string prompt = 1;
|
||||
string negative_prompt = 2; // Negative prompt for video generation
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
# recipe is a make target (not a prepare.sh) so 'make purge && make' is a clean
|
||||
# rebuild and so the bump bot can see the pin.
|
||||
|
||||
AUDIO_CPP_VERSION?=4e3aea2fd99aeaa5924e71c51eb2793846045332
|
||||
AUDIO_CPP_VERSION?=f78227c52736a4792a50aa3f82ead7e7385c891b
|
||||
AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
|
||||
# Pinned to the HEAD of the `prism` branch on https://github.com/PrismML-Eng/llama.cpp.
|
||||
# Auto-bumped nightly by .github/workflows/bump_deps.yaml.
|
||||
BONSAI_VERSION?=9ca265a57f85f2117942490f421f64a226dd9847
|
||||
BONSAI_VERSION?=7529fdaaf99ffdc5ca71ace9c7409a56b27ad92f
|
||||
LLAMA_REPO?=https://github.com/PrismML-Eng/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=60389410a1ff01f9d37dcc6261db33b3183bdea2
|
||||
IK_LLAMA_VERSION?=3f53a059024039358e9fef75b5dc0c99dbcb40f9
|
||||
LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=221f0f6356efe2260023208365705ec5d5a7c8f5
|
||||
LLAMA_VERSION?=1cbfd1988311775425d36c0ce066590f7d3049cf
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
# MiniMax-M3 chat-template parser, vendored from upstream llama.cpp PR #24523.
|
||||
#
|
||||
# Upstream has since merged the *model* half of #24523 (LLM_ARCH_MINIMAX_M3,
|
||||
# src/models/minimax-m3.cpp, the gguf-py constants and conversion/minimax.py), so
|
||||
# only the chat half is carried here: M3's namespace token "]<]minimax[>[" collides
|
||||
# with the autoparser's markup delimiters, so common/chat.cpp needs a dedicated
|
||||
# template detection + PEG parser that upstream does not have yet.
|
||||
#
|
||||
# Rebased against LLAMA_VERSION 0d47ea7427463093e69128bf2c2f9cd06b3ee5b3, which also
|
||||
# renamed common_chat_params::thinking_end_tag to thinking_end_tags (a vector).
|
||||
# LLAMA_VERSION is auto-bumped nightly; if a bump rejects this patch, re-vendor from
|
||||
# #24523 — or, once the chat half merges upstream, delete this file.
|
||||
# See https://github.com/mudler/LocalAI/issues/10820 and PR #10837.
|
||||
diff --git a/common/chat.cpp b/common/chat.cpp
|
||||
index 7a6e7238c..2dd015a2e 100644
|
||||
--- a/common/chat.cpp
|
||||
+++ b/common/chat.cpp
|
||||
@@ -2121,6 +2121,191 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
return data;
|
||||
}
|
||||
|
||||
+static common_chat_params common_chat_params_init_minimax_m3(const common_chat_template & tmpl,
|
||||
+ const autoparser::generation_params & inputs) {
|
||||
+ common_chat_params data;
|
||||
+
|
||||
+ data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs);
|
||||
+ data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs);
|
||||
+ data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||
+ data.supports_thinking = true;
|
||||
+ data.thinking_start_tag = "<mm:think>";
|
||||
+ data.thinking_end_tags = {"</mm:think>"};
|
||||
+
|
||||
+ // M3 prefixes every tool tag with the namespace token "]<]minimax[>[";
|
||||
+ // params use the parameter name as the tag (<file_path>...</file_path>).
|
||||
+ const std::string NS = "]<]minimax[>[";
|
||||
+ const std::string THINK_START = "<mm:think>";
|
||||
+ const std::string THINK_END = "</mm:think>";
|
||||
+ const std::string FC_START = NS + "<tool_call>";
|
||||
+ const std::string FC_END = NS + "</tool_call>";
|
||||
+ const std::string INVOKE_END = NS + "</invoke>";
|
||||
+
|
||||
+ data.preserved_tokens = {
|
||||
+ NS,
|
||||
+ "<tool_call>",
|
||||
+ "</tool_call>",
|
||||
+ THINK_START,
|
||||
+ THINK_END,
|
||||
+ };
|
||||
+
|
||||
+ auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
||||
+ auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
|
||||
+ auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
+ auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
+
|
||||
+ const std::string GEN_PROMPT = data.generation_prompt;
|
||||
+
|
||||
+ if (inputs.has_continuation()) {
|
||||
+ const auto & msg = inputs.continue_msg;
|
||||
+
|
||||
+ data.generation_prompt = GEN_PROMPT + THINK_START + msg.reasoning_content;
|
||||
+ if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) {
|
||||
+ data.generation_prompt += THINK_END + msg.render_content();
|
||||
+ }
|
||||
+
|
||||
+ data.prompt += data.generation_prompt;
|
||||
+ }
|
||||
+
|
||||
+ auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
+ auto generation_prompt = p.literal(GEN_PROMPT);
|
||||
+ auto end = p.end();
|
||||
+
|
||||
+ auto reasoning = p.eps();
|
||||
+ // M3 can emit a bare </mm:think> (no opener) after tool results; keep the opener optional.
|
||||
+ if (extract_reasoning && inputs.enable_thinking) {
|
||||
+ reasoning = p.optional(p.optional(p.literal(THINK_START)) + p.reasoning(p.until(THINK_END)) + THINK_END);
|
||||
+ } else if (extract_reasoning) {
|
||||
+ reasoning = p.optional(p.optional(p.literal(THINK_START)) + p.until(THINK_END) + p.literal(THINK_END));
|
||||
+ }
|
||||
+
|
||||
+ if (has_response_format) {
|
||||
+ auto response_format = p.rule("response-format",
|
||||
+ p.literal("```json") + p.space() +
|
||||
+ p.content(p.schema(p.json(), "response-format-schema", inputs.json_schema)) +
|
||||
+ p.space() + p.literal("```"));
|
||||
+ return generation_prompt + reasoning + response_format + end;
|
||||
+ }
|
||||
+
|
||||
+ if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) {
|
||||
+ return generation_prompt + reasoning + p.content(p.rest()) + end;
|
||||
+ }
|
||||
+
|
||||
+ auto tool_choice = p.choice();
|
||||
+ foreach_function(inputs.tools, [&](const json & tool) {
|
||||
+ const auto & function = tool.at("function");
|
||||
+ std::string name = function.at("name");
|
||||
+ auto params = function.contains("parameters") ? function.at("parameters") : json::object();
|
||||
+ const auto & props = params.contains("properties") ? params.at("properties") : json::object();
|
||||
+
|
||||
+ std::set<std::string> required;
|
||||
+ if (params.contains("required")) {
|
||||
+ params.at("required").get_to(required);
|
||||
+ }
|
||||
+
|
||||
+ auto schema_info = common_schema_info();
|
||||
+ schema_info.resolve_refs(params);
|
||||
+
|
||||
+ std::vector<common_peg_parser> required_parsers;
|
||||
+ std::vector<common_peg_parser> optional_parsers;
|
||||
+ for (const auto & [param_name, param_schema] : props.items()) {
|
||||
+ bool is_required = required.find(param_name) != required.end();
|
||||
+ bool is_string = schema_info.resolves_to_string(param_schema);
|
||||
+
|
||||
+ const std::string p_close = NS + "</" + param_name + ">";
|
||||
+
|
||||
+ auto arg = p.tool_arg(
|
||||
+ p.tool_arg_open(
|
||||
+ p.literal(NS + "<") +
|
||||
+ p.tool_arg_name(p.literal(param_name)) +
|
||||
+ p.literal(">")) +
|
||||
+ (is_string
|
||||
+ ? p.ac(p.tool_arg_string_value(p.until(p_close)) +
|
||||
+ p.tool_arg_close(p.literal(p_close)), p_close)
|
||||
+ : p.tool_arg_json_value(p.schema(p.json(),
|
||||
+ "tool-" + name + "-arg-" + param_name + "-schema",
|
||||
+ param_schema, false)) +
|
||||
+ p.tool_arg_close(p.literal(p_close))));
|
||||
+
|
||||
+ auto named_arg = p.rule("tool-" + name + "-arg-" + param_name, arg);
|
||||
+ if (is_required) {
|
||||
+ required_parsers.push_back(named_arg);
|
||||
+ } else {
|
||||
+ optional_parsers.push_back(named_arg);
|
||||
+ }
|
||||
+ }
|
||||
+
|
||||
+ common_peg_parser args_seq = p.eps();
|
||||
+ for (size_t i = 0; i < required_parsers.size(); i++) {
|
||||
+ if (i > 0) {
|
||||
+ args_seq = args_seq + p.space();
|
||||
+ }
|
||||
+ args_seq = args_seq + required_parsers[i];
|
||||
+ }
|
||||
+
|
||||
+ if (!optional_parsers.empty()) {
|
||||
+ common_peg_parser any_opt = p.choice();
|
||||
+ for (const auto & opt : optional_parsers) {
|
||||
+ any_opt |= opt;
|
||||
+ }
|
||||
+ args_seq = args_seq + p.repeat(p.space() + any_opt, 0, -1);
|
||||
+ }
|
||||
+
|
||||
+ common_peg_parser invoke_body = args_seq;
|
||||
+ auto func_parser = p.tool(
|
||||
+ p.tool_open(p.literal(NS + "<invoke name=\"") +
|
||||
+ p.tool_name(p.literal(name)) + p.literal("\">")) +
|
||||
+ p.space() + invoke_body + p.space() +
|
||||
+ p.tool_close(p.literal(INVOKE_END)));
|
||||
+
|
||||
+ tool_choice |= p.rule("tool-" + name, func_parser);
|
||||
+ });
|
||||
+
|
||||
+ auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
+
|
||||
+ common_peg_parser tool_calls = p.eps();
|
||||
+ if (inputs.parallel_tool_calls) {
|
||||
+ tool_calls = p.trigger_rule("tool-call",
|
||||
+ p.literal(FC_START) + p.space() + tool_choice +
|
||||
+ p.zero_or_more(p.space() + tool_choice) + p.space() + p.literal(FC_END));
|
||||
+ } else {
|
||||
+ tool_calls = p.trigger_rule("tool-call",
|
||||
+ p.literal(FC_START) + p.space() + tool_choice + p.space() + p.literal(FC_END));
|
||||
+ }
|
||||
+
|
||||
+ if (!require_tools) {
|
||||
+ tool_calls = p.optional(tool_calls);
|
||||
+ }
|
||||
+
|
||||
+ auto content_before_tools = p.content(p.until(FC_START));
|
||||
+ return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
+ });
|
||||
+
|
||||
+ data.parser = parser.save();
|
||||
+
|
||||
+ if (include_grammar) {
|
||||
+ data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
+ data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
+ foreach_function(inputs.tools, [&](const json & tool) {
|
||||
+ const auto & function = tool.at("function");
|
||||
+ auto schema = function.contains("parameters") ? function.at("parameters") : json::object();
|
||||
+ builder.resolve_refs(schema);
|
||||
+ });
|
||||
+ if (has_response_format) {
|
||||
+ auto schema = inputs.json_schema;
|
||||
+ builder.resolve_refs(schema);
|
||||
+ }
|
||||
+ parser.build_grammar(builder, data.grammar_lazy);
|
||||
+ });
|
||||
+
|
||||
+ data.grammar_triggers = {
|
||||
+ { COMMON_GRAMMAR_TRIGGER_TYPE_WORD, FC_START },
|
||||
+ };
|
||||
+ }
|
||||
+
|
||||
+ return data;
|
||||
+}
|
||||
+
|
||||
// Cohere2 MoE (a.k.a. "North Code") parser.
|
||||
//
|
||||
// The assistant turn is fully marker-wrapped:
|
||||
@@ -2707,6 +2892,15 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
return common_chat_params_init_gigachat_v3(tmpl, params);
|
||||
}
|
||||
|
||||
+ // MiniMax-M3: the namespace token "]<]minimax[>[" collides with the autoparser's
|
||||
+ // markup delimiters, so detect the template and use a dedicated parser.
|
||||
+ if (src.find("]<]minimax[>[") != std::string::npos &&
|
||||
+ src.find("<tool_call>") != std::string::npos &&
|
||||
+ src.find("<invoke name=") != std::string::npos) {
|
||||
+ LOG_DBG("Using specialized template: MiniMax-M3\n");
|
||||
+ return common_chat_params_init_minimax_m3(tmpl, params);
|
||||
+ }
|
||||
+
|
||||
// DeepSeek V3.2/V4 format detection: template defines dsml_token and uses it for tool calls.
|
||||
// The template source contains the token as a variable assignment, not as a literal in markup.
|
||||
// V3.2 names the tool call block "function_calls", V4 names it "tool_calls".
|
||||
@@ -12,11 +12,10 @@ grep -e "flags" /proc/cpuinfo | head -1
|
||||
|
||||
BINARY=llama-cpp-fallback
|
||||
|
||||
# CPU images and most x86 GPU images ship a single llama-cpp-cpu-all built with ggml
|
||||
# CPU images (x86, arm64, darwin) ship a single llama-cpp-cpu-all built with ggml
|
||||
# CPU_ALL_VARIANTS: ggml's backend registry dlopens the best libggml-cpu-*.so for this
|
||||
# host, so no shell-side AVX probing. GPU arm64 images still ship llama-cpp-fallback
|
||||
# until their builder toolchains support ggml's complete arm variant matrix, and so do
|
||||
# the SYCL images, whose icpx compiler hangs on the sapphirerapids variant.
|
||||
# host, so no shell-side AVX probing. GPU images (cublas/sycl/vulkan/hipblas) ship only
|
||||
# llama-cpp-fallback (the accelerator does the compute), so fall back to it when absent.
|
||||
if [ -e "$CURDIR"/llama-cpp-cpu-all ]; then
|
||||
BINARY=llama-cpp-cpu-all
|
||||
fi
|
||||
@@ -77,4 +76,4 @@ echo "Using binary: $BINARY"
|
||||
exec "$CURDIR"/$BINARY "$@"
|
||||
|
||||
# We should never reach this point, however just in case we do, run fallback
|
||||
exec "$CURDIR"/llama-cpp-fallback "$@"
|
||||
exec "$CURDIR"/llama-cpp-fallback "$@"
|
||||
@@ -1,7 +1,7 @@
|
||||
|
||||
# Pinned to the HEAD of feature/turboquant-kv-cache on https://github.com/TheTom/llama-cpp-turboquant.
|
||||
# Auto-bumped nightly by .github/workflows/bump_deps.yaml.
|
||||
TURBOQUANT_VERSION?=8a891f4b566efdbd3cea92fafee3227a0a267683
|
||||
TURBOQUANT_VERSION?=c26cbdffcf6fc9b7430cd6b117757e9a3f70b7ea
|
||||
LLAMA_REPO?=https://github.com/TheTom/llama-cpp-turboquant
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -12,12 +12,9 @@ grep -e "flags" /proc/cpuinfo | head -1
|
||||
|
||||
BINARY=turboquant-fallback
|
||||
|
||||
# CPU images and most x86 GPU images ship a single turboquant-cpu-all built with ggml
|
||||
# CPU_ALL_VARIANTS: ggml's
|
||||
# x86/arm64 ship a single turboquant-cpu-all built with ggml CPU_ALL_VARIANTS: ggml's
|
||||
# backend registry dlopens the best libggml-cpu-*.so for this host, so no shell-side
|
||||
# probing. GPU arm64 images still ship turboquant-fallback until their builder toolchains
|
||||
# support ggml's complete arm variant matrix, and so do the SYCL images, whose icpx
|
||||
# compiler hangs on the sapphirerapids variant.
|
||||
# probing. ROCm ships only turboquant-fallback, so fall back to it when cpu-all is absent.
|
||||
if [ -e "$CURDIR"/turboquant-cpu-all ]; then
|
||||
BINARY=turboquant-cpu-all
|
||||
fi
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# CrispASR version (release tag)
|
||||
CRISPASR_REPO?=https://github.com/CrispStrobe/CrispASR
|
||||
CRISPASR_VERSION?=fe3caf8e363b27572dbdd1a9d37083f25e6decda
|
||||
CRISPASR_VERSION?=b5211ac635489049ee8ce86a82d69faa18e8d8da
|
||||
SO_TARGET?=libgocrispasr.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -7,18 +7,8 @@ GO_TAGS?=
|
||||
JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# qwentts.cpp version
|
||||
#
|
||||
# Held at 35ebe537 rather than tracking latest: abab6b3 hangs in synthesis.
|
||||
# TTS() never returns from the native call, so tests-qwen3-tts-cpp goes from
|
||||
# ~5 minutes to the 20 minute Go test timeout. Reproduced on master on
|
||||
# 2026-08-01 and again on re-run, and the bump PR (#11241) was merged with
|
||||
# this same check already red.
|
||||
#
|
||||
# The regression is in 35ebe537..abab6b3, three upstream commits whose only
|
||||
# functional change is 26dd8adb, "predictor: unroll the frame into one cgraph
|
||||
# and sample in standard ops". Restore the bump once that is fixed upstream.
|
||||
QWEN3TTS_REPO?=https://github.com/ServeurpersoCom/qwentts.cpp
|
||||
QWEN3TTS_CPP_VERSION?=35ebe5376b82a0a59d008586d55bbe623d449011
|
||||
QWEN3TTS_CPP_VERSION?=abab6b3bf317cfa1b788efce1d25f4f9239395ad
|
||||
SO_TARGET?=libgoqwen3ttscpp.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# stablediffusion.cpp (ggml)
|
||||
STABLEDIFFUSION_GGML_REPO?=https://github.com/leejet/stable-diffusion.cpp
|
||||
STABLEDIFFUSION_GGML_VERSION?=db99efdd6d2a43c7937fd55b3359206c680a75b0
|
||||
STABLEDIFFUSION_GGML_VERSION?=e31a86ce9110b11a98bd5990c329093244c2d1e3
|
||||
|
||||
CMAKE_ARGS+=-DGGML_MAX_NAME=128
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ JOBS?=$(shell nproc --ignore=1 2>/dev/null || sysctl -n hw.ncpu 2>/dev/null || e
|
||||
|
||||
# vllm.cpp version
|
||||
VLLM_CPP_REPO?=https://github.com/mudler/vllm.cpp
|
||||
VLLM_CPP_VERSION?=a42b8187caff02c570c28e19e4dc2b1d7f55ed14
|
||||
VLLM_CPP_VERSION?=9e1c9025ae61167a3335454d7cc0de6093c21845
|
||||
|
||||
# The backend consumes only the stable C ABI (libvllm + include/vllm.h), so the
|
||||
# server, examples and tests of the engine are never built here.
|
||||
@@ -56,12 +56,6 @@ endif
|
||||
UNAME_S := $(shell uname -s)
|
||||
ifeq ($(UNAME_S),Darwin)
|
||||
LIB=libvllm.dylib
|
||||
# Apple Clang diagnoses a pair of constant-folded array bounds in the Metal
|
||||
# build as a GNU extension. Disable that diagnostic for both Objective-C and
|
||||
# C++ because vllm.cpp appends target-local -Werror after these global flags.
|
||||
CMAKE_ARGS+=-DCMAKE_CXX_FLAGS=-Wno-gnu-folding-constant
|
||||
CMAKE_ARGS+=-DCMAKE_OBJC_FLAGS=-Wno-gnu-folding-constant
|
||||
CMAKE_ARGS+=-DCMAKE_OBJCXX_FLAGS=-Wno-gnu-folding-constant
|
||||
else
|
||||
LIB=libvllm.so
|
||||
endif
|
||||
|
||||
@@ -109,16 +109,6 @@ func (v *VllmCpp) Load(opts *pb.ModelOptions) error {
|
||||
|
||||
v.opts = parseOptions(opts)
|
||||
|
||||
// A DFlash draft is a second checkpoint the engine opens by path, and the
|
||||
// engine never downloads one. Resolve it against LocalAI's models directory
|
||||
// now so a repo-id spelling works, and so a missing draft fails here with an
|
||||
// actionable message rather than as an HF-cache miss inside the load.
|
||||
resolvedSpec, err := resolveDraftModelPath(v.opts.speculativeConfig, opts.ModelPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
v.opts.speculativeConfig = resolvedSpec
|
||||
|
||||
mp := defaultModelParams()
|
||||
if v.opts.blockSize > 0 {
|
||||
mp.BlockSize = v.opts.blockSize
|
||||
@@ -126,62 +116,34 @@ func (v *VllmCpp) Load(opts *pb.ModelOptions) error {
|
||||
if v.opts.numBlocks > 0 {
|
||||
mp.NumBlocks = v.opts.numBlocks
|
||||
}
|
||||
// Sequence-length precedence, narrowest source last: context_size is the
|
||||
// generic LocalAI knob every backend honours, max_model_len is the
|
||||
// vLLM-specific one, and engine_args.max_model_len is the explicit
|
||||
// vllm-cpp override.
|
||||
if opts.ContextSize > 0 {
|
||||
mp.MaxModelLen = opts.ContextSize
|
||||
}
|
||||
if opts.MaxModelLen > 0 {
|
||||
mp.MaxModelLen = opts.MaxModelLen
|
||||
}
|
||||
if v.opts.maxModelLen > 0 {
|
||||
mp.MaxModelLen = v.opts.maxModelLen
|
||||
}
|
||||
if v.opts.maxNumSeqs > 0 {
|
||||
mp.MaxNumSeqs = v.opts.maxNumSeqs
|
||||
}
|
||||
if v.opts.maxNumBatchedTokens > 0 {
|
||||
mp.MaxNumBatchedTokens = v.opts.maxNumBatchedTokens
|
||||
}
|
||||
mp.EnablePrefixCaching = v.opts.enablePrefixCaching
|
||||
mp.EnableJumpForward = v.opts.enableJumpForward
|
||||
|
||||
// Every string below is borrowed by C for the duration of the load call
|
||||
// only (the library copies what it keeps), so the backing slices just have
|
||||
// to outlive vllmEngineLoad - hence the single KeepAlive after it.
|
||||
modelC := cString(model)
|
||||
mp.ModelPath = uintptr(unsafe.Pointer(&modelC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
keep := [][]byte{modelC}
|
||||
setStr := func(dst *uintptr, s string) {
|
||||
if s == "" {
|
||||
return
|
||||
}
|
||||
b := cString(s)
|
||||
keep = append(keep, b)
|
||||
*dst = uintptr(unsafe.Pointer(&b[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
var toolParserC, reasoningParserC []byte
|
||||
if v.opts.toolParser != "" {
|
||||
toolParserC = cString(v.opts.toolParser)
|
||||
mp.ToolParser = uintptr(unsafe.Pointer(&toolParserC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
}
|
||||
if v.opts.reasoningParser != "" {
|
||||
reasoningParserC = cString(v.opts.reasoningParser)
|
||||
mp.ReasoningParser = uintptr(unsafe.Pointer(&reasoningParserC[0])) // #nosec G103 -- borrowed by C for the load call only
|
||||
}
|
||||
setStr(&mp.ToolParser, v.opts.toolParser)
|
||||
setStr(&mp.ReasoningParser, v.opts.reasoningParser)
|
||||
setStr(&mp.SpeculativeConfig, v.opts.speculativeConfig)
|
||||
setStr(&mp.KVTransferConfig, v.opts.kvTransferConfig)
|
||||
setStr(&mp.SchedulingPolicy, v.opts.schedulingPolicy)
|
||||
setStr(&mp.TokenizerConfigPath, v.opts.tokenizerConfigPath)
|
||||
|
||||
xlog.Info("[vllm-cpp] Load", "model", model, "engine", vllmVersion(),
|
||||
"blockSize", mp.BlockSize, "numBlocks", mp.NumBlocks,
|
||||
"maxModelLen", mp.MaxModelLen, "maxNumSeqs", mp.MaxNumSeqs,
|
||||
"maxNumBatchedTokens", mp.MaxNumBatchedTokens,
|
||||
"prefixCaching", triStateName(mp.EnablePrefixCaching),
|
||||
"jumpForward", triStateName(mp.EnableJumpForward),
|
||||
"schedulingPolicy", v.opts.schedulingPolicy,
|
||||
"speculativeConfig", v.opts.speculativeConfig,
|
||||
"kvTransferConfig", v.opts.kvTransferConfig)
|
||||
"maxModelLen", mp.MaxModelLen, "maxNumSeqs", mp.MaxNumSeqs)
|
||||
|
||||
var engine uintptr
|
||||
rc := vllmEngineLoad(unsafe.Pointer(&mp), unsafe.Pointer(&engine)) // #nosec G103 -- POD out-params
|
||||
runtime.KeepAlive(keep)
|
||||
runtime.KeepAlive(modelC)
|
||||
runtime.KeepAlive(toolParserC)
|
||||
runtime.KeepAlive(reasoningParserC)
|
||||
if rc != vllmOK {
|
||||
return fmt.Errorf("vllm-cpp: engine load failed: %s", vllmLastError())
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package main
|
||||
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v10).
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v2).
|
||||
//
|
||||
// The structs below are hand-mirrored PODs of the C declarations, with
|
||||
// explicit padding so the Go layout matches the C layout on linux/darwin
|
||||
@@ -18,56 +18,23 @@ import (
|
||||
)
|
||||
|
||||
// abiVersion is the VLLM_ABI_VERSION this file mirrors (vllm.h).
|
||||
const abiVersion = 10
|
||||
|
||||
// The ABI's tri-state toggles (enable_prefix_caching ABI v7,
|
||||
// enable_jump_forward ABI v10) share one encoding: 0 is NOT "off", it is
|
||||
// "defer" - to the model capability for prefix caching, to the environment for
|
||||
// jump forward. Only 2 is an explicit off.
|
||||
const (
|
||||
triStateDefer int32 = 0
|
||||
triStateOn int32 = 1
|
||||
triStateOff int32 = 2
|
||||
)
|
||||
|
||||
// triStateName renders a tri-state for the load log line, where "0" would
|
||||
// otherwise read as "off" rather than "whatever the default resolves to".
|
||||
func triStateName(state int32) string {
|
||||
switch state {
|
||||
case triStateOn:
|
||||
return "on"
|
||||
case triStateOff:
|
||||
return "off"
|
||||
default:
|
||||
return "model-default"
|
||||
}
|
||||
}
|
||||
const abiVersion = 5
|
||||
|
||||
// vllm_status (vllm.h).
|
||||
const (
|
||||
vllmOK = 0
|
||||
)
|
||||
|
||||
// cModelParams mirrors vllm_model_params. The int32 fields sit in pairs so the
|
||||
// interior needs no padding on LP64, but the struct is 8-aligned (it holds
|
||||
// pointers) and ends on a lone int32, so the trailing pad is explicit. Offsets
|
||||
// and total size are asserted in vllmcpp_test.go.
|
||||
// cModelParams mirrors vllm_model_params.
|
||||
type cModelParams struct {
|
||||
ModelPath uintptr // const char*
|
||||
TokenizerConfigPath uintptr // const char*; NULL = <model_dir>/... (ABI v9)
|
||||
TokenizerConfigPath uintptr // const char*
|
||||
BlockSize int32
|
||||
NumBlocks int32
|
||||
MaxModelLen int32
|
||||
MaxNumSeqs int32
|
||||
ToolParser uintptr // const char*; NULL = auto-detect (ABI v4)
|
||||
ReasoningParser uintptr // const char*; NULL = auto-detect (ABI v5)
|
||||
SpeculativeConfig uintptr // const char* JSON; NULL = no speculation (ABI v6)
|
||||
EnablePrefixCaching int32 // tri-state 0/1/2 (ABI v7)
|
||||
MaxNumBatchedTokens int32 // <= 0 = per-arch default (ABI v9)
|
||||
SchedulingPolicy uintptr // const char*; NULL = "fcfs" (ABI v9)
|
||||
KVTransferConfig uintptr // const char* JSON; NULL = no connector (ABI v9)
|
||||
EnableJumpForward int32 // tri-state 0/1/2 (ABI v10)
|
||||
_ [4]byte // trailing pad to the struct's 8-byte alignment
|
||||
}
|
||||
|
||||
// cSamplingParams mirrors vllm_sampling_params (ABI v2, structured fields
|
||||
@@ -98,12 +65,6 @@ type cSamplingParams struct {
|
||||
StructuredGrammar uintptr // const char*
|
||||
StructuredJSONObject int32
|
||||
_ [4]byte
|
||||
// ABI v8 tail. LocalAI installs no custom logits processor, but the fields
|
||||
// MUST be mirrored: the C side reads them off the pointer we hand it, so a
|
||||
// Go struct that stopped at StructuredJSONObject would have the engine read
|
||||
// 16 bytes past our allocation and call whatever garbage sat there.
|
||||
LogitsProcessor uintptr // vllm_logits_processor; NULL = none
|
||||
LogitsProcessorUserData uintptr // void*
|
||||
}
|
||||
|
||||
// cCompletion mirrors vllm_completion.
|
||||
|
||||
@@ -1,80 +1,30 @@
|
||||
package main
|
||||
|
||||
// Load-time engine configuration, from two config surfaces:
|
||||
//
|
||||
// - `engine_args:` (ModelOptions.EngineArgs, a JSON object) is the canonical
|
||||
// one. Keys are spelled exactly as vLLM's own CLI flags, so a config written
|
||||
// against vLLM works verbatim here - `speculative_config` and
|
||||
// `kv_transfer_config` in particular take the same JSON documents vLLM's
|
||||
// --speculative-config / --kv-transfer-config accept, and are handed to the
|
||||
// engine unparsed.
|
||||
// - `options:` (the free-form "key:value" list) is the older surface this
|
||||
// backend shipped with. It is still honoured so existing configs keep
|
||||
// working; engine_args wins on any key set in both.
|
||||
//
|
||||
// Anything unrecognised is ignored rather than fatal: the engine validates the
|
||||
// documents it is given and reports a precise error at load, and a config that
|
||||
// also carries knobs for a different backend must not fail the load here.
|
||||
// Engine-sizing knobs carried through the model config's free-form
|
||||
// `options:` list ("key:value" entries), mirroring how the other in-house
|
||||
// backends pass engine-specific settings that have no proto field.
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
type loadOptions struct {
|
||||
blockSize int32 // KV block size (tokens/block); engine default 32.
|
||||
numBlocks int32 // KV blocks to allocate; engine default 256.
|
||||
maxNumSeqs int32 // max concurrent sequences; engine default 8.
|
||||
// Max sequence length. Also settable through the model config's
|
||||
// context_size / max_model_len; see Load for the precedence.
|
||||
maxModelLen int32
|
||||
// Per-step chunked-prefill token budget (ABI v9). 0 = the engine's
|
||||
// bounded per-arch default.
|
||||
maxNumBatchedTokens int32
|
||||
// Automatic prefix caching tri-state (ABI v7): 0 = the model-capability
|
||||
// default, 1 = force on, 2 = force off.
|
||||
enablePrefixCaching int32
|
||||
// Jump-forward decoding tri-state (ABI v10), SGLang's grammar-speed subset:
|
||||
// 0 = defer to the environment (VT_ENABLE_JUMP_FORWARD, default off),
|
||||
// 1 = force on, 2 = force off.
|
||||
enableJumpForward int32
|
||||
// Scheduler admission policy (ABI v9): "" = fcfs, else fcfs|priority|lpm.
|
||||
schedulingPolicy string
|
||||
// Engine-side parser selection (ABI v4/v5). Empty = the engine
|
||||
// auto-detects from the chat template; "none" disables the reasoning
|
||||
// split; unknown names fail the first chat call.
|
||||
toolParser string
|
||||
reasoningParser string
|
||||
// Speculative decoding (ABI v6), as vLLM's --speculative-config JSON:
|
||||
// {"method":"mtp"|"dflash"|"ngram", ...}. Empty = no speculation.
|
||||
speculativeConfig string
|
||||
// External KV connector / LMCache (ABI v9), as vLLM's --kv-transfer-config
|
||||
// JSON. Empty = no connector.
|
||||
kvTransferConfig string
|
||||
// Override for the tokenizer_config.json the chat template is read from
|
||||
// (ABI v9). Empty = <model_dir>/tokenizer_config.json.
|
||||
tokenizerConfigPath string
|
||||
}
|
||||
|
||||
func parseOptions(opts *pb.ModelOptions) loadOptions {
|
||||
lo := loadOptions{}
|
||||
applyOptionsList(&lo, opts.GetOptions())
|
||||
applyEngineArgs(&lo, opts.GetEngineArgs())
|
||||
return lo
|
||||
}
|
||||
|
||||
// applyOptionsList reads the legacy free-form "key:value" list. strings.Cut
|
||||
// splits on the FIRST colon only, so a JSON object value survives intact.
|
||||
func applyOptionsList(lo *loadOptions, options []string) {
|
||||
for _, o := range options {
|
||||
for _, o := range opts.GetOptions() {
|
||||
k, v, found := strings.Cut(o, ":")
|
||||
if !found {
|
||||
continue
|
||||
@@ -86,211 +36,13 @@ func applyOptionsList(lo *loadOptions, options []string) {
|
||||
lo.numBlocks = parseInt32(v, lo.numBlocks)
|
||||
case "max_num_seqs":
|
||||
lo.maxNumSeqs = parseInt32(v, lo.maxNumSeqs)
|
||||
case "max_num_batched_tokens":
|
||||
lo.maxNumBatchedTokens = parseInt32(v, lo.maxNumBatchedTokens)
|
||||
case "max_model_len":
|
||||
lo.maxModelLen = parseInt32(v, lo.maxModelLen)
|
||||
case "scheduling_policy", "schedule_policy":
|
||||
lo.schedulingPolicy = strings.TrimSpace(v)
|
||||
case "tool_parser", "tool_call_parser":
|
||||
case "tool_parser":
|
||||
lo.toolParser = strings.TrimSpace(v)
|
||||
case "reasoning_parser":
|
||||
lo.reasoningParser = strings.TrimSpace(v)
|
||||
case "speculative_config":
|
||||
lo.speculativeConfig = strings.TrimSpace(v)
|
||||
case "kv_transfer_config":
|
||||
lo.kvTransferConfig = strings.TrimSpace(v)
|
||||
case "tokenizer_config", "tokenizer_config_path":
|
||||
lo.tokenizerConfigPath = strings.TrimSpace(v)
|
||||
case "enable_prefix_caching", "enable_radix_attention":
|
||||
if b, err := strconv.ParseBool(strings.TrimSpace(v)); err == nil {
|
||||
lo.enablePrefixCaching = boolTriState(b)
|
||||
}
|
||||
case "enable_jump_forward":
|
||||
if b, err := strconv.ParseBool(strings.TrimSpace(v)); err == nil {
|
||||
lo.enableJumpForward = boolTriState(b)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// applyEngineArgs overlays the `engine_args:` JSON object. A document that does
|
||||
// not parse is logged and skipped: engine_args is shared with the other engines
|
||||
// (the vLLM and SGLang backends read the same field), so a stray key must not
|
||||
// take the model down.
|
||||
func applyEngineArgs(lo *loadOptions, engineArgs string) {
|
||||
if strings.TrimSpace(engineArgs) == "" {
|
||||
return
|
||||
}
|
||||
var args map[string]any
|
||||
if err := json.Unmarshal([]byte(engineArgs), &args); err != nil {
|
||||
xlog.Warn("[vllm-cpp] ignoring unparseable engine_args", "error", err)
|
||||
return
|
||||
}
|
||||
for k, v := range args {
|
||||
switch k {
|
||||
case "block_size":
|
||||
lo.blockSize = jsonInt32(v, lo.blockSize)
|
||||
case "num_blocks":
|
||||
lo.numBlocks = jsonInt32(v, lo.numBlocks)
|
||||
case "max_num_seqs":
|
||||
lo.maxNumSeqs = jsonInt32(v, lo.maxNumSeqs)
|
||||
case "max_num_batched_tokens":
|
||||
lo.maxNumBatchedTokens = jsonInt32(v, lo.maxNumBatchedTokens)
|
||||
case "max_model_len":
|
||||
lo.maxModelLen = jsonInt32(v, lo.maxModelLen)
|
||||
case "scheduling_policy", "schedule_policy":
|
||||
lo.schedulingPolicy = jsonString(v, lo.schedulingPolicy)
|
||||
case "tool_parser", "tool_call_parser":
|
||||
lo.toolParser = jsonString(v, lo.toolParser)
|
||||
case "reasoning_parser":
|
||||
lo.reasoningParser = jsonString(v, lo.reasoningParser)
|
||||
case "tokenizer_config", "tokenizer_config_path":
|
||||
lo.tokenizerConfigPath = jsonString(v, lo.tokenizerConfigPath)
|
||||
case "speculative_config":
|
||||
lo.speculativeConfig = jsonDocument(v, lo.speculativeConfig, k)
|
||||
case "kv_transfer_config":
|
||||
lo.kvTransferConfig = jsonDocument(v, lo.kvTransferConfig, k)
|
||||
case "enable_prefix_caching", "enable_radix_attention":
|
||||
if b, ok := v.(bool); ok {
|
||||
lo.enablePrefixCaching = boolTriState(b)
|
||||
}
|
||||
case "enable_jump_forward":
|
||||
if b, ok := v.(bool); ok {
|
||||
lo.enableJumpForward = boolTriState(b)
|
||||
}
|
||||
default:
|
||||
xlog.Debug("[vllm-cpp] ignoring unknown engine_args key", "key", k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// boolTriState maps a YAML/JSON boolean onto the ABI's tri-state encoding. An
|
||||
// explicit `false` must reach the engine as force-OFF (2), NOT as the 0 that
|
||||
// means "defer". The difference is real in both directions: prefix caching
|
||||
// defaults ON for dense archs and OFF for hybrid ones, and jump forward defers
|
||||
// to VT_ENABLE_JUMP_FORWARD.
|
||||
func boolTriState(on bool) int32 {
|
||||
if on {
|
||||
return triStateOn
|
||||
}
|
||||
return triStateOff
|
||||
}
|
||||
|
||||
// jsonDocument normalises an object-valued engine_args entry to a JSON string
|
||||
// for the C ABI. YAML nesting arrives as a map (the natural spelling); a
|
||||
// pre-encoded JSON string is accepted too, since a config round-tripped through
|
||||
// a flat store may carry it that way.
|
||||
func jsonDocument(v any, fallback string, key string) string {
|
||||
switch t := v.(type) {
|
||||
case string:
|
||||
if strings.TrimSpace(t) == "" {
|
||||
return fallback
|
||||
}
|
||||
return t
|
||||
default:
|
||||
buf, err := json.Marshal(t)
|
||||
if err != nil {
|
||||
xlog.Warn("[vllm-cpp] ignoring unencodable engine_args value", "key", key, "error", err)
|
||||
return fallback
|
||||
}
|
||||
return string(buf)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonString(v any, fallback string) string {
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return fallback
|
||||
}
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
|
||||
// jsonInt32 accepts the float64 a JSON number decodes to, plus the string
|
||||
// spelling a YAML config may produce. Non-positive values keep the fallback:
|
||||
// every knob this covers uses "<= 0 means the engine default".
|
||||
func jsonInt32(v any, fallback int32) int32 {
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
if t <= 0 || t > 1<<31-1 {
|
||||
return fallback
|
||||
}
|
||||
return int32(t)
|
||||
case string:
|
||||
return parseInt32(t, fallback)
|
||||
default:
|
||||
return fallback
|
||||
}
|
||||
}
|
||||
|
||||
// resolveDraftModelPath rewrites a DFlash draft reference into an absolute path
|
||||
// the engine can actually open.
|
||||
//
|
||||
// The engine resolves `speculative_config.model` against a directory containing
|
||||
// config.json, or against ~/.cache/huggingface/hub/models--<org>--<repo>/
|
||||
// snapshots/* - and it NEVER downloads. LocalAI keeps models in its own
|
||||
// directory, so a bare HF repo id (the spelling the vLLM docs teach) misses the
|
||||
// HF cache and dies deep in the load with "draft checkpoint not found", which
|
||||
// reads like a broken checkpoint rather than a missing download.
|
||||
//
|
||||
// So: try the reference as given, then the last path segment under the models
|
||||
// dir (`z-lab/Qwen3.6-27B-DFlash` -> `<models>/Qwen3.6-27B-DFlash`, which is
|
||||
// what LocalAI's own downloader produces), then the whole reference under the
|
||||
// models dir. If none exist, fail HERE with a message naming both what was
|
||||
// asked for and where we looked.
|
||||
//
|
||||
// mtp and ngram carry no separate draft checkpoint, so they pass through. A
|
||||
// document that does not parse also passes through: the engine owns config
|
||||
// validation and produces the better error.
|
||||
func resolveDraftModelPath(speculativeConfig, modelsDir string) (string, error) {
|
||||
if strings.TrimSpace(speculativeConfig) == "" {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
var spec map[string]any
|
||||
if err := json.Unmarshal([]byte(speculativeConfig), &spec); err != nil {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
if method, _ := spec["method"].(string); !strings.EqualFold(method, "dflash") {
|
||||
return speculativeConfig, nil
|
||||
}
|
||||
|
||||
ref, _ := spec["model"].(string)
|
||||
ref = strings.TrimSpace(ref)
|
||||
if ref == "" {
|
||||
return "", fmt.Errorf(
|
||||
"vllm-cpp: speculative_config method %q requires a \"model\" key naming the draft checkpoint", "dflash")
|
||||
}
|
||||
|
||||
candidates := []string{ref}
|
||||
if modelsDir != "" {
|
||||
if base := path.Base(filepath.ToSlash(ref)); base != "" && base != "." && base != "/" {
|
||||
candidates = append(candidates, filepath.Join(modelsDir, base))
|
||||
}
|
||||
candidates = append(candidates, filepath.Join(modelsDir, filepath.FromSlash(ref)))
|
||||
}
|
||||
|
||||
for _, c := range candidates {
|
||||
if _, err := os.Stat(filepath.Join(c, "config.json")); err != nil {
|
||||
continue
|
||||
}
|
||||
abs, err := filepath.Abs(c)
|
||||
if err != nil {
|
||||
abs = c
|
||||
}
|
||||
spec["model"] = abs
|
||||
out, err := json.Marshal(spec)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("vllm-cpp: re-encoding speculative_config: %w", err)
|
||||
}
|
||||
xlog.Info("[vllm-cpp] resolved DFlash draft checkpoint", "reference", ref, "path", abs)
|
||||
return string(out), nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf(
|
||||
"vllm-cpp: DFlash draft checkpoint %q not found (looked in: %s). "+
|
||||
"The engine does not download drafts - install the draft model into LocalAI first, "+
|
||||
"or set speculative_config.model to an absolute path to a directory containing config.json",
|
||||
ref, strings.Join(candidates, ", "))
|
||||
return lo
|
||||
}
|
||||
|
||||
func parseInt32(s string, fallback int32) int32 {
|
||||
|
||||
@@ -16,7 +16,7 @@ func TestVllmCpp(t *testing.T) {
|
||||
RunSpecs(t, "vllm-cpp suite")
|
||||
}
|
||||
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v9)
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v2)
|
||||
// byte-for-byte: these offsets are the C offsets on LP64 (linux/darwin
|
||||
// amd64+arm64). A failure here means govllmcpp.go drifted from vllm.h.
|
||||
var _ = Describe("C ABI struct mirrors", func() {
|
||||
@@ -30,18 +30,10 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
Expect(unsafe.Offsetof(p.MaxNumSeqs)).To(Equal(uintptr(28)))
|
||||
Expect(unsafe.Offsetof(p.ToolParser)).To(Equal(uintptr(32)))
|
||||
Expect(unsafe.Offsetof(p.ReasoningParser)).To(Equal(uintptr(40)))
|
||||
Expect(unsafe.Offsetof(p.SpeculativeConfig)).To(Equal(uintptr(48)))
|
||||
Expect(unsafe.Offsetof(p.EnablePrefixCaching)).To(Equal(uintptr(56)))
|
||||
Expect(unsafe.Offsetof(p.MaxNumBatchedTokens)).To(Equal(uintptr(60)))
|
||||
Expect(unsafe.Offsetof(p.SchedulingPolicy)).To(Equal(uintptr(64)))
|
||||
Expect(unsafe.Offsetof(p.KVTransferConfig)).To(Equal(uintptr(72)))
|
||||
Expect(unsafe.Offsetof(p.EnableJumpForward)).To(Equal(uintptr(80)))
|
||||
// 88, not 84: the struct is 8-aligned (it holds pointers), so the
|
||||
// trailing int32 is padded out. Go pads identically.
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(88)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(48)))
|
||||
})
|
||||
|
||||
It("cSamplingParams matches vllm_sampling_params (ABI v8)", func() {
|
||||
It("cSamplingParams matches vllm_sampling_params (ABI v2)", func() {
|
||||
var p cSamplingParams
|
||||
Expect(unsafe.Offsetof(p.Temperature)).To(Equal(uintptr(0)))
|
||||
Expect(unsafe.Offsetof(p.TopP)).To(Equal(uintptr(4)))
|
||||
@@ -63,9 +55,7 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
Expect(unsafe.Offsetof(p.NStructuredChoice)).To(Equal(uintptr(96)))
|
||||
Expect(unsafe.Offsetof(p.StructuredGrammar)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.StructuredJSONObject)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Offsetof(p.LogitsProcessor)).To(Equal(uintptr(120)))
|
||||
Expect(unsafe.Offsetof(p.LogitsProcessorUserData)).To(Equal(uintptr(128)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(136)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(120)))
|
||||
})
|
||||
|
||||
It("cCompletion matches vllm_completion", func() {
|
||||
@@ -78,23 +68,6 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// Pin/mirror skew is the failure mode this backend is most exposed to: the Go
|
||||
// PODs above are hand-written against one VLLM_ABI_VERSION, and the Makefile
|
||||
// pins the vllm.cpp commit that produces it. This spec catches drift without
|
||||
// needing model weights - set VLLM_CPP_LIBRARY to a built libvllm and it binds
|
||||
// every symbol and compares the library's reported ABI against the mirrors'.
|
||||
var _ = Describe("real library ABI handshake", func() {
|
||||
It("binds every symbol and reports the ABI the mirrors were written against", func() {
|
||||
lib := os.Getenv("VLLM_CPP_LIBRARY")
|
||||
if lib == "" {
|
||||
Skip("VLLM_CPP_LIBRARY not set; skipping the real-library handshake")
|
||||
}
|
||||
Expect(registerLib(lib)).To(Succeed())
|
||||
Expect(vllmABIVersion()).To(Equal(int32(abiVersion)))
|
||||
Expect(vllmVersion()).NotTo(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("parseOptions", func() {
|
||||
It("extracts the engine sizing knobs", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{
|
||||
@@ -110,129 +83,6 @@ var _ = Describe("parseOptions", func() {
|
||||
}})
|
||||
Expect(lo).To(Equal(loadOptions{}))
|
||||
})
|
||||
|
||||
It("carries a speculative_config JSON value through the legacy options list", func() {
|
||||
// strings.Cut splits on the FIRST colon only, so a JSON object value
|
||||
// survives the "key:value" spelling intact.
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{
|
||||
`speculative_config:{"method":"mtp","num_speculative_tokens":1}`,
|
||||
}})
|
||||
Expect(lo.speculativeConfig).To(Equal(`{"method":"mtp","num_speculative_tokens":1}`))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("engine_args", func() {
|
||||
It("maps every load knob onto the C model params", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"block_size": 64,
|
||||
"num_blocks": 1024,
|
||||
"max_model_len": 16384,
|
||||
"max_num_seqs": 32,
|
||||
"max_num_batched_tokens": 8192,
|
||||
"enable_prefix_caching": true,
|
||||
"scheduling_policy": "lpm",
|
||||
"tool_parser": "qwen3",
|
||||
"reasoning_parser": "deepseek_r1",
|
||||
"tokenizer_config": "/models/tok/tokenizer_config.json"
|
||||
}`})
|
||||
Expect(lo.blockSize).To(Equal(int32(64)))
|
||||
Expect(lo.numBlocks).To(Equal(int32(1024)))
|
||||
Expect(lo.maxModelLen).To(Equal(int32(16384)))
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(32)))
|
||||
Expect(lo.maxNumBatchedTokens).To(Equal(int32(8192)))
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(1)))
|
||||
Expect(lo.schedulingPolicy).To(Equal("lpm"))
|
||||
Expect(lo.toolParser).To(Equal("qwen3"))
|
||||
Expect(lo.reasoningParser).To(Equal("deepseek_r1"))
|
||||
Expect(lo.tokenizerConfigPath).To(Equal("/models/tok/tokenizer_config.json"))
|
||||
})
|
||||
|
||||
It("re-marshals a nested speculative_config object to JSON for the engine", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"speculative_config": {"method": "mtp", "num_speculative_tokens": 1}
|
||||
}`})
|
||||
Expect(lo.speculativeConfig).To(MatchJSON(`{"method":"mtp","num_speculative_tokens":1}`))
|
||||
})
|
||||
|
||||
It("re-marshals a nested kv_transfer_config object (LMCache) to JSON", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"kv_transfer_config": {
|
||||
"kv_connector": "LMCacheConnector",
|
||||
"kv_role": "kv_both",
|
||||
"kv_connector_extra_config": {"host": "127.0.0.1", "port": 65432}
|
||||
}
|
||||
}`})
|
||||
Expect(lo.kvTransferConfig).To(MatchJSON(`{
|
||||
"kv_connector":"LMCacheConnector",
|
||||
"kv_role":"kv_both",
|
||||
"kv_connector_extra_config":{"host":"127.0.0.1","port":65432}
|
||||
}`))
|
||||
})
|
||||
|
||||
It("accepts a pre-encoded JSON string for the object-valued knobs", func() {
|
||||
// A config written by hand (or round-tripped through a flat store) may
|
||||
// carry the object as a string; both spellings reach the engine the same.
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{
|
||||
"speculative_config": "{\"method\":\"ngram\",\"num_speculative_tokens\":4}"
|
||||
}`})
|
||||
Expect(lo.speculativeConfig).To(MatchJSON(`{"method":"ngram","num_speculative_tokens":4}`))
|
||||
})
|
||||
|
||||
It("maps enable_prefix_caching false onto the force-OFF tri-state", func() {
|
||||
// The C ABI tri-state is 0=model default, 1=on, 2=off, so an explicit
|
||||
// `false` must NOT collapse to the 0 that means "let the model decide".
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_prefix_caching": false}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(2)))
|
||||
})
|
||||
|
||||
It("leaves the prefix-caching tri-state at the model default when unset", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"max_num_seqs": 4}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(0)))
|
||||
})
|
||||
|
||||
It("accepts the radix-attention alias upstream documents for prefix caching", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_radix_attention": true}`})
|
||||
Expect(lo.enablePrefixCaching).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("maps enable_jump_forward onto its own tri-state", func() {
|
||||
// ABI v10. Same tri-state shape as prefix caching, and the same trap:
|
||||
// an explicit false must be force-OFF (2), not the 0 that defers to the
|
||||
// environment.
|
||||
on := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_jump_forward": true}`})
|
||||
Expect(on.enableJumpForward).To(Equal(int32(1)))
|
||||
off := parseOptions(&pb.ModelOptions{EngineArgs: `{"enable_jump_forward": false}`})
|
||||
Expect(off.enableJumpForward).To(Equal(int32(2)))
|
||||
unset := parseOptions(&pb.ModelOptions{EngineArgs: `{"max_num_seqs": 4}`})
|
||||
Expect(unset.enableJumpForward).To(Equal(int32(0)))
|
||||
})
|
||||
|
||||
It("reads enable_jump_forward from the legacy options list too", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{Options: []string{"enable_jump_forward:true"}})
|
||||
Expect(lo.enableJumpForward).To(Equal(int32(1)))
|
||||
})
|
||||
|
||||
It("lets engine_args override the legacy options list", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"max_num_seqs:8", "block_size:16"},
|
||||
EngineArgs: `{"max_num_seqs": 64}`,
|
||||
})
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(64))) // engine_args wins
|
||||
Expect(lo.blockSize).To(Equal(int32(16))) // untouched keys survive
|
||||
})
|
||||
|
||||
It("ignores malformed engine_args rather than failing the load", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"max_num_seqs:8"},
|
||||
EngineArgs: `{not json`,
|
||||
})
|
||||
Expect(lo.maxNumSeqs).To(Equal(int32(8)))
|
||||
})
|
||||
|
||||
It("ignores unknown keys", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{EngineArgs: `{"gpu_memory_utilization": 0.9}`})
|
||||
Expect(lo).To(Equal(loadOptions{}))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("samplingFromPredict", func() {
|
||||
@@ -285,91 +135,6 @@ var _ = Describe("samplingFromPredict", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// The engine resolves speculative_config.model against a local directory or
|
||||
// ~/.cache/huggingface/hub ONLY - it never downloads. LocalAI keeps models in
|
||||
// its own directory, so a bare repo id would miss the HF cache and fail deep in
|
||||
// the load with a confusing "draft checkpoint not found". Resolve it here.
|
||||
var _ = Describe("resolveDraftModelPath", func() {
|
||||
var modelsDir string
|
||||
|
||||
BeforeEach(func() {
|
||||
modelsDir = GinkgoT().TempDir()
|
||||
})
|
||||
|
||||
// draftDir creates a plausible draft checkpoint under models/.
|
||||
draftDir := func(name string) string {
|
||||
d := filepath.Join(modelsDir, name)
|
||||
Expect(os.MkdirAll(d, 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(d, "config.json"), []byte("{}"), 0o600)).To(Succeed())
|
||||
return d
|
||||
}
|
||||
|
||||
It("rewrites a repo id to the matching directory in the models dir", func() {
|
||||
want := draftDir("Qwen3.6-27B-DFlash")
|
||||
spec := `{"method":"dflash","model":"z-lab/Qwen3.6-27B-DFlash"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(`{"method":"dflash","model":"` + want + `"}`))
|
||||
})
|
||||
|
||||
It("rewrites a models-dir-relative path", func() {
|
||||
want := draftDir("drafts__dflash")
|
||||
spec := `{"method":"dflash","model":"drafts__dflash"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(ContainSubstring(want))
|
||||
})
|
||||
|
||||
It("leaves an absolute path that already resolves alone", func() {
|
||||
abs := draftDir("elsewhere")
|
||||
spec := `{"method":"dflash","model":"` + abs + `"}`
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(spec))
|
||||
})
|
||||
|
||||
It("fails with an actionable error when the draft is nowhere on disk", func() {
|
||||
// Silently passing the repo id through would surface as an HF-cache
|
||||
// miss inside the engine, which reads as "your model is broken".
|
||||
spec := `{"method":"dflash","model":"z-lab/Not-Downloaded"}`
|
||||
_, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("z-lab/Not-Downloaded"))
|
||||
Expect(err.Error()).To(ContainSubstring(modelsDir))
|
||||
})
|
||||
|
||||
It("requires a model key for dflash", func() {
|
||||
_, err := resolveDraftModelPath(`{"method":"dflash"}`, modelsDir)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("model"))
|
||||
})
|
||||
|
||||
It("leaves mtp and ngram configs untouched", func() {
|
||||
// Neither has a separate draft checkpoint to resolve.
|
||||
for _, spec := range []string{
|
||||
`{"method":"mtp"}`,
|
||||
`{"method":"ngram","num_speculative_tokens":4}`,
|
||||
} {
|
||||
out, err := resolveDraftModelPath(spec, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(MatchJSON(spec))
|
||||
}
|
||||
})
|
||||
|
||||
It("passes a malformed document through for the engine to reject", func() {
|
||||
// The engine owns config validation and produces the better message.
|
||||
out, err := resolveDraftModelPath(`{not json`, modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(Equal(`{not json`))
|
||||
})
|
||||
|
||||
It("is a no-op on an empty config", func() {
|
||||
out, err := resolveDraftModelPath("", modelsDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("validModelPath", func() {
|
||||
It("accepts a .gguf file", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# whisper.cpp version
|
||||
WHISPER_REPO?=https://github.com/ggml-org/whisper.cpp
|
||||
WHISPER_CPP_VERSION?=64d57d3df5c8dacee098577257edcaa154bf5ef3
|
||||
WHISPER_CPP_VERSION?=2ca53bb45e38748d07b310eeb36245a7157ac882
|
||||
SO_TARGET?=libgowhisper.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -883,34 +883,6 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
|
||||
return backend_pb2.Result(message="Media generated", success=True)
|
||||
|
||||
def UpscaleImage(self, request, context):
|
||||
try:
|
||||
if not request.src:
|
||||
return backend_pb2.Result(success=False, message="No source image provided")
|
||||
if not request.dst:
|
||||
return backend_pb2.Result(success=False, message="No destination path provided")
|
||||
|
||||
scale = request.scale if request.scale > 0 else 2
|
||||
image = Image.open(request.src).convert("RGB")
|
||||
|
||||
# If the loaded pipeline supports upscaling (e.g. StableDiffusionUpscalePipeline),
|
||||
# use it; otherwise fall back to high-quality Lanczos resize.
|
||||
if self.pipe is not None and self.PipelineType in ("StableDiffusionUpscalePipeline", "StableDiffusionLatentUpscalePipeline"):
|
||||
print(f"UpscaleImage: using diffusers upscale pipeline ({self.PipelineType})", file=sys.stderr)
|
||||
upscaled = self.pipe(prompt="", image=image).images[0]
|
||||
else:
|
||||
# Fallback: high-quality Lanczos resize
|
||||
print(f"UpscaleImage: no upscale pipeline loaded, using Lanczos resize (scale={scale})", file=sys.stderr)
|
||||
new_w = image.width * scale
|
||||
new_h = image.height * scale
|
||||
upscaled = image.resize((new_w, new_h), Image.LANCZOS)
|
||||
|
||||
upscaled.save(request.dst)
|
||||
return backend_pb2.Result(message="Image upscaled", success=True)
|
||||
except Exception as e:
|
||||
print(f"UpscaleImage error: {e}", file=sys.stderr)
|
||||
return backend_pb2.Result(success=False, message=str(e))
|
||||
|
||||
def GenerateVideo(self, request, context):
|
||||
try:
|
||||
prompt = request.prompt
|
||||
|
||||
@@ -10,4 +10,11 @@ else
|
||||
source $backend_dir/../common/libbackend.sh
|
||||
fi
|
||||
|
||||
# CUDA 13 has no prebuilt FlashAttention wheel, so the fallback source build
|
||||
# exceeds the CI runner's memory when ninja compiles multiple units at once.
|
||||
if [ "x${BUILD_PROFILE}" = "xcublas13" ]; then
|
||||
export MAX_JOBS="${MAX_JOBS:-1}"
|
||||
export NVCC_THREADS="${NVCC_THREADS:-1}"
|
||||
fi
|
||||
|
||||
installRequirements
|
||||
|
||||
1
backend/python/qwen-tts/requirements-cublas13-after.txt
Normal file
1
backend/python/qwen-tts/requirements-cublas13-after.txt
Normal file
@@ -0,0 +1 @@
|
||||
flash-attn
|
||||
@@ -2,6 +2,15 @@
|
||||
set -e
|
||||
|
||||
backend_dir=$(dirname $0)
|
||||
|
||||
for cuda_version in 12 13; do
|
||||
grep -qx "flash-attn" "$backend_dir/requirements-cublas${cuda_version}-after.txt"
|
||||
done
|
||||
|
||||
grep -q 'BUILD_PROFILE.*cublas13' "$backend_dir/install.sh"
|
||||
grep -q 'MAX_JOBS.*1' "$backend_dir/install.sh"
|
||||
grep -q 'NVCC_THREADS.*1' "$backend_dir/install.sh"
|
||||
|
||||
if [ -d $backend_dir/common ]; then
|
||||
source $backend_dir/common/libbackend.sh
|
||||
else
|
||||
|
||||
@@ -15,12 +15,3 @@ sglang[all]>=0.5.11
|
||||
# load-bearing for flash-attn-4, and this is the narrower change. Raise the
|
||||
# bound once 0.46.0 final ships.
|
||||
nvidia-modelopt<0.46
|
||||
|
||||
# Same failure mode as the nvidia-modelopt bound above, via a different
|
||||
# package. sglang -> flashinfer-python -> cuda-tile, unbounded, and the
|
||||
# global --prerelease=allow resolves it to 1.6.0rc3, whose build backend
|
||||
# imports wheel_stub without declaring it in build-system.requires. With
|
||||
# --no-build-isolation nothing installs it and the build dies with
|
||||
# "No module named 'wheel_stub'". 1.5.0 is the newest stable release.
|
||||
# Raise the bound once 1.6.0 final ships.
|
||||
cuda-tile<1.6
|
||||
|
||||
@@ -15,12 +15,3 @@ sglang[all]>=0.5.11
|
||||
# load-bearing for flash-attn-4, and this is the narrower change. Raise the
|
||||
# bound once 0.46.0 final ships.
|
||||
nvidia-modelopt<0.46
|
||||
|
||||
# Same failure mode as the nvidia-modelopt bound above, via a different
|
||||
# package. sglang -> flashinfer-python -> cuda-tile, unbounded, and the
|
||||
# global --prerelease=allow resolves it to 1.6.0rc3, whose build backend
|
||||
# imports wheel_stub without declaring it in build-system.requires. With
|
||||
# --no-build-isolation nothing installs it and the build dies with
|
||||
# "No module named 'wheel_stub'". 1.5.0 is the newest stable release.
|
||||
# Raise the bound once 1.6.0 final ships.
|
||||
cuda-tile<1.6
|
||||
|
||||
@@ -13,12 +13,3 @@
|
||||
# FunctionCallParser, ReasoningParser); the [all] extras are optional
|
||||
# accelerators not required at import time.
|
||||
sglang>=0.5.11
|
||||
|
||||
# Same failure mode the cublas profiles carry an nvidia-modelopt bound for,
|
||||
# reached through a different package. sglang -> flashinfer-python ->
|
||||
# cuda-tile, unbounded, and the global --prerelease=allow resolves it to
|
||||
# 1.6.0rc3, whose build backend imports wheel_stub without declaring it in
|
||||
# build-system.requires. With --no-build-isolation nothing installs it and
|
||||
# the build dies with "No module named 'wheel_stub'". 1.5.0 is the newest
|
||||
# stable release. Raise the bound once 1.6.0 final ships.
|
||||
cuda-tile<1.6
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
@@ -108,13 +107,6 @@ For documentation and support:
|
||||
// Run the thing!
|
||||
err = ctx.Run(&cli.CLI.Context)
|
||||
if err != nil {
|
||||
// A command that has already told the user what went wrong returns
|
||||
// only a status. Logging it as well would print a bare "exit status 1"
|
||||
// underneath the explanation they just read.
|
||||
var reported cli.ExitCodeError
|
||||
if errors.As(err, &reported) {
|
||||
os.Exit(reported.Code)
|
||||
}
|
||||
xlog.Fatal("Error running the application", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -553,17 +553,12 @@ func (a *Application) start() error {
|
||||
// once at startup and reused across chat sessions that opt in via metadata.
|
||||
if !a.applicationConfig.DisableLocalAIAssistant {
|
||||
holder := mcpTools.NewLocalAIAssistantHolder()
|
||||
var nodeRegistry *nodes.NodeRegistry
|
||||
if a.distributed != nil {
|
||||
nodeRegistry = a.distributed.Registry
|
||||
}
|
||||
assistantClient := localaiInproc.New(
|
||||
a.applicationConfig,
|
||||
a.applicationConfig.SystemState,
|
||||
a.backendLoader,
|
||||
a.modelLoader,
|
||||
a.galleryService,
|
||||
nodeRegistry,
|
||||
)
|
||||
// Wire usage tracking so the assistant's get_usage_stats tool
|
||||
// returns real data; nil values keep the tool returning a clear
|
||||
|
||||
@@ -444,13 +444,6 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
// when gallery data refreshes instead of using a fixed TTL.
|
||||
vram.SetGalleryGenerationFunc(gallery.GalleryGeneration)
|
||||
|
||||
// Fill those caches ahead of the first visitor. An estimate for an entry
|
||||
// nobody has asked about yet costs a remote probe of its weight files, and
|
||||
// the model gallery asks for one per row, so without this the first page
|
||||
// spends seconds filling in its own sizes while somebody watches it.
|
||||
// Non-blocking, and bounded: see DefaultEstimateWarmConfig.
|
||||
gallery.WarmEstimateCache(options.Context, options.Galleries, options.SystemState, gallery.EstimateWarmConfigFromEnv())
|
||||
|
||||
if options.ConfigFile != "" {
|
||||
if err := application.ModelConfigLoader().LoadMultipleModelConfigsSingleFile(options.ConfigFile, configLoaderOpts...); err != nil {
|
||||
xlog.Error("error loading config file", "error", err)
|
||||
|
||||
@@ -1,37 +0,0 @@
|
||||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
model "github.com/mudler/LocalAI/pkg/model"
|
||||
)
|
||||
|
||||
// ImageUpscale loads the model specified in modelConfig and calls UpscaleImage
|
||||
// on the backend, writing the result to dst.
|
||||
func ImageUpscale(ctx context.Context, src, dst string, scale int, loader *model.ModelLoader, modelConfig config.ModelConfig, appConfig *config.ApplicationConfig) (func() error, error) {
|
||||
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
|
||||
inferenceModel, err := loader.Load(opts...)
|
||||
if err != nil {
|
||||
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
fn := func() error {
|
||||
_, err := inferenceModel.UpscaleImage(
|
||||
ctx,
|
||||
&proto.UpscaleImageRequest{
|
||||
Src: src,
|
||||
Dst: dst,
|
||||
Scale: int32(scale),
|
||||
},
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
return fn, nil
|
||||
}
|
||||
|
||||
// ImageUpscaleFunc is a test-friendly indirection.
|
||||
var ImageUpscaleFunc = ImageUpscale
|
||||
30
core/cli/chat/chat.go
Normal file
30
core/cli/chat/chat.go
Normal file
@@ -0,0 +1,30 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type Options struct {
|
||||
Model string
|
||||
BaseURL string
|
||||
APIKey string
|
||||
In io.Reader
|
||||
Out io.Writer
|
||||
}
|
||||
|
||||
func Run(ctx context.Context, opts Options) error {
|
||||
if opts.In == nil {
|
||||
opts.In = strings.NewReader("")
|
||||
}
|
||||
if opts.Out == nil {
|
||||
opts.Out = io.Discard
|
||||
}
|
||||
|
||||
session, err := newChatSession(ctx, newLocalAIChatClient(opts.BaseURL, opts.APIKey), opts.Model)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return runTerminalChat(ctx, session, opts.In, opts.Out)
|
||||
}
|
||||
172
core/cli/chat/chat_test.go
Normal file
172
core/cli/chat/chat_test.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Run chat", func() {
|
||||
It("streams a single chat response", func() {
|
||||
var capturedModel string
|
||||
var capturedAuth string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v1/models" {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
writeResponse(w, `{"object":"list","data":[{"id":"test-model","object":"model"}]}`)
|
||||
return
|
||||
}
|
||||
|
||||
Expect(r.URL.Path).To(Equal("/v1/chat/completions"))
|
||||
capturedAuth = r.Header.Get("Authorization")
|
||||
|
||||
var body struct {
|
||||
Model string `json:"model"`
|
||||
Messages []struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
} `json:"messages"`
|
||||
}
|
||||
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
|
||||
capturedModel = body.Model
|
||||
Expect(body.Messages).To(HaveLen(1))
|
||||
Expect(body.Messages[0].Role).To(Equal("user"))
|
||||
Expect(body.Messages[0].Content).To(Equal("hello"))
|
||||
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
writeResponse(w, "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hi\"}}]}\n\n")
|
||||
writeResponse(w, "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"!\"}}]}\n\n")
|
||||
writeResponse(w, "data: [DONE]\n\n")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
var out bytes.Buffer
|
||||
err := Run(GinkgoT().Context(), Options{
|
||||
Model: "test-model",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "secret",
|
||||
In: strings.NewReader("hello\n/exit\n"),
|
||||
Out: &out,
|
||||
})
|
||||
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(capturedModel).To(Equal("test-model"))
|
||||
Expect(capturedAuth).To(Equal("Bearer secret"))
|
||||
Expect(out.String()).To(ContainSubstring("assistant: hi!"))
|
||||
Expect(out.String()).To(ContainSubstring("bye"))
|
||||
})
|
||||
|
||||
It("auto-selects the only available model", func() {
|
||||
server := chatTestServer([]string{"solo"}, nil)
|
||||
defer server.Close()
|
||||
|
||||
var out bytes.Buffer
|
||||
err := Run(GinkgoT().Context(), Options{
|
||||
BaseURL: server.URL + "/v1",
|
||||
In: strings.NewReader("/exit\n"),
|
||||
Out: &out,
|
||||
})
|
||||
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out.String()).To(ContainSubstring("LocalAI chat (solo)"))
|
||||
})
|
||||
|
||||
It("returns an actionable error when no models are installed", func() {
|
||||
server := chatTestServer(nil, nil)
|
||||
defer server.Close()
|
||||
|
||||
err := Run(GinkgoT().Context(), Options{
|
||||
BaseURL: server.URL + "/v1",
|
||||
In: strings.NewReader(""),
|
||||
})
|
||||
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("no chat models are installed"))
|
||||
Expect(err.Error()).To(ContainSubstring("local-ai models install <model>"))
|
||||
})
|
||||
|
||||
It("returns an actionable error when multiple models are available without a selection", func() {
|
||||
server := chatTestServer([]string{"alpha", "beta"}, nil)
|
||||
defer server.Close()
|
||||
|
||||
err := Run(GinkgoT().Context(), Options{
|
||||
BaseURL: server.URL + "/v1",
|
||||
In: strings.NewReader(""),
|
||||
})
|
||||
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("multiple models are available"))
|
||||
Expect(err.Error()).To(ContainSubstring("--model"))
|
||||
Expect(err.Error()).To(ContainSubstring("alpha"))
|
||||
Expect(err.Error()).To(ContainSubstring("beta"))
|
||||
})
|
||||
|
||||
It("lists and switches models inside the chat", func() {
|
||||
requestedModels := []string{}
|
||||
server := chatTestServer([]string{"alpha", "beta"}, func(model string) {
|
||||
requestedModels = append(requestedModels, model)
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
var out bytes.Buffer
|
||||
err := Run(GinkgoT().Context(), Options{
|
||||
Model: "alpha",
|
||||
BaseURL: server.URL + "/v1",
|
||||
In: strings.NewReader("/models\n/model beta\nhello\n/exit\n"),
|
||||
Out: &out,
|
||||
})
|
||||
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out.String()).To(ContainSubstring("* alpha"))
|
||||
Expect(out.String()).To(ContainSubstring(" beta"))
|
||||
Expect(out.String()).To(ContainSubstring("switched to beta; conversation cleared"))
|
||||
Expect(requestedModels).To(Equal([]string{"beta"}))
|
||||
})
|
||||
})
|
||||
|
||||
func chatTestServer(models []string, onChat func(model string)) *httptest.Server {
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/v1/models":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
writeResponse(w, `{"object":"list","data":[`)
|
||||
for i, model := range models {
|
||||
if i > 0 {
|
||||
writeResponse(w, ",")
|
||||
}
|
||||
writeResponsef(w, `{"id":%q,"object":"model"}`, model)
|
||||
}
|
||||
writeResponse(w, `]}`)
|
||||
case "/v1/chat/completions":
|
||||
var body struct {
|
||||
Model string `json:"model"`
|
||||
}
|
||||
Expect(json.NewDecoder(r.Body).Decode(&body)).To(Succeed())
|
||||
if onChat != nil {
|
||||
onChat(body.Model)
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
writeResponse(w, "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"}}]}\n\n")
|
||||
writeResponse(w, "data: [DONE]\n\n")
|
||||
default:
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
func writeResponse(w io.Writer, text string) {
|
||||
_, err := fmt.Fprint(w, text)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
}
|
||||
|
||||
func writeResponsef(w io.Writer, format string, args ...any) {
|
||||
_, err := fmt.Fprintf(w, format, args...)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
}
|
||||
114
core/cli/chat/client.go
Normal file
114
core/cli/chat/client.go
Normal file
@@ -0,0 +1,114 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
openai "github.com/sashabaranov/go-openai"
|
||||
)
|
||||
|
||||
type chatClient interface {
|
||||
ListModels(ctx context.Context) ([]string, error)
|
||||
StreamChat(ctx context.Context, model string, messages []chatMessage, out io.Writer) (string, error)
|
||||
}
|
||||
|
||||
type localAIChatClient struct {
|
||||
client *openai.Client
|
||||
}
|
||||
|
||||
func newLocalAIChatClient(baseURL string, apiKey string) *localAIChatClient {
|
||||
cfg := openai.DefaultConfig(apiKey)
|
||||
cfg.BaseURL = baseURL
|
||||
return &localAIChatClient{client: openai.NewClientWithConfig(cfg)}
|
||||
}
|
||||
|
||||
func (c *localAIChatClient) ListModels(ctx context.Context) ([]string, error) {
|
||||
resp, err := c.client.ListModels(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
models := make([]string, 0, len(resp.Models))
|
||||
for _, model := range resp.Models {
|
||||
if model.ID != "" {
|
||||
models = append(models, model.ID)
|
||||
}
|
||||
}
|
||||
sort.Strings(models)
|
||||
return models, nil
|
||||
}
|
||||
|
||||
func (c *localAIChatClient) StreamChat(ctx context.Context, model string, messages []chatMessage, out io.Writer) (string, error) {
|
||||
stream, err := c.client.CreateChatCompletionStream(ctx, openai.ChatCompletionRequest{
|
||||
Model: model,
|
||||
Messages: openAIChatMessages(messages),
|
||||
})
|
||||
if err != nil {
|
||||
return "", friendlyChatError(err, model)
|
||||
}
|
||||
defer func() {
|
||||
_ = stream.Close()
|
||||
}()
|
||||
|
||||
var answer strings.Builder
|
||||
for {
|
||||
resp, err := stream.Recv()
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return answer.String(), friendlyChatError(err, model)
|
||||
}
|
||||
if len(resp.Choices) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
token := resp.Choices[0].Delta.Content
|
||||
if token == "" {
|
||||
continue
|
||||
}
|
||||
answer.WriteString(token)
|
||||
if _, err := fmt.Fprint(out, token); err != nil {
|
||||
return answer.String(), err
|
||||
}
|
||||
}
|
||||
|
||||
return answer.String(), nil
|
||||
}
|
||||
|
||||
func openAIChatMessages(messages []chatMessage) []openai.ChatCompletionMessage {
|
||||
converted := make([]openai.ChatCompletionMessage, len(messages))
|
||||
for i, message := range messages {
|
||||
converted[i] = openai.ChatCompletionMessage{
|
||||
Role: message.Role,
|
||||
Content: message.Content,
|
||||
}
|
||||
}
|
||||
return converted
|
||||
}
|
||||
|
||||
func friendlyChatError(err error, model string) error {
|
||||
var apiErr *openai.APIError
|
||||
if errors.As(err, &apiErr) {
|
||||
switch apiErr.HTTPStatusCode {
|
||||
case 404:
|
||||
return fmt.Errorf("model %q is not available. Run `local-ai models list`, install a model with `local-ai models install <model>`, or switch with `/model <name>`", model)
|
||||
case 403:
|
||||
return fmt.Errorf("model %q is disabled. Enable it from LocalAI settings or choose another model with `/model <name>`", model)
|
||||
}
|
||||
if apiErr.Message != "" {
|
||||
return errors.New(apiErr.Message)
|
||||
}
|
||||
}
|
||||
|
||||
msg := err.Error()
|
||||
if strings.Contains(msg, "model") && strings.Contains(msg, "not found") {
|
||||
return fmt.Errorf("model %q is not available. Run `local-ai models list`, install a model with `local-ai models install <model>`, or switch with `/model <name>`", model)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
17
core/cli/chat/models.go
Normal file
17
core/cli/chat/models.go
Normal file
@@ -0,0 +1,17 @@
|
||||
package chat
|
||||
|
||||
import "strings"
|
||||
|
||||
func formatChatModelList(models []string, current string) string {
|
||||
var b strings.Builder
|
||||
for _, model := range models {
|
||||
prefix := " "
|
||||
if model == current {
|
||||
prefix = "* "
|
||||
}
|
||||
b.WriteString(prefix)
|
||||
b.WriteString(model)
|
||||
b.WriteByte('\n')
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
@@ -1,153 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// stateDirMode matches the mode nib uses for the same directory. The directory
|
||||
// holds an API key, so it stays owner-only.
|
||||
const stateDirMode = 0o700
|
||||
|
||||
// configFileMode keeps the config owner-only: nib stores the user's API key in
|
||||
// it alongside the keys written here.
|
||||
const configFileMode = 0o600
|
||||
|
||||
// StateDir resolves where the chat agent keeps its config, plugins, and
|
||||
// skills. This is user-scoped rather than server-scoped: chat is a client that
|
||||
// may target a remote LocalAI, so it does not belong under LOCALAI_CONFIG_DIR.
|
||||
func StateDir(override string) (string, error) {
|
||||
if override != "" {
|
||||
return override, nil
|
||||
}
|
||||
if xdg := os.Getenv("XDG_CONFIG_HOME"); xdg != "" {
|
||||
return filepath.Join(xdg, "localai", "chat"), nil
|
||||
}
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolving home directory for the agent state dir: %w", err)
|
||||
}
|
||||
return filepath.Join(home, ".config", "localai", "chat"), nil
|
||||
}
|
||||
|
||||
// ConfigPath is the agent's config file inside dir.
|
||||
func ConfigPath(dir string) string { return filepath.Join(dir, "config.yaml") }
|
||||
|
||||
// EnsureStateDir creates dir and, on first run only, seeds a config file
|
||||
// pointing at baseURL. It deliberately does not seed a model: a baked-in model
|
||||
// name goes stale as soon as the user installs a different one.
|
||||
//
|
||||
// The config file is machine-managed from here on: nib rewrites it whenever it
|
||||
// self-configures, so hand-written comments in it do not survive.
|
||||
func EnsureStateDir(dir, baseURL string) error {
|
||||
if err := os.MkdirAll(dir, stateDirMode); err != nil {
|
||||
return fmt.Errorf("creating agent state dir %s: %w", dir, err)
|
||||
}
|
||||
path := ConfigPath(dir)
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return nil // already configured; never overwrite the user's file
|
||||
} else if !os.IsNotExist(err) {
|
||||
return fmt.Errorf("checking agent config %s: %w", path, err)
|
||||
}
|
||||
|
||||
seed := map[string]string{"base_url": baseURL}
|
||||
data, err := yaml.Marshal(seed)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encoding seed agent config: %w", err)
|
||||
}
|
||||
if err := writeConfigFile(path, data); err != nil {
|
||||
return fmt.Errorf("writing seed agent config: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// PersistModel records the chosen model in the agent config, preserving every
|
||||
// other key the user may have set, including the api_key nib writes there.
|
||||
//
|
||||
// The file is machine-managed: this overlays the model onto the parsed keys and
|
||||
// re-marshals, which drops comments. That is deliberate rather than an
|
||||
// oversight, because nib's own save path does the same thing and would erase
|
||||
// them on its next write regardless.
|
||||
func PersistModel(dir, model string) error {
|
||||
// PersistModel is callable before EnsureStateDir, so it cannot assume the
|
||||
// directory exists.
|
||||
if err := os.MkdirAll(dir, stateDirMode); err != nil {
|
||||
return fmt.Errorf("creating agent state dir %s: %w", dir, err)
|
||||
}
|
||||
path := ConfigPath(dir)
|
||||
|
||||
values := map[string]any{}
|
||||
// #nosec G304 -- path is the fixed config.yaml name under the user-selected
|
||||
// chat state directory; selecting that directory is the documented override.
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("reading agent config %s: %w", path, err)
|
||||
}
|
||||
if err == nil {
|
||||
if err := yaml.Unmarshal(data, &values); err != nil {
|
||||
return fmt.Errorf("parsing agent config %s: %w", path, err)
|
||||
}
|
||||
}
|
||||
values["model"] = model
|
||||
|
||||
out, err := yaml.Marshal(values)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encoding agent config: %w", err)
|
||||
}
|
||||
if err := writeConfigFile(path, out); err != nil {
|
||||
return fmt.Errorf("writing agent config: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeConfigFile replaces path with data atomically: it writes a temporary
|
||||
// file next to the target and renames it over the target. Writing the target in
|
||||
// place would truncate it first, so an interrupted or out-of-disk write would
|
||||
// leave a half-written config and destroy the api_key nib keeps in the same
|
||||
// file. The temporary file must share the directory because rename is only
|
||||
// atomic within one filesystem.
|
||||
func writeConfigFile(path string, data []byte) error {
|
||||
dir := filepath.Dir(path)
|
||||
|
||||
// A randomized name rather than a fixed config.yaml.tmp, so two concurrent
|
||||
// writers cannot corrupt each other's temporary file.
|
||||
tmp, err := os.CreateTemp(dir, "config.yaml.*.tmp")
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating temp file in %s: %w", dir, err)
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
renamed := false
|
||||
defer func() {
|
||||
if !renamed {
|
||||
// Leave no litter behind on any failure path.
|
||||
_ = os.Remove(tmpPath)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := tmp.Write(data); err != nil {
|
||||
_ = tmp.Close()
|
||||
return fmt.Errorf("writing %s: %w", tmpPath, err)
|
||||
}
|
||||
// Flush before the rename: renaming a file whose contents are still only in
|
||||
// the page cache can still lose them across a crash.
|
||||
if err := tmp.Sync(); err != nil {
|
||||
_ = tmp.Close()
|
||||
return fmt.Errorf("syncing %s: %w", tmpPath, err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return fmt.Errorf("closing %s: %w", tmpPath, err)
|
||||
}
|
||||
// CreateTemp already asks for 0600, but the umask can only ever clear bits,
|
||||
// so set the mode explicitly rather than inheriting whatever survived.
|
||||
if err := os.Chmod(tmpPath, configFileMode); err != nil {
|
||||
return fmt.Errorf("setting mode on %s: %w", tmpPath, err)
|
||||
}
|
||||
if err := os.Rename(tmpPath, path); err != nil {
|
||||
return fmt.Errorf("replacing %s: %w", path, err)
|
||||
}
|
||||
renamed = true
|
||||
return nil
|
||||
}
|
||||
@@ -1,186 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// richConfig stands in for a config nib has already taken ownership of: a
|
||||
// comment, a secret, and a nested block. A flat scalar alone would not catch a
|
||||
// writer that mangles structure or drops a key it does not know about.
|
||||
const richConfig = `# hand written note
|
||||
base_url: http://x.invalid/v1
|
||||
api_key: secret-token
|
||||
mcp_servers:
|
||||
files:
|
||||
command: mcp-files
|
||||
args:
|
||||
- --root
|
||||
- /tmp
|
||||
`
|
||||
|
||||
var _ = Describe("Agent state directory", func() {
|
||||
Describe("StateDir", func() {
|
||||
It("prefers an explicit override", func() {
|
||||
Expect(StateDir("/custom/dir")).To(Equal("/custom/dir"))
|
||||
})
|
||||
|
||||
It("uses XDG_CONFIG_HOME when set", func() {
|
||||
tmp := GinkgoT().TempDir()
|
||||
GinkgoT().Setenv("XDG_CONFIG_HOME", tmp)
|
||||
Expect(StateDir("")).To(Equal(filepath.Join(tmp, "localai", "chat")))
|
||||
})
|
||||
|
||||
It("falls back to ~/.config/localai/chat", func() {
|
||||
tmp := GinkgoT().TempDir()
|
||||
GinkgoT().Setenv("XDG_CONFIG_HOME", "")
|
||||
GinkgoT().Setenv("HOME", tmp)
|
||||
Expect(StateDir("")).To(Equal(filepath.Join(tmp, ".config", "localai", "chat")))
|
||||
})
|
||||
|
||||
It("fails when neither XDG_CONFIG_HOME nor a home directory is resolvable", func() {
|
||||
GinkgoT().Setenv("XDG_CONFIG_HOME", "")
|
||||
GinkgoT().Setenv("HOME", "")
|
||||
|
||||
dir, err := StateDir("")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("agent state dir"))
|
||||
// No silent fallback to a relative path: writing an API key into the
|
||||
// working directory would be worse than refusing.
|
||||
Expect(dir).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("EnsureStateDir", func() {
|
||||
It("creates the directory and seeds base_url on first run", func() {
|
||||
dir := filepath.Join(GinkgoT().TempDir(), "chat")
|
||||
Expect(EnsureStateDir(dir, "http://127.0.0.1:8080/v1")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("base_url: http://127.0.0.1:8080/v1"))
|
||||
// A model must NOT be seeded: it goes stale as soon as the user
|
||||
// installs a different one.
|
||||
Expect(string(data)).ToNot(ContainSubstring("model:"))
|
||||
})
|
||||
|
||||
It("keeps the seeded config and its directory owner-only", func() {
|
||||
dir := filepath.Join(GinkgoT().TempDir(), "chat")
|
||||
Expect(EnsureStateDir(dir, "http://127.0.0.1:8080/v1")).To(Succeed())
|
||||
|
||||
// nib writes the user's api_key into this same file, so the modes are
|
||||
// load-bearing, not cosmetic.
|
||||
config, err := os.Stat(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(config.Mode().Perm()).To(Equal(os.FileMode(0o600)))
|
||||
|
||||
state, err := os.Stat(dir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(state.Mode().Perm()).To(Equal(os.FileMode(0o700)))
|
||||
})
|
||||
|
||||
It("leaves an existing config byte-for-byte untouched", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte(richConfig), 0o600)).To(Succeed())
|
||||
|
||||
Expect(EnsureStateDir(dir, "http://127.0.0.1:8080/v1")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
// Byte-exact against a fixture carrying a comment and a nested block:
|
||||
// an implementation that "preserves" by re-marshaling through a map
|
||||
// fails here rather than passing on a flat scalar.
|
||||
Expect(string(data)).To(Equal(richConfig))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("PersistModel", func() {
|
||||
It("adds a model to an existing config, preserving other keys", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte("base_url: http://x.invalid/v1\n"), 0o600)).To(Succeed())
|
||||
|
||||
Expect(PersistModel(dir, "chosen-model")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("base_url: http://x.invalid/v1"))
|
||||
Expect(string(data)).To(ContainSubstring("model: chosen-model"))
|
||||
})
|
||||
|
||||
It("replaces an existing model rather than duplicating the key", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte("model: old\nbase_url: http://x.invalid/v1\n"), 0o600)).To(Succeed())
|
||||
|
||||
Expect(PersistModel(dir, "new")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("model: new"))
|
||||
Expect(string(data)).ToNot(ContainSubstring("model: old"))
|
||||
})
|
||||
|
||||
It("preserves secrets and nested blocks it does not understand", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte(richConfig), 0o600)).To(Succeed())
|
||||
|
||||
Expect(PersistModel(dir, "chosen-model")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
var got map[string]any
|
||||
Expect(yaml.Unmarshal(data, &got)).To(Succeed())
|
||||
Expect(got).To(HaveKeyWithValue("model", "chosen-model"))
|
||||
Expect(got).To(HaveKeyWithValue("base_url", "http://x.invalid/v1"))
|
||||
// Losing this key logs the user out of their own server.
|
||||
Expect(got).To(HaveKeyWithValue("api_key", "secret-token"))
|
||||
Expect(got).To(HaveKeyWithValue("mcp_servers",
|
||||
HaveKeyWithValue("files", And(
|
||||
HaveKeyWithValue("command", "mcp-files"),
|
||||
HaveKeyWithValue("args", ConsistOf("--root", "/tmp")),
|
||||
)),
|
||||
))
|
||||
|
||||
// Documented, accepted behavior rather than an aspiration: the overlay
|
||||
// re-marshals, so comments do not survive. nib's own save path erases
|
||||
// them too, so preserving them here would buy nothing.
|
||||
Expect(string(data)).ToNot(ContainSubstring("# hand written note"))
|
||||
})
|
||||
|
||||
It("keeps the rewritten config owner-only and leaves no temp file behind", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte(richConfig), 0o600)).To(Succeed())
|
||||
|
||||
Expect(PersistModel(dir, "chosen-model")).To(Succeed())
|
||||
|
||||
info, err := os.Stat(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(info.Mode().Perm()).To(Equal(os.FileMode(0o600)))
|
||||
|
||||
// The atomic write stages through a sibling temp file; it must not
|
||||
// survive a successful write.
|
||||
entries, err := os.ReadDir(dir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
names := []string{}
|
||||
for _, entry := range entries {
|
||||
names = append(names, entry.Name())
|
||||
}
|
||||
Expect(names).To(ConsistOf("config.yaml"))
|
||||
})
|
||||
|
||||
It("creates the state directory when it does not exist yet", func() {
|
||||
// Task 4 may persist a picked model before anything else has run.
|
||||
dir := filepath.Join(GinkgoT().TempDir(), "chat")
|
||||
|
||||
Expect(PersistModel(dir, "chosen-model")).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("model: chosen-model"))
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -1,86 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
||||
openai "github.com/sashabaranov/go-openai"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrUnreachable means nothing answered at the endpoint. Callers use this
|
||||
// to decide whether offering to start a server makes sense.
|
||||
ErrUnreachable = errors.New("no LocalAI server reachable")
|
||||
// ErrUnauthorized means the server answered but rejected the credentials.
|
||||
ErrUnauthorized = errors.New("LocalAI server rejected the API key")
|
||||
)
|
||||
|
||||
// Probe lists the models the endpoint advertises. It classifies the two
|
||||
// failures that need different advice: nothing listening, and bad credentials.
|
||||
//
|
||||
// The returned list is what the server advertises, verbatim and in server
|
||||
// order. LocalAI happily lists non-model entries it finds in the models
|
||||
// directory (stray archives, dotfiles), and guessing which advertised IDs are
|
||||
// real belongs to whoever presents them, not here.
|
||||
func Probe(ctx context.Context, baseURL, apiKey string) ([]string, error) {
|
||||
cfg := openai.DefaultConfig(apiKey)
|
||||
cfg.BaseURL = baseURL
|
||||
|
||||
resp, err := openai.NewClientWithConfig(cfg).ListModels(ctx)
|
||||
if err != nil {
|
||||
if status, answered := responseStatus(err); answered {
|
||||
if status == http.StatusUnauthorized || status == http.StatusForbidden {
|
||||
return nil, fmt.Errorf("%w: %w", ErrUnauthorized, err)
|
||||
}
|
||||
// The server answered, so it is up; surface its error as-is.
|
||||
return nil, fmt.Errorf("listing models at %s: %w", baseURL, err)
|
||||
}
|
||||
// A caller who cancelled the probe learned nothing about the endpoint,
|
||||
// so claiming it is unreachable would send them to fix a server that
|
||||
// may be fine. A deadline is left alone: an endpoint that cannot answer
|
||||
// within the probe's budget is unreachable for our purposes.
|
||||
var urlErr *url.Error
|
||||
if errors.As(err, &urlErr) && !errors.Is(err, context.Canceled) {
|
||||
// Only a failure to complete the round trip means nothing is
|
||||
// listening. A reply we could not parse is a different problem,
|
||||
// so it falls through to the generic error below.
|
||||
return nil, fmt.Errorf("%w at %s: %w", ErrUnreachable, baseURL, err)
|
||||
}
|
||||
return nil, fmt.Errorf("listing models at %s: %w", baseURL, err)
|
||||
}
|
||||
|
||||
models := make([]string, 0, len(resp.Models))
|
||||
for _, m := range resp.Models {
|
||||
if m.ID != "" {
|
||||
models = append(models, m.ID)
|
||||
}
|
||||
}
|
||||
return models, nil
|
||||
}
|
||||
|
||||
// responseStatus reports the HTTP status a failed call came back with, and
|
||||
// whether there was one at all.
|
||||
//
|
||||
// go-openai splits this across two types depending on the error body, and both
|
||||
// occur against a real LocalAI: it returns *openai.APIError when the body
|
||||
// parses as an OpenAI error envelope, which is what LocalAI's normal error
|
||||
// handler sends, and *openai.RequestError when it does not, which is what
|
||||
// LocalAI sends when started with opaque errors, since that handler replies
|
||||
// with a bare status and no body.
|
||||
func responseStatus(err error) (int, bool) {
|
||||
// *RequestError is checked first because it is the outer type when
|
||||
// go-openai nests one error inside the other; the inner value in that case
|
||||
// carries no status.
|
||||
var reqErr *openai.RequestError
|
||||
if errors.As(err, &reqErr) {
|
||||
return reqErr.HTTPStatusCode, true
|
||||
}
|
||||
var apiErr *openai.APIError
|
||||
if errors.As(err, &apiErr) {
|
||||
return apiErr.HTTPStatusCode, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
@@ -1,169 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Probe", func() {
|
||||
It("returns the advertised models", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{
|
||||
"object": "list",
|
||||
"data": []map[string]string{
|
||||
{"id": "model-a", "object": "model"},
|
||||
{"id": "model-b", "object": "model"},
|
||||
},
|
||||
})).To(Succeed())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
models, err := Probe(context.Background(), srv.URL+"/v1", "")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(models).To(Equal([]string{"model-a", "model-b"}))
|
||||
})
|
||||
|
||||
It("reports an unreachable server distinguishably", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||
url := srv.URL
|
||||
srv.Close() // nothing is listening now
|
||||
|
||||
_, err := Probe(context.Background(), url+"/v1", "")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnreachable)).To(BeTrue(), "want ErrUnreachable, got %v", err)
|
||||
})
|
||||
|
||||
It("reports an auth failure distinguishably", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := Probe(context.Background(), srv.URL+"/v1", "bad-key")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnauthorized)).To(BeTrue(), "want ErrUnauthorized, got %v", err)
|
||||
})
|
||||
|
||||
// LocalAI's normal error handler replies with an OpenAI error envelope, and
|
||||
// its opaque-errors handler replies with a bare status and no body. Those
|
||||
// reach the client as two different go-openai types, so both have to be
|
||||
// classified the same way.
|
||||
It("reports an auth failure carrying an error envelope distinguishably", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{
|
||||
"error": map[string]any{"message": "invalid api key", "code": http.StatusUnauthorized},
|
||||
})).To(Succeed())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := Probe(context.Background(), srv.URL+"/v1", "bad-key")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnauthorized)).To(BeTrue(), "want ErrUnauthorized, got %v", err)
|
||||
})
|
||||
|
||||
It("does not call a server that answered with an error unreachable", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := Probe(context.Background(), srv.URL+"/v1", "")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnreachable)).To(BeFalse(), "a server that replied is not unreachable, got %v", err)
|
||||
Expect(errors.Is(err, ErrUnauthorized)).To(BeFalse(), "500 is not an auth failure, got %v", err)
|
||||
})
|
||||
|
||||
// Pointing chat at some other service that happens to be listening is a
|
||||
// different problem from nothing listening, and needs different advice.
|
||||
It("does not call a reply it could not parse unreachable", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
_, err := w.Write([]byte("<html><body>not LocalAI</body></html>"))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := Probe(context.Background(), srv.URL+"/v1", "")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnreachable)).To(BeFalse(), "something answered, got %v", err)
|
||||
})
|
||||
|
||||
It("returns every advertised id, including ones that are not models", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{
|
||||
"object": "list",
|
||||
"data": []map[string]string{
|
||||
{"id": "zeta", "object": "model"},
|
||||
{"id": ".gitignore", "object": "model"},
|
||||
{"id": "alpha", "object": "model"},
|
||||
{"id": "voice.tar.bz2", "object": "model"},
|
||||
},
|
||||
})).To(Succeed())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
// Verbatim and in server order: deciding which of these are real, and
|
||||
// what order to show them in, belongs to the caller.
|
||||
models, err := Probe(context.Background(), srv.URL+"/v1", "")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(models).To(Equal([]string{"zeta", ".gitignore", "alpha", "voice.tar.bz2"}))
|
||||
})
|
||||
|
||||
It("stops early when the context is already cancelled", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{"object": "list", "data": []any{}})).To(Succeed())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
_, err := Probe(ctx, srv.URL+"/v1", "")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, context.Canceled)).To(BeTrue(), "want the cancellation preserved, got %v", err)
|
||||
// A cancelled probe learned nothing about the endpoint, so it must not
|
||||
// send the caller off to start a server that may already be running.
|
||||
Expect(errors.Is(err, ErrUnreachable)).To(BeFalse(), "cancelling is not a verdict on the server, got %v", err)
|
||||
})
|
||||
|
||||
It("reports a server that never answers as unreachable", func() {
|
||||
release := make(chan struct{})
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
<-release
|
||||
}))
|
||||
defer srv.Close()
|
||||
defer close(release)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
_, err := Probe(ctx, srv.URL+"/v1", "")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrUnreachable)).To(BeTrue(), "want ErrUnreachable, got %v", err)
|
||||
Expect(errors.Is(err, context.DeadlineExceeded)).To(BeTrue(), "want the deadline preserved, got %v", err)
|
||||
})
|
||||
|
||||
It("returns an empty list when the server has no models", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{"object": "list", "data": []any{}})).To(Succeed())
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
models, err := Probe(context.Background(), srv.URL+"/v1", "")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(models).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -1,95 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
// ModelChooser asks the user to pick one of models. It is nil when the session
|
||||
// is not interactive.
|
||||
type ModelChooser func(models []string) (string, error)
|
||||
|
||||
// ModelRequest is everything model resolution needs.
|
||||
type ModelRequest struct {
|
||||
Flag string // --model
|
||||
Configured string // model recorded in the agent config
|
||||
Available []string // models the server advertises
|
||||
StateDir string // where an interactive choice is persisted
|
||||
Choose ModelChooser // nil means non-interactive
|
||||
// Notify reports a problem that is worth telling the user about but not
|
||||
// worth failing over. Nil discards it. It exists because the one such
|
||||
// problem here, a choice that could not be saved, changes what the user
|
||||
// should expect next: they will be asked again. A log line does not reach
|
||||
// them, since the agent runs at log level error by default.
|
||||
Notify func(message string)
|
||||
}
|
||||
|
||||
// ResolveModel picks the model for this invocation. A flag or a configured
|
||||
// value wins outright and is not persisted; only an interactive choice is
|
||||
// written back, so the prompt appears at most once.
|
||||
//
|
||||
// Available is used exactly as the server gave it. LocalAI advertises stray
|
||||
// files it finds in the models directory alongside real models, but real model
|
||||
// IDs contain dots too (lfm2.5-8b-a1b), so any client-side "looks like a
|
||||
// filename" heuristic would eventually hide a model the user has. Deciding
|
||||
// which advertised IDs are real belongs to the endpoint, not to a guess here.
|
||||
func ResolveModel(req ModelRequest) (string, error) {
|
||||
if req.Flag != "" {
|
||||
return req.Flag, nil
|
||||
}
|
||||
if req.Configured != "" {
|
||||
return req.Configured, nil
|
||||
}
|
||||
|
||||
// The server's /v1/models ordering is not stable between calls, so sort
|
||||
// before showing or listing: the same number must mean the same model on
|
||||
// the next run. Sort a copy; the caller's slice is not ours to reorder.
|
||||
available := append([]string(nil), req.Available...)
|
||||
sort.Strings(available)
|
||||
|
||||
switch len(available) {
|
||||
case 0:
|
||||
return "", errors.New("the LocalAI server has no models installed. Install one with 'local-ai models install <name>', then run 'local-ai chat' again")
|
||||
case 1:
|
||||
return available[0], nil
|
||||
}
|
||||
|
||||
if req.Choose == nil {
|
||||
return "", fmt.Errorf(
|
||||
"several models are available; pick one with --model. Available: %s",
|
||||
strings.Join(available, ", "),
|
||||
)
|
||||
}
|
||||
|
||||
chosen, err := req.Choose(available)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// Choose is an interface, so its answer is checked rather than trusted.
|
||||
// What comes back is persisted and every later run starts against it, so a
|
||||
// chooser that returns an empty string or a name of its own would record a
|
||||
// model the server never offered and there would be nothing left to catch
|
||||
// it.
|
||||
if !slices.Contains(available, chosen) {
|
||||
return "", fmt.Errorf(
|
||||
"the model chooser answered %q, which is not one of the available models: %s",
|
||||
chosen, strings.Join(available, ", "),
|
||||
)
|
||||
}
|
||||
if req.StateDir != "" {
|
||||
if err := PersistModel(req.StateDir, chosen); err != nil {
|
||||
// A failure to remember the choice must not block the session: the
|
||||
// user picked a model, so honour it and say what will happen.
|
||||
xlog.Warn("could not save the model choice", "error", err, "model", chosen)
|
||||
if req.Notify != nil {
|
||||
req.Notify(fmt.Sprintf("Your choice of %s could not be saved, so this question comes back next time: %v", chosen, err))
|
||||
}
|
||||
}
|
||||
}
|
||||
return chosen, nil
|
||||
}
|
||||
@@ -1,156 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("ResolveModel", func() {
|
||||
It("prefers the flag over everything", func() {
|
||||
got, err := ResolveModel(ModelRequest{
|
||||
Flag: "from-flag",
|
||||
Configured: "from-config",
|
||||
Available: []string{"a", "b"},
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("from-flag"))
|
||||
})
|
||||
|
||||
It("uses the configured model when no flag is given", func() {
|
||||
got, err := ResolveModel(ModelRequest{
|
||||
Configured: "from-config",
|
||||
Available: []string{"a", "b"},
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("from-config"))
|
||||
})
|
||||
|
||||
It("auto-selects when the server offers exactly one model", func() {
|
||||
got, err := ResolveModel(ModelRequest{Available: []string{"only-one"}})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("only-one"))
|
||||
})
|
||||
|
||||
It("errors and lists the options when several models exist and there is no chooser", func() {
|
||||
_, err := ResolveModel(ModelRequest{Available: []string{"a", "b"}})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("a"))
|
||||
Expect(err.Error()).To(ContainSubstring("b"))
|
||||
Expect(err.Error()).To(ContainSubstring("--model"))
|
||||
})
|
||||
|
||||
It("sorts before offering, so the same number means the same model next run", func() {
|
||||
var offered []string
|
||||
available := []string{"zeta", "alpha", "mid"}
|
||||
_, err := ResolveModel(ModelRequest{
|
||||
Available: available,
|
||||
StateDir: GinkgoT().TempDir(),
|
||||
Choose: func(models []string) (string, error) {
|
||||
offered = models
|
||||
return models[0], nil
|
||||
},
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
// The server's /v1/models ordering is unstable between calls.
|
||||
Expect(offered).To(Equal([]string{"alpha", "mid", "zeta"}))
|
||||
// Sorting must happen on a copy: the caller still owns this slice, and
|
||||
// reordering it under them would move whatever they index into it.
|
||||
Expect(available).To(Equal([]string{"zeta", "alpha", "mid"}))
|
||||
})
|
||||
|
||||
It("lists models in sorted order in the several-models error", func() {
|
||||
_, err := ResolveModel(ModelRequest{Available: []string{"zeta", "alpha"}})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("alpha, zeta"))
|
||||
})
|
||||
|
||||
It("asks the chooser when several models exist, and persists the answer", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
got, err := ResolveModel(ModelRequest{
|
||||
Available: []string{"a", "b"},
|
||||
StateDir: dir,
|
||||
Choose: func(models []string) (string, error) { return models[1], nil },
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("b"))
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("model: b"))
|
||||
})
|
||||
|
||||
// The answer is persisted and every later run starts against it, and
|
||||
// ModelChooser is exported, so the invariant has to hold for choosers this
|
||||
// package did not write.
|
||||
DescribeTable("refuses an answer the chooser was not offered",
|
||||
func(answer string) {
|
||||
dir := GinkgoT().TempDir()
|
||||
got, err := ResolveModel(ModelRequest{
|
||||
Available: []string{"alpha", "zeta"},
|
||||
StateDir: dir,
|
||||
Choose: func([]string) (string, error) { return answer, nil },
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(got).To(BeEmpty())
|
||||
Expect(err.Error()).To(ContainSubstring("alpha, zeta"))
|
||||
|
||||
_, statErr := os.Stat(ConfigPath(dir))
|
||||
Expect(os.IsNotExist(statErr)).To(BeTrue(), "nothing may be recorded for an answer that was refused")
|
||||
},
|
||||
Entry("nothing at all", ""),
|
||||
Entry("a model the server never offered", "gamma"),
|
||||
Entry("an offered model with stray whitespace", " alpha"),
|
||||
Entry("an offered model in the wrong case", "Alpha"),
|
||||
)
|
||||
|
||||
It("notifies, and still honours the choice, when it cannot be persisted", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
// A directory where the config file belongs: the write fails for any
|
||||
// user, including root.
|
||||
Expect(os.MkdirAll(ConfigPath(dir), 0o700)).To(Succeed())
|
||||
|
||||
var notices []string
|
||||
got, err := ResolveModel(ModelRequest{
|
||||
Available: []string{"a", "b"},
|
||||
StateDir: dir,
|
||||
Choose: func(models []string) (string, error) { return models[0], nil },
|
||||
Notify: func(message string) { notices = append(notices, message) },
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(got).To(Equal("a"))
|
||||
Expect(notices).To(HaveLen(1))
|
||||
Expect(notices[0]).To(ContainSubstring("a"))
|
||||
Expect(notices[0]).To(ContainSubstring("could not be saved"))
|
||||
})
|
||||
|
||||
It("says nothing when the choice was saved", func() {
|
||||
var notices []string
|
||||
_, err := ResolveModel(ModelRequest{
|
||||
Available: []string{"a", "b"},
|
||||
StateDir: GinkgoT().TempDir(),
|
||||
Choose: func(models []string) (string, error) { return models[0], nil },
|
||||
Notify: func(message string) { notices = append(notices, message) },
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(notices).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("propagates a chooser cancellation", func() {
|
||||
cancelled := errors.New("cancelled")
|
||||
_, err := ResolveModel(ModelRequest{
|
||||
Available: []string{"a", "b"},
|
||||
StateDir: GinkgoT().TempDir(),
|
||||
Choose: func([]string) (string, error) { return "", cancelled },
|
||||
})
|
||||
Expect(errors.Is(err, cancelled)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("errors with an install hint when the server has no models", func() {
|
||||
_, err := ResolveModel(ModelRequest{Available: nil})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("local-ai models install"))
|
||||
})
|
||||
})
|
||||
@@ -1,475 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/nib/app"
|
||||
nibcmd "github.com/mudler/nib/cmd"
|
||||
nibconfig "github.com/mudler/nib/config"
|
||||
nibtypes "github.com/mudler/nib/types"
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
// Options is everything the chat command passes down from its flags.
|
||||
type Options struct {
|
||||
Args []string // forwarded to the agent verbatim
|
||||
Endpoint string // the server root, e.g. http://127.0.0.1:8080
|
||||
BaseURL string // the API base, e.g. http://127.0.0.1:8080/v1
|
||||
APIKey string
|
||||
Model string
|
||||
StateDir string
|
||||
TraceDir string
|
||||
Yolo bool
|
||||
// ProbeTimeout bounds each check of the server. Zero means
|
||||
// defaultProbeTimeout.
|
||||
ProbeTimeout time.Duration
|
||||
|
||||
In io.Reader
|
||||
Out io.Writer
|
||||
ErrOut io.Writer
|
||||
}
|
||||
|
||||
// ExitStatus reports the status the process should exit with for an agent run
|
||||
// that failed, and whether err is such a failure.
|
||||
//
|
||||
// nib writes what went wrong to the error stream itself and hands back nothing
|
||||
// but a code, so an error that satisfies this has already been explained to the
|
||||
// user and must not be reported a second time. The refusal to open a
|
||||
// full-screen session on a stdin that cannot be read arrives this way, and it
|
||||
// is the one a user is most likely to meet: 'echo q | local-ai chat' names
|
||||
// --cli, and burying that under a second message would hide the fix.
|
||||
func ExitStatus(err error) (int, bool) {
|
||||
var exit app.ExitError
|
||||
if errors.As(err, &exit) {
|
||||
return exit.Code, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// shutdownSignals end the session. SIGHUP is one of them because this is a
|
||||
// terminal program: once the terminal is gone there is nobody left to talk to,
|
||||
// and a server started for the session has to go with it.
|
||||
var shutdownSignals = []os.Signal{os.Interrupt, syscall.SIGTERM, syscall.SIGHUP}
|
||||
|
||||
// shutdownContext derives a context that is cancelled when the process is
|
||||
// asked to stop.
|
||||
//
|
||||
// Without it a signal kills this process where it stands, skipping every
|
||||
// deferred call, and a 'local-ai run' started for the session is reparented to
|
||||
// init with nothing left that knows to shut it down. An interactive Ctrl+C is
|
||||
// safe on its own, because the child shares this process' foreground process
|
||||
// group and the terminal signals all of it, but a SIGTERM from a supervisor or
|
||||
// a script reaches only this process.
|
||||
//
|
||||
// Since nib v0.5.1 cancelling this context does end the session: RunTUI passes
|
||||
// it to bubbletea, which unwinds the program and reports the context's own
|
||||
// error. The server is still stopped on cancellation rather than on the way
|
||||
// out (see runSession), because registering here removes SIGHUP's default
|
||||
// terminate disposition, and a guarantee about a server this process owns is
|
||||
// not worth resting on how promptly a third party unwinds its interface.
|
||||
//
|
||||
// A handler rather than SysProcAttr.Pdeathsig on the child: Pdeathsig is
|
||||
// Linux-only, and in Go it is delivered when the OS thread that forked exits
|
||||
// rather than when the process does, so it can fire on a perfectly healthy
|
||||
// parent. Setpgid is not an alternative either, since taking the child out of
|
||||
// the foreground process group is what would break the Ctrl+C that works
|
||||
// today. SIGKILL stays uncovered, as it must: nothing in the process can
|
||||
// observe it.
|
||||
func shutdownContext(parent context.Context) (context.Context, context.CancelFunc) {
|
||||
return signal.NotifyContext(parent, shutdownSignals...)
|
||||
}
|
||||
|
||||
// Run starts the agent: resolve where state lives, make sure a server is
|
||||
// reachable, pick a model, then hand off to nib.
|
||||
func Run(ctx context.Context, opts Options) error {
|
||||
ctx, stop := shutdownContext(ctx)
|
||||
defer stop()
|
||||
|
||||
p, err := prepare(ctx, opts, isTerminal(opts.In))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// A server this process started belongs to this session, and Stop is
|
||||
// nil-safe and idempotent, so one defer covers both cases and costs nothing
|
||||
// when runSession has already stopped it.
|
||||
defer p.server.Stop()
|
||||
|
||||
return runSession(ctx, p.server, func(ctx context.Context) error {
|
||||
return runAgent(ctx, p.dir, p.model, opts)
|
||||
})
|
||||
}
|
||||
|
||||
// runSession hands the terminal to agent, and stops a server started for this
|
||||
// session as soon as the context is cancelled rather than when agent returns.
|
||||
//
|
||||
// The difference matters because the deferred Stop in Run is only reached once
|
||||
// agent returns, and how long that takes is nib's business rather than ours.
|
||||
// nib v0.5.1 does unwind the TUI on a cancelled context, so it does return; a
|
||||
// SIGHUP no longer leaves the interface on screen with the server behind it,
|
||||
// which it did before, when bubbletea's own SIGINT and SIGTERM handler was the
|
||||
// only thing that ever quit the program and registering for SIGHUP had removed
|
||||
// the default disposition that used to end the process. Watching the context
|
||||
// keeps the guarantee independent of what the agent does with it.
|
||||
func runSession(ctx context.Context, server *StartedServer, agent func(context.Context) error) error {
|
||||
returned := make(chan struct{})
|
||||
defer close(returned)
|
||||
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
server.Stop()
|
||||
case <-returned:
|
||||
}
|
||||
}()
|
||||
|
||||
return agent(ctx)
|
||||
}
|
||||
|
||||
// preparation is what the agent needs once the environment is ready: where its
|
||||
// state lives, which model to talk to, and the server this process started on
|
||||
// the user's behalf, if any.
|
||||
type preparation struct {
|
||||
dir string
|
||||
model string
|
||||
server *StartedServer
|
||||
}
|
||||
|
||||
// prepare does everything that has to happen before the agent takes over the
|
||||
// terminal. It is split out of Run because all of it is testable and none of
|
||||
// what follows is: once app.Run has the terminal there is no seam left.
|
||||
//
|
||||
// interactive says whether there is a user to prompt. It is a parameter rather
|
||||
// than a second read of opts.In so the prompts can be driven over a pipe.
|
||||
func prepare(ctx context.Context, opts Options, interactive bool) (_ *preparation, err error) {
|
||||
dir, dirErr := StateDir(opts.StateDir)
|
||||
if dirErr != nil {
|
||||
return nil, dirErr
|
||||
}
|
||||
if err := EnsureStateDir(dir, opts.BaseURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if isLocalOnlyArgs(opts.Args) {
|
||||
return &preparation{dir: dir}, nil
|
||||
}
|
||||
|
||||
// One prompter for every question this run asks; see its doc comment for
|
||||
// why the reader cannot be rebuilt per question.
|
||||
var prompts *prompter
|
||||
if interactive {
|
||||
prompts = newPrompter(opts.In, opts.ErrOut)
|
||||
}
|
||||
|
||||
var started *StartedServer
|
||||
defer func() {
|
||||
// Nothing after the spawn may leave a server behind: the caller only
|
||||
// learns about it through a successful return.
|
||||
if err != nil {
|
||||
started.Stop()
|
||||
}
|
||||
}()
|
||||
|
||||
models, err := probeModels(ctx, opts)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrUnauthorized) {
|
||||
return nil, fmt.Errorf("the LocalAI server at %s rejected the API key. Pass --api-key or set LOCALAI_API_KEY", opts.Endpoint)
|
||||
}
|
||||
if !errors.Is(err, ErrUnreachable) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var confirm Confirmer
|
||||
if interactive {
|
||||
confirm = prompts.yesNo
|
||||
}
|
||||
var startErr error
|
||||
started, startErr = OfferToStart(ctx, StartOptions{
|
||||
Endpoint: opts.Endpoint,
|
||||
Confirm: confirm,
|
||||
Stderr: opts.ErrOut,
|
||||
})
|
||||
if startErr != nil {
|
||||
err = startErr
|
||||
if errors.Is(startErr, ErrDeclined) {
|
||||
err = fmt.Errorf("no LocalAI server at %s. Start one with 'local-ai run', or point elsewhere with --endpoint", opts.Endpoint)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
say(opts.ErrOut, "Started a temporary LocalAI server; it stops when you exit. Use 'local-ai run' for a persistent one.\n")
|
||||
|
||||
if models, err = probeModels(ctx, opts); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
var chooser ModelChooser
|
||||
if interactive {
|
||||
chooser = prompts.choose
|
||||
}
|
||||
model, err := ResolveModel(ModelRequest{
|
||||
Flag: opts.Model,
|
||||
Configured: configuredModel(dir),
|
||||
Available: models,
|
||||
StateDir: dir,
|
||||
Choose: chooser,
|
||||
Notify: func(message string) { say(opts.ErrOut, "%s\n", message) },
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &preparation{dir: dir, model: model, server: started}, nil
|
||||
}
|
||||
|
||||
func runAgent(ctx context.Context, dir, model string, opts Options) error {
|
||||
return app.Run(ctx, agentOptions(dir, model, opts))
|
||||
}
|
||||
|
||||
// agentOptions builds the request handed to nib. It is split out of runAgent
|
||||
// because app.Run takes the terminal and cannot be called from a test, while
|
||||
// what is asked of it is exactly the part worth pinning.
|
||||
//
|
||||
// The stream fields are the interesting ones, and they are not symmetric.
|
||||
//
|
||||
// nib reads a non-nil stream as "the embedder wants this used", and refuses
|
||||
// every mode but --cli when such a stream is not a terminal, because the
|
||||
// full-screen interface renders on /dev/tty and would otherwise ignore it in
|
||||
// silence. Nil means "not injected": nib falls back to the process stream and
|
||||
// behaves as standalone nib does.
|
||||
//
|
||||
// Stdin is passed through as it comes. A piped or redirected stdin really is
|
||||
// ignored by the interface, so the refusal is the honest answer there, and it
|
||||
// is the one users meet: 'echo q | local-ai chat' says to re-run with --cli
|
||||
// rather than opening a full-screen session that will never read the question.
|
||||
//
|
||||
// Stdout is different, and the process stream is deliberately sent as nil. The
|
||||
// interface does write to stdout even when it is a pipe: that is the whole of
|
||||
// nib's shell-capture idiom, out=$(local-ai chat --height 50%), which is what
|
||||
// the Ctrl+Space widget emitted by --init is built on. Injecting os.Stdout
|
||||
// there would refuse the widget for a stream nib was going to use anyway.
|
||||
//
|
||||
// The test is identity with os.Stdout rather than whether it happens to be a
|
||||
// terminal, which means a shell redirect goes the same way as the widget:
|
||||
// 'local-ai chat > out.txt' no longer refuses either, and renders on /dev/tty
|
||||
// with the capture line landing in the file. That is not a second decision, it
|
||||
// is the same one. Both are the process stdout as the shell handed it over,
|
||||
// differing only in being a pipe rather than a regular file, which nib's gate
|
||||
// does not look at and should not. Refusing one would refuse the other.
|
||||
//
|
||||
// What stays injected, and so stays subject to the refusal, is a writer some
|
||||
// in-process caller chose for itself rather than inherited: a bytes.Buffer, or
|
||||
// an *os.File it opened. The specs rely on that.
|
||||
//
|
||||
// Stderr is never gated by nib, so it is passed through unchanged.
|
||||
//
|
||||
// The config values go through Overrides rather than Defaults, and that is not
|
||||
// a detail. Defaults are seeds: they sit BENEATH the config file, so the file
|
||||
// silently undoes them. Everything here is a decision this invocation already
|
||||
// made on the user's behalf, and a flag that the file can undo is not a flag.
|
||||
// It was not a rare case either, since EnsureStateDir writes base_url on the
|
||||
// first run and an interactive choice writes model, so from the second run on
|
||||
// the file carried a value for both and --endpoint and --model did nothing.
|
||||
//
|
||||
// The one asymmetry to plan around is that nib cannot tell "set to the zero
|
||||
// value" from "not set", so an override only ever raises a field. --yolo can
|
||||
// turn approval off, but nothing on the command line can turn it back on over
|
||||
// an approval_mode: auto in the file; that needs a config edit. Same shape for
|
||||
// the strings, which is what makes an unset --api-key or --trace-dir leave the
|
||||
// file's value standing, as it should.
|
||||
//
|
||||
// nib's own --trace-dir and --yolo, and their NIB_TRACE_DIR and NIB_YOLO twins,
|
||||
// are resolved after the config load and so still outrank these. That is
|
||||
// deliberate upstream: they are instructions to nib rather than ambient
|
||||
// environment.
|
||||
func agentOptions(dir, model string, opts Options) app.Options {
|
||||
// Model is the model this run resolved, which already prefers --model and
|
||||
// falls back to the file's own model, so the override restates the file's
|
||||
// value rather than fighting it whenever no flag was given.
|
||||
//
|
||||
// BaseURL is the endpoint this run probed, offered to start a server for,
|
||||
// and seeded the config with. Handing nib a different one is precisely the
|
||||
// split that made --endpoint a no-op, so the agent talks to the server
|
||||
// LocalAI checked. Pointing somewhere else for good is LOCALAI_CHAT_ENDPOINT
|
||||
// or --endpoint, not a hand-edited base_url the probe never reads.
|
||||
//
|
||||
// APIKey and TraceDir are the flags as given, empty when they were not, and
|
||||
// an empty override leaves the file alone. TraceDir is runtime-only in nib
|
||||
// (yaml:"-"), so no file value exists for it to beat today; it belongs here
|
||||
// with the other flags rather than one rung down for a reason that could
|
||||
// quietly stop being true.
|
||||
overrides := nibtypes.Config{
|
||||
Model: model,
|
||||
APIKey: opts.APIKey,
|
||||
BaseURL: opts.BaseURL,
|
||||
TraceDir: opts.TraceDir,
|
||||
}
|
||||
if opts.Yolo {
|
||||
overrides.ApprovalMode = "auto"
|
||||
}
|
||||
|
||||
return app.Options{
|
||||
Args: opts.Args,
|
||||
ProgramName: "local-ai chat",
|
||||
BaseDir: dir,
|
||||
Overrides: overrides,
|
||||
SkipSetup: true,
|
||||
SkipBareEnv: true,
|
||||
Stdin: opts.In,
|
||||
Stdout: ownStdout(opts.Out),
|
||||
Stderr: opts.ErrOut,
|
||||
}
|
||||
}
|
||||
|
||||
// ownStdout reports the writer as nib's own rather than as an injected one when
|
||||
// it is the process stdout, by answering nil for it. See agentOptions for why
|
||||
// that distinction is the difference between a working Ctrl+Space widget and a
|
||||
// refused one.
|
||||
func ownStdout(w io.Writer) io.Writer {
|
||||
if f, ok := w.(*os.File); ok && f == os.Stdout {
|
||||
return nil
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// defaultProbeTimeout bounds a check of the server. Listing models is cheap,
|
||||
// so this is long enough that a loaded server is never given up on and short
|
||||
// enough that a hung one does not leave the user staring at nothing.
|
||||
const defaultProbeTimeout = 30 * time.Second
|
||||
|
||||
// probeModels lists what the endpoint offers, under a budget.
|
||||
func probeModels(ctx context.Context, opts Options) ([]string, error) {
|
||||
timeout := opts.ProbeTimeout
|
||||
if timeout <= 0 {
|
||||
timeout = defaultProbeTimeout
|
||||
}
|
||||
// A real deadline rather than a cancel plus a timer. Probe reads
|
||||
// context.Canceled as "the caller gave up", which is a statement about the
|
||||
// caller and not about the endpoint, and only a deadline as "nothing
|
||||
// answered in time". Expiring the budget as a cancellation would stop
|
||||
// ErrUnreachable firing for precisely the hung servers that the offer to
|
||||
// start one exists for.
|
||||
probeCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
return Probe(probeCtx, opts.BaseURL, opts.APIKey)
|
||||
}
|
||||
|
||||
// isLocalOnlyArgs reports whether the forwarded arguments do their work
|
||||
// without ever reaching a model, in which case demanding a running server (and
|
||||
// offering to start one) would be an obstacle rather than a service.
|
||||
//
|
||||
// Two groups qualify. The management subcommands edit nib's own state: plugin,
|
||||
// skill, and the mcp verbs that add or remove configured servers, which is
|
||||
// asked of nib rather than restated, because bare 'mcp' and its transport
|
||||
// flags do serve the agent and do need a model. The other group is the flags
|
||||
// that only print something, above all --init: its shell snippet goes into an
|
||||
// rc file, typically long before any server exists.
|
||||
func isLocalOnlyArgs(args []string) bool {
|
||||
if len(args) == 0 {
|
||||
return false
|
||||
}
|
||||
// A scan rather than a look at args[0]: the mode flags this command
|
||||
// translates are prepended, so --init is not necessarily first. Positional
|
||||
// text cannot be mistaken for a flag here, since nib ignores what is left
|
||||
// after flag parsing.
|
||||
for _, a := range args {
|
||||
switch {
|
||||
case a == "--init", a == "-init", strings.HasPrefix(a, "--init="), strings.HasPrefix(a, "-init="):
|
||||
return true
|
||||
case a == "--version", a == "-version":
|
||||
return true
|
||||
}
|
||||
}
|
||||
switch args[0] {
|
||||
case "plugin", "skill":
|
||||
return true
|
||||
case "mcp":
|
||||
return len(args) >= 2 && nibcmd.IsMCPManageSubcommand(args[1])
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// configuredModel reads the model already recorded in the agent config, if any.
|
||||
func configuredModel(dir string) string {
|
||||
cfg := nibconfig.LoadWith(nibconfig.LoadOptions{BaseDir: dir, SkipBareEnv: true})
|
||||
return cfg.Model
|
||||
}
|
||||
|
||||
func isTerminal(in io.Reader) bool {
|
||||
f, ok := in.(*os.File)
|
||||
return ok && term.IsTerminal(int(f.Fd()))
|
||||
}
|
||||
|
||||
// say writes a line of interactive chatter: a question, or a notice about
|
||||
// something that did not stop the session. A write that fails is not worth
|
||||
// failing over, and when the terminal really is gone the read that follows the
|
||||
// question says so.
|
||||
func say(w io.Writer, format string, args ...any) {
|
||||
_, _ = fmt.Fprintf(w, format, args...)
|
||||
}
|
||||
|
||||
// prompter asks this run's questions on the user's terminal.
|
||||
//
|
||||
// It owns the buffered reader rather than wrapping opts.In per question,
|
||||
// because bufio reads ahead: a throwaway reader for the "start a server?"
|
||||
// question swallows the model choice that was typed behind it, and the next
|
||||
// question then sees EOF. A real run asks both, one after the other.
|
||||
type prompter struct {
|
||||
in *bufio.Reader
|
||||
out io.Writer
|
||||
}
|
||||
|
||||
func newPrompter(in io.Reader, out io.Writer) *prompter {
|
||||
return &prompter{in: bufio.NewReader(in), out: out}
|
||||
}
|
||||
|
||||
// yesNo satisfies Confirmer. Anything that is not an explicit yes is a no, so
|
||||
// a closed stream declines rather than proceeding on the user's behalf.
|
||||
func (p *prompter) yesNo(question string) (bool, error) {
|
||||
say(p.out, "%s [y/N]: ", question)
|
||||
line, err := p.in.ReadString('\n')
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return false, fmt.Errorf("reading the answer: %w", err)
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(line)) {
|
||||
case "y", "yes":
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// choose satisfies ModelChooser. It answers with a list index rather than with
|
||||
// what the user typed, so the result can only ever be one of the models it was
|
||||
// offered: a model name is not something to accept unvalidated here, since
|
||||
// ResolveModel persists whatever comes back and every later run then starts
|
||||
// against it.
|
||||
func (p *prompter) choose(models []string) (string, error) {
|
||||
if len(models) == 0 {
|
||||
return "", errors.New("there is nothing to choose from")
|
||||
}
|
||||
say(p.out, "Several models are available:\n")
|
||||
for i, m := range models {
|
||||
say(p.out, " %d) %s\n", i+1, m)
|
||||
}
|
||||
say(p.out, "Pick one [1-%d]: ", len(models))
|
||||
|
||||
line, err := p.in.ReadString('\n')
|
||||
if err != nil && !errors.Is(err, io.EOF) {
|
||||
return "", fmt.Errorf("reading the choice: %w", err)
|
||||
}
|
||||
answer := strings.TrimSpace(line)
|
||||
n, err := strconv.Atoi(answer)
|
||||
if err != nil || n < 1 || n > len(models) {
|
||||
return "", fmt.Errorf("not a valid choice: %q. Pick a number between 1 and %d, or pass --model", answer, len(models))
|
||||
}
|
||||
return models[n-1], nil
|
||||
}
|
||||
@@ -1,629 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/nib/app"
|
||||
nibconfig "github.com/mudler/nib/config"
|
||||
nibtypes "github.com/mudler/nib/types"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
// modelServer answers /v1/models with the given ids, as LocalAI does.
|
||||
func modelServer(ids ...string) *httptest.Server {
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
data := make([]map[string]string, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
data = append(data, map[string]string{"id": id, "object": "model"})
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
Expect(json.NewEncoder(w).Encode(map[string]any{"object": "list", "data": data})).To(Succeed())
|
||||
}))
|
||||
}
|
||||
|
||||
var _ = Describe("prepare", func() {
|
||||
var (
|
||||
dir string
|
||||
errOut *bytes.Buffer
|
||||
)
|
||||
|
||||
BeforeEach(func() {
|
||||
dir = GinkgoT().TempDir()
|
||||
errOut = &bytes.Buffer{}
|
||||
})
|
||||
|
||||
// optionsFor points a run at srv, with no input to read: the default is a
|
||||
// session nobody can be asked anything in.
|
||||
optionsFor := func(srv *httptest.Server) Options {
|
||||
endpoint := "http://127.0.0.1:0"
|
||||
base := endpoint + "/v1"
|
||||
if srv != nil {
|
||||
endpoint, base = srv.URL, srv.URL+"/v1"
|
||||
}
|
||||
return Options{
|
||||
Endpoint: endpoint,
|
||||
BaseURL: base,
|
||||
StateDir: dir,
|
||||
In: strings.NewReader(""),
|
||||
Out: &bytes.Buffer{},
|
||||
ErrOut: errOut,
|
||||
}
|
||||
}
|
||||
|
||||
It("uses the only model the server offers", func() {
|
||||
srv := modelServer("the-only-model")
|
||||
defer srv.Close()
|
||||
|
||||
p, err := prepare(context.Background(), optionsFor(srv), false)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(p.model).To(Equal("the-only-model"))
|
||||
Expect(p.dir).To(Equal(dir))
|
||||
Expect(p.server).To(BeNil(), "nothing was started, so nothing is owned")
|
||||
})
|
||||
|
||||
It("seeds the agent config with the endpoint on first run", func() {
|
||||
srv := modelServer("m")
|
||||
defer srv.Close()
|
||||
|
||||
_, err := prepare(context.Background(), optionsFor(srv), false)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring(srv.URL + "/v1"))
|
||||
})
|
||||
|
||||
It("lets --model win over what the server offers", func() {
|
||||
srv := modelServer("a", "b")
|
||||
defer srv.Close()
|
||||
|
||||
opts := optionsFor(srv)
|
||||
opts.Model = "not-listed-yet"
|
||||
p, err := prepare(context.Background(), opts, false)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(p.model).To(Equal("not-listed-yet"))
|
||||
})
|
||||
|
||||
It("advises about the API key when the server rejects it", func() {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, err := prepare(context.Background(), optionsFor(srv), false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("--api-key"))
|
||||
Expect(err.Error()).To(ContainSubstring(srv.URL))
|
||||
})
|
||||
|
||||
// Not interactive means nobody can answer the offer, so the advice has to
|
||||
// stand on its own.
|
||||
It("advises how to start a server when none is reachable", func() {
|
||||
srv := modelServer()
|
||||
url := srv.URL
|
||||
srv.Close() // nothing is listening now
|
||||
|
||||
opts := optionsFor(nil)
|
||||
opts.Endpoint, opts.BaseURL = url, url+"/v1"
|
||||
_, err := prepare(context.Background(), opts, false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("local-ai run"))
|
||||
Expect(err.Error()).To(ContainSubstring(url))
|
||||
})
|
||||
|
||||
// A server that accepts the connection and then never replies is the case
|
||||
// the offer to start one exists for, so the budget has to expire as a
|
||||
// deadline: Probe reads a cancellation as "the caller gave up" and refuses
|
||||
// to call the endpoint unreachable on the strength of it.
|
||||
It("treats a server that never answers as one that is not there", func(ctx SpecContext) {
|
||||
release := make(chan struct{})
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
select {
|
||||
case <-release:
|
||||
case <-r.Context().Done():
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
defer close(release)
|
||||
|
||||
opts := optionsFor(srv)
|
||||
opts.ProbeTimeout = 100 * time.Millisecond
|
||||
_, err := prepare(context.Background(), opts, false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("local-ai run"), "want the offer-a-server advice, got %v", err)
|
||||
}, SpecTimeout(30*time.Second))
|
||||
|
||||
It("asks which model to use and remembers the answer", func() {
|
||||
srv := modelServer("zeta", "alpha")
|
||||
defer srv.Close()
|
||||
|
||||
opts := optionsFor(srv)
|
||||
opts.In = strings.NewReader("2\n")
|
||||
p, err := prepare(context.Background(), opts, true)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
// The list is sorted before it is shown, so 2 is zeta, not the second
|
||||
// thing the server happened to name.
|
||||
Expect(p.model).To(Equal("zeta"))
|
||||
Expect(errOut.String()).To(ContainSubstring("1) alpha"))
|
||||
Expect(errOut.String()).To(ContainSubstring("2) zeta"))
|
||||
|
||||
data, err := os.ReadFile(ConfigPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(data)).To(ContainSubstring("zeta"))
|
||||
})
|
||||
|
||||
// The choice is prompted for once and remembered. When remembering it fails
|
||||
// the user is about to be asked again on every future run, so they have to
|
||||
// be told here: a log line is invisible at the default log level.
|
||||
It("says so on the prompt when the choice cannot be remembered", func() {
|
||||
srv := modelServer("zeta", "alpha")
|
||||
defer srv.Close()
|
||||
|
||||
// A directory where the config file belongs: writable state dir,
|
||||
// unwritable config, on any platform and as any user.
|
||||
Expect(os.MkdirAll(ConfigPath(dir), 0o700)).To(Succeed())
|
||||
|
||||
opts := optionsFor(srv)
|
||||
opts.In = strings.NewReader("1\n")
|
||||
p, err := prepare(context.Background(), opts, true)
|
||||
|
||||
// Failing to remember the choice must not cost the user their session.
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(p.model).To(Equal("alpha"))
|
||||
Expect(errOut.String()).To(ContainSubstring("could not be saved"), "the user has to learn they will be asked again")
|
||||
})
|
||||
|
||||
It("does not ask again once a model is recorded", func() {
|
||||
srv := modelServer("zeta", "alpha")
|
||||
defer srv.Close()
|
||||
|
||||
Expect(PersistModel(dir, "alpha")).To(Succeed())
|
||||
|
||||
opts := optionsFor(srv)
|
||||
opts.In = strings.NewReader("") // an answer would have nothing to read
|
||||
p, err := prepare(context.Background(), opts, true)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(p.model).To(Equal("alpha"))
|
||||
Expect(errOut.String()).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("says what to install when the server has no models", func() {
|
||||
srv := modelServer()
|
||||
defer srv.Close()
|
||||
|
||||
_, err := prepare(context.Background(), optionsFor(srv), false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("models install"))
|
||||
})
|
||||
|
||||
Describe("arguments that only touch local state", func() {
|
||||
unreachable := func(args ...string) Options {
|
||||
opts := optionsFor(nil) // port 0: nothing can ever answer here
|
||||
opts.Args = args
|
||||
return opts
|
||||
}
|
||||
|
||||
DescribeTable("skips the server entirely",
|
||||
func(args ...string) {
|
||||
p, err := prepare(context.Background(), unreachable(args...), false)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(p.model).To(BeEmpty())
|
||||
Expect(p.server).To(BeNil())
|
||||
},
|
||||
Entry("plugin", "plugin", "list"),
|
||||
Entry("skill", "skill", "list"),
|
||||
Entry("mcp add", "mcp", "add", "srv"),
|
||||
Entry("mcp list", "mcp", "list"),
|
||||
// The shell snippet is what a user puts in their rc file, long
|
||||
// before any server exists.
|
||||
Entry("the shell integration script", "--init", "zsh"),
|
||||
Entry("the version", "--version"),
|
||||
)
|
||||
|
||||
// Bare 'mcp' and its transport flags serve the agent over MCP, so they
|
||||
// need a model like any other session. Only the verbs that edit the
|
||||
// configured servers are local.
|
||||
DescribeTable("still needs a server",
|
||||
func(args ...string) {
|
||||
_, err := prepare(context.Background(), unreachable(args...), false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("local-ai run"))
|
||||
},
|
||||
Entry("mcp over stdio", "mcp", "--stdio"),
|
||||
Entry("bare mcp", "mcp"),
|
||||
)
|
||||
})
|
||||
|
||||
// A reader per question would read ahead into a buffer it then discards, so
|
||||
// the second question would see EOF whenever both answers were typed ahead.
|
||||
// That is the shape of a real run: the offer to start a server is followed
|
||||
// by the model prompt.
|
||||
It("keeps reading answers from the same stream across questions", func() {
|
||||
out := &bytes.Buffer{}
|
||||
p := newPrompter(strings.NewReader("y\n2\n"), out)
|
||||
|
||||
yes, err := p.yesNo("Start one now?")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(yes).To(BeTrue())
|
||||
|
||||
chosen, err := p.choose([]string{"alpha", "zeta"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(chosen).To(Equal("zeta"))
|
||||
})
|
||||
|
||||
// Whatever the chooser returns is persisted and used for every later run,
|
||||
// so an answer that is not one of the offered models must never come back
|
||||
// as one.
|
||||
Describe("the model prompt", func() {
|
||||
offered := []string{"alpha", "zeta"}
|
||||
|
||||
DescribeTable("refuses an answer that is not one of the numbers shown",
|
||||
func(answer string) {
|
||||
chosen, err := newPrompter(strings.NewReader(answer), &bytes.Buffer{}).choose(offered)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(chosen).To(BeEmpty())
|
||||
},
|
||||
Entry("nothing at all", ""),
|
||||
Entry("a blank line", "\n"),
|
||||
Entry("only spaces", " \n"),
|
||||
Entry("zero", "0\n"),
|
||||
Entry("past the end", "3\n"),
|
||||
Entry("negative", "-1\n"),
|
||||
Entry("a model name", "zeta\n"),
|
||||
Entry("a number with a suffix", "1x\n"),
|
||||
)
|
||||
|
||||
It("says how to answer when the answer was not a number", func() {
|
||||
_, err := newPrompter(strings.NewReader("banana\n"), &bytes.Buffer{}).choose(offered)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("between 1 and 2"))
|
||||
Expect(err.Error()).To(ContainSubstring("--model"))
|
||||
})
|
||||
|
||||
It("returns the model shown against the number", func() {
|
||||
chosen, err := newPrompter(strings.NewReader("1\n"), &bytes.Buffer{}).choose(offered)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(chosen).To(Equal("alpha"))
|
||||
})
|
||||
|
||||
It("refuses to ask when there is nothing to offer", func() {
|
||||
chosen, err := newPrompter(strings.NewReader("1\n"), &bytes.Buffer{}).choose(nil)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(chosen).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
// A server started for this session is stopped by a deferred call, which a
|
||||
// signal skips: the process dies where it stands and leaves 'local-ai run'
|
||||
// reparented to init.
|
||||
Describe("shutdown signals", func() {
|
||||
It("ends the session when the terminal goes away", func() {
|
||||
ctx, stop := shutdownContext(context.Background())
|
||||
defer stop()
|
||||
|
||||
self, err := os.FindProcess(os.Getpid())
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(self.Signal(syscall.SIGHUP)).To(Succeed())
|
||||
|
||||
Eventually(ctx.Done()).WithTimeout(5 * time.Second).Should(BeClosed())
|
||||
Expect(ctx.Err()).To(MatchError(context.Canceled))
|
||||
})
|
||||
|
||||
// SIGINT and SIGTERM cannot be delivered here to prove the same thing:
|
||||
// Ginkgo registers for both to abort the suite, and a signal goes to
|
||||
// every registered listener.
|
||||
It("also listens for an interrupt and a terminate", func() {
|
||||
Expect(shutdownSignals).To(ContainElements(os.Signal(os.Interrupt), os.Signal(syscall.SIGTERM)))
|
||||
})
|
||||
})
|
||||
|
||||
// Cancelling the context does unwind nib's TUI since v0.5.1, but how long
|
||||
// that takes is nib's business, and the deferred Stop in Run is only reached
|
||||
// once the agent returns. A server this process started is ours to end, so
|
||||
// the guarantee is made here instead, where it does not depend on the agent
|
||||
// at all. Before v0.5.1 there was no guarantee to be had on the SIGHUP path:
|
||||
// bubbletea's own SIGINT and SIGTERM handler was the only thing that ever
|
||||
// quit the program, and registering for SIGHUP took away the default
|
||||
// disposition that used to end the process.
|
||||
Describe("runSession", func() {
|
||||
It("stops the session's server on cancellation, without waiting for the agent", func() {
|
||||
server, proc := stoppableServer()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
err := runSession(ctx, server, func(ctx context.Context) error {
|
||||
cancel()
|
||||
Eventually(func() int32 { return proc.interrupts.Load() }).
|
||||
WithTimeout(5 * time.Second).
|
||||
Should(BeNumerically(">", 0), "the server has to be stopped while the agent is still running")
|
||||
return nil
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(proc.lastSignal.Load()).To(Equal(os.Interrupt))
|
||||
})
|
||||
|
||||
It("leaves the server alone for as long as the session lasts", func() {
|
||||
server, proc := stoppableServer()
|
||||
|
||||
Expect(runSession(context.Background(), server, func(context.Context) error {
|
||||
return nil
|
||||
})).To(Succeed())
|
||||
Expect(proc.interrupts.Load()).To(BeZero())
|
||||
Expect(proc.kills.Load()).To(BeZero())
|
||||
})
|
||||
|
||||
It("returns what the agent returned", func() {
|
||||
failed := errors.New("the agent gave up")
|
||||
server, _ := stoppableServer()
|
||||
|
||||
Expect(runSession(context.Background(), server, func(context.Context) error {
|
||||
return failed
|
||||
})).To(MatchError(failed))
|
||||
})
|
||||
|
||||
// Most sessions run against a server the user already had, and there is
|
||||
// nothing to stop then.
|
||||
It("copes with a session that started no server", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
Expect(runSession(ctx, nil, func(context.Context) error {
|
||||
return nil
|
||||
})).To(Succeed())
|
||||
})
|
||||
})
|
||||
|
||||
// Which streams reach nib decides two user-visible behaviours at once, and
|
||||
// they pull in opposite directions, so both are pinned here rather than left
|
||||
// to whoever next edits the literal.
|
||||
//
|
||||
// nib refuses every mode but --cli when a stream it was handed is not a
|
||||
// terminal. That refusal is wanted for stdin, where it is what tells someone
|
||||
// piping a question to re-run with --cli. It is not wanted for the process
|
||||
// stdout, where it would refuse the Ctrl+Space widget that --init emits:
|
||||
// out=$(local-ai chat --height 50%) puts a pipe on stdout by construction,
|
||||
// and writing the chosen command into that pipe is the entire point.
|
||||
Describe("agentOptions", func() {
|
||||
// optionsWithStreams is a request that differs from the next only in
|
||||
// what it was told to read and write.
|
||||
optionsWithStreams := func(in io.Reader, out, errOut io.Writer) Options {
|
||||
return Options{
|
||||
BaseURL: "http://127.0.0.1:8080/v1",
|
||||
In: in,
|
||||
Out: out,
|
||||
ErrOut: errOut,
|
||||
}
|
||||
}
|
||||
|
||||
Describe("stdout", func() {
|
||||
// The regression this exists to catch: reinstating
|
||||
// 'Stdout: opts.Out' breaks Ctrl+Space and nothing else notices.
|
||||
It("hands nib nothing for the process stdout, so the capture widget is not refused", func() {
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr))
|
||||
Expect(o.Stdout).To(BeNil(), "injecting os.Stdout is what refuses out=$(local-ai chat)")
|
||||
})
|
||||
|
||||
It("keeps a stdout the caller chose, which the refusal still guards", func() {
|
||||
out := &bytes.Buffer{}
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, out, os.Stderr))
|
||||
Expect(o.Stdout).To(BeIdenticalTo(out))
|
||||
})
|
||||
|
||||
// Being an *os.File is not what makes a stream nib's own; being the
|
||||
// process stdout is. This is a file an in-process caller opened for
|
||||
// itself, not one a shell redirect handed over as stdout, which
|
||||
// still arrives as os.Stdout and is still nil-ed. It was never going
|
||||
// to receive the interface, so it stays injected and stays refused.
|
||||
It("keeps a file that is not the process stdout", func() {
|
||||
f, err := os.CreateTemp(GinkgoT().TempDir(), "captured")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(f.Close)
|
||||
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, f, os.Stderr))
|
||||
Expect(o.Stdout).To(BeIdenticalTo(f))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("stdin", func() {
|
||||
// The opposite regression: nilling stdin the way stdout is nilled
|
||||
// would silently drop the refusal that names --cli.
|
||||
It("hands the process stdin over, so a piped session is still refused", func() {
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr))
|
||||
Expect(o.Stdin).To(BeIdenticalTo(os.Stdin))
|
||||
})
|
||||
|
||||
It("hands over a stdin the caller chose", func() {
|
||||
in := strings.NewReader("a question")
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(in, os.Stdout, os.Stderr))
|
||||
Expect(o.Stdin).To(BeIdenticalTo(in))
|
||||
})
|
||||
})
|
||||
|
||||
// nib gates stdin and stdout and nothing else, so there is no reason to
|
||||
// hide the error stream from it.
|
||||
It("hands the error stream over whatever it is", func() {
|
||||
errOut := &bytes.Buffer{}
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, errOut))
|
||||
Expect(o.Stderr).To(BeIdenticalTo(errOut))
|
||||
|
||||
o = agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr))
|
||||
Expect(o.Stderr).To(BeIdenticalTo(os.Stderr))
|
||||
})
|
||||
|
||||
It("names the command a user would type, not the binary nib ships as", func() {
|
||||
o := agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr))
|
||||
Expect(o.ProgramName).To(Equal("local-ai chat"),
|
||||
"the --init widget invokes this name, so a user has to be able to run it")
|
||||
})
|
||||
|
||||
It("carries the resolved session through to nib", func() {
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
opts.Args = []string{"--cli"}
|
||||
opts.APIKey = "a-key"
|
||||
opts.TraceDir = "/traces"
|
||||
|
||||
o := agentOptions(dir, "the-model", opts)
|
||||
Expect(o.Args).To(Equal([]string{"--cli"}))
|
||||
Expect(o.BaseDir).To(Equal(dir))
|
||||
Expect(o.Overrides.Model).To(Equal("the-model"))
|
||||
Expect(o.Overrides.APIKey).To(Equal("a-key"))
|
||||
Expect(o.Overrides.BaseURL).To(Equal("http://127.0.0.1:8080/v1"))
|
||||
Expect(o.Overrides.TraceDir).To(Equal("/traces"))
|
||||
// The model and the server are settled before nib starts, and the
|
||||
// bare MODEL and API_KEY variables belong to some other tool.
|
||||
Expect(o.SkipSetup).To(BeTrue())
|
||||
Expect(o.SkipBareEnv).To(BeTrue())
|
||||
})
|
||||
|
||||
// Defaults sit beneath the config file. Anything routed through them is
|
||||
// accepted from the command line and then thrown away the moment the
|
||||
// file carries the same key, which is the normal state rather than an
|
||||
// edge case. Nothing this command resolves belongs there, so the channel
|
||||
// stays empty and this says so: it is what fails if the block is moved
|
||||
// back a rung.
|
||||
It("seeds nothing, because a seed is not a flag", func() {
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
opts.APIKey = "a-key"
|
||||
opts.TraceDir = "/traces"
|
||||
opts.Yolo = true
|
||||
|
||||
Expect(agentOptions(dir, "the-model", opts).Defaults).To(Equal(nibtypes.Config{}),
|
||||
"Defaults lose to the config file, so a value placed there is a flag that does nothing")
|
||||
})
|
||||
|
||||
It("asks for automatic approval only when --yolo was given", func() {
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
Expect(agentOptions(dir, "a-model", opts).Overrides.ApprovalMode).To(BeEmpty())
|
||||
|
||||
opts.Yolo = true
|
||||
Expect(agentOptions(dir, "a-model", opts).Overrides.ApprovalMode).To(Equal("auto"))
|
||||
})
|
||||
|
||||
// The specs above pin what is handed over. These pin what nib does with
|
||||
// it, which is the part that was wrong: every value below reached
|
||||
// app.Options intact and was then discarded by the config load, so a
|
||||
// spec that stops at the struct cannot see the bug. Resolving the config
|
||||
// the way app.Run resolves it can.
|
||||
Describe("the config nib actually resolves", func() {
|
||||
// writeConfig puts a config file where nib will read it, with values
|
||||
// that disagree with every flag under test.
|
||||
writeConfig := func(body string) {
|
||||
Expect(os.WriteFile(ConfigPath(dir), []byte(body), 0o600)).To(Succeed())
|
||||
}
|
||||
|
||||
// resolve loads the config exactly as app.Run does, so the precedence
|
||||
// under test is nib's own rather than a restatement of it here.
|
||||
resolve := func(o app.Options) nibtypes.Config {
|
||||
return nibconfig.LoadWith(nibconfig.LoadOptions{
|
||||
BaseDir: o.BaseDir,
|
||||
Defaults: o.Defaults,
|
||||
Overrides: o.Overrides,
|
||||
SkipBareEnv: o.SkipBareEnv,
|
||||
})
|
||||
}
|
||||
|
||||
It("sends the requests to the endpoint the flag named, not the one on disk", func() {
|
||||
writeConfig("base_url: http://127.0.0.1:9999/v1\n")
|
||||
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
opts.BaseURL = "http://127.0.0.1:8080/v1"
|
||||
|
||||
cfg := resolve(agentOptions(dir, "a-model", opts))
|
||||
Expect(cfg.BaseURL).To(Equal("http://127.0.0.1:8080/v1"),
|
||||
"--endpoint probed 8080; every turn has to go there too")
|
||||
})
|
||||
|
||||
It("uses the model the flag named, not the one the picker recorded", func() {
|
||||
writeConfig("model: recorded-model\n")
|
||||
|
||||
cfg := resolve(agentOptions(dir, "flag-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)))
|
||||
Expect(cfg.Model).To(Equal("flag-model"))
|
||||
})
|
||||
|
||||
It("uses the key the flag named, not the one nib saved", func() {
|
||||
writeConfig("api_key: saved-key\n")
|
||||
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
opts.APIKey = "flag-key"
|
||||
|
||||
cfg := resolve(agentOptions(dir, "a-model", opts))
|
||||
Expect(cfg.APIKey).To(Equal("flag-key"))
|
||||
})
|
||||
|
||||
It("turns approval off for --yolo even when the file demands it", func() {
|
||||
writeConfig("approval_mode: prompt\n")
|
||||
|
||||
opts := optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)
|
||||
opts.Yolo = true
|
||||
|
||||
cfg := resolve(agentOptions(dir, "a-model", opts))
|
||||
Expect(cfg.ApprovalMode).To(Equal("auto"))
|
||||
})
|
||||
|
||||
// The other half of the same rule, and the reason an unset flag is
|
||||
// not a demand for the empty string: an override only ever raises a
|
||||
// field, so what the user configured survives a run that said
|
||||
// nothing about it.
|
||||
It("leaves what the file configured alone when no flag was given", func() {
|
||||
writeConfig("api_key: saved-key\napproval_mode: prompt\n")
|
||||
|
||||
cfg := resolve(agentOptions(dir, "a-model", optionsWithStreams(os.Stdin, os.Stdout, os.Stderr)))
|
||||
Expect(cfg.APIKey).To(Equal("saved-key"))
|
||||
Expect(cfg.ApprovalMode).To(Equal("prompt"))
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
// nib reports its own failures on the error stream and returns nothing but
|
||||
// a status, so anything that reaches here as one has already been explained
|
||||
// once. The refusal to open a full-screen session on a stdin that cannot be
|
||||
// read is the one users meet: 'echo q | local-ai chat' names --cli, and a
|
||||
// second message on top would bury the fix.
|
||||
Describe("ExitStatus", func() {
|
||||
It("recognises a status the agent already explained", func() {
|
||||
code, reported := ExitStatus(app.ExitError{Code: 2})
|
||||
Expect(reported).To(BeTrue())
|
||||
Expect(code).To(Equal(2))
|
||||
})
|
||||
|
||||
It("finds one that has been wrapped", func() {
|
||||
code, reported := ExitStatus(fmt.Errorf("running the agent: %w", app.ExitError{Code: 1}))
|
||||
Expect(reported).To(BeTrue())
|
||||
Expect(code).To(Equal(1))
|
||||
})
|
||||
|
||||
It("leaves an ordinary failure to be reported", func() {
|
||||
_, reported := ExitStatus(errors.New("no LocalAI server at http://127.0.0.1:8080"))
|
||||
Expect(reported).To(BeFalse())
|
||||
})
|
||||
|
||||
It("says nothing about a run that succeeded", func() {
|
||||
_, reported := ExitStatus(nil)
|
||||
Expect(reported).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
It("reports a state dir it cannot create", func() {
|
||||
blocked := filepath.Join(dir, "a-file")
|
||||
Expect(os.WriteFile(blocked, []byte("not a dir"), 0o600)).To(Succeed())
|
||||
|
||||
opts := optionsFor(nil)
|
||||
opts.StateDir = filepath.Join(blocked, "chat")
|
||||
_, err := prepare(context.Background(), opts, false)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("agent state dir"))
|
||||
})
|
||||
})
|
||||
@@ -1,276 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
)
|
||||
|
||||
// ErrDeclined means no server was started, either because the session is not
|
||||
// interactive or because the user said no.
|
||||
var ErrDeclined = errors.New("no server started")
|
||||
|
||||
// errServerExited means the process we spawned died before it ever reported
|
||||
// ready, so there is no point in polling out the rest of the budget.
|
||||
var errServerExited = errors.New("the LocalAI server exited before it became ready")
|
||||
|
||||
const (
|
||||
// defaultReadyTimeout bounds the wait for a freshly spawned server. A cold
|
||||
// start probes hardware and may pull a backend, so the budget is generous.
|
||||
defaultReadyTimeout = 2 * time.Minute
|
||||
// readyPollInterval is how long to wait between readiness polls.
|
||||
readyPollInterval = 500 * time.Millisecond
|
||||
// readyProbeTimeout bounds a single readiness request, so one connection
|
||||
// that hangs cannot swallow the whole budget.
|
||||
readyProbeTimeout = 5 * time.Second
|
||||
// shutdownGrace is how long a server we started gets to unload models and
|
||||
// stop its backends after SIGINT before it is killed outright.
|
||||
shutdownGrace = 10 * time.Second
|
||||
// childOutputDrainDelay bounds how long cmd.Wait keeps copying the child's
|
||||
// output after the child itself has exited.
|
||||
//
|
||||
// This is not a theoretical guard for LocalAI. 'local-ai run' spawns backend
|
||||
// subprocesses, and they inherit the write end of the pipe exec created for
|
||||
// the child's stderr. A backend that outlives its parent holds that pipe
|
||||
// open, so an unbounded cmd.Wait would block on the copy goroutine long
|
||||
// after the server itself is gone: exited would never close, Stop would burn
|
||||
// its whole grace period even on a clean shutdown, and the waiter goroutine
|
||||
// would leak.
|
||||
//
|
||||
// The value is long enough that a legitimate final burst of logs is never
|
||||
// truncated even on a loaded machine, where the copy itself takes
|
||||
// microseconds. It must stay strictly below shutdownGrace: at or above it,
|
||||
// every wedged-pipe shutdown would exhaust the grace period and then SIGKILL
|
||||
// a process that had already exited cleanly.
|
||||
childOutputDrainDelay = 5 * time.Second
|
||||
)
|
||||
|
||||
// Confirmer asks a yes/no question. Nil means the session is not interactive.
|
||||
type Confirmer func(question string) (bool, error)
|
||||
|
||||
// StartOptions configures OfferToStart.
|
||||
type StartOptions struct {
|
||||
// Endpoint is the address the user expected a server on, used in the
|
||||
// question and polled for readiness. This is the endpoint root, not the
|
||||
// /v1 API base URL: readiness is served at the root.
|
||||
Endpoint string
|
||||
// Confirm asks whether to start a server. Nil means never start.
|
||||
Confirm Confirmer
|
||||
// Stderr receives the child's output.
|
||||
Stderr io.Writer
|
||||
// Executable overrides the binary to run. Empty means os.Executable().
|
||||
Executable string
|
||||
// ReadyTimeout bounds the wait for readiness. Zero means defaultReadyTimeout.
|
||||
ReadyTimeout time.Duration
|
||||
}
|
||||
|
||||
// StartedServer is a server this process started and is responsible for.
|
||||
type StartedServer struct {
|
||||
// exited is closed once the child has been reaped. One background waiter
|
||||
// owns cmd.Wait: it may only be called once, and it is what closes the
|
||||
// pipes exec created for Stdout/Stderr and joins the goroutines copying
|
||||
// them, so calling os.Process.Wait directly instead would leak both.
|
||||
exited chan struct{}
|
||||
// waitErr is the child's exit status. It is written before exited is
|
||||
// closed and must only be read after that channel is observed closed.
|
||||
waitErr error
|
||||
|
||||
// proc is the child. It is an interface rather than *os.Process so that
|
||||
// Stop's contract, in particular that the child is asked to stop exactly
|
||||
// once however often Stop is called, can be pinned without a live process
|
||||
// to signal. Nil means nothing was ever started.
|
||||
proc processControl
|
||||
|
||||
stopOnce sync.Once
|
||||
}
|
||||
|
||||
// processControl is the part of *os.Process that Stop needs.
|
||||
//
|
||||
// One interface rather than a pair of independent function fields: two fields
|
||||
// can be wired to each other's operation, or one left nil, and no test can tell,
|
||||
// because a fake satisfies any combination. There is nothing to swap or forget
|
||||
// here, since the sole implementation is the real process and the method names
|
||||
// carry the meaning.
|
||||
type processControl interface {
|
||||
Signal(os.Signal) error
|
||||
Kill() error
|
||||
}
|
||||
|
||||
// *os.Process satisfies processControl unmodified, so production needs no
|
||||
// adapter and no nil branch: the wiring is a single assignment.
|
||||
var _ processControl = (*os.Process)(nil)
|
||||
|
||||
// newServerCommand builds the child process. Split out from OfferToStart so the
|
||||
// process' configuration can be asserted on without spawning anything.
|
||||
func newServerCommand(bin string, stderr io.Writer) *exec.Cmd {
|
||||
cmd := exec.Command(bin, "run")
|
||||
// Stdin is left nil, so the child gets /dev/null: it is a background
|
||||
// server, and sharing the terminal would have it stealing keystrokes from
|
||||
// the agent.
|
||||
cmd.Stdout = stderr // the child's logs are diagnostics, not chat output
|
||||
cmd.Stderr = stderr
|
||||
// Bound the wait for the child's output pipes; see childOutputDrainDelay.
|
||||
cmd.WaitDelay = childOutputDrainDelay
|
||||
return cmd
|
||||
}
|
||||
|
||||
// OfferToStart asks whether to start a LocalAI server and, if allowed, spawns
|
||||
// one and waits for it to report ready.
|
||||
//
|
||||
// A child process rather than an in-process boot: RunCMD.Run installs its own
|
||||
// signal handling and blocks until shutdown, so re-entering it from a chat
|
||||
// session would entangle two lifecycles in one process.
|
||||
func OfferToStart(ctx context.Context, opts StartOptions) (*StartedServer, error) {
|
||||
if opts.Confirm == nil {
|
||||
// Not interactive. Spawning a server nobody asked for is the one thing
|
||||
// this function must never do: in CI, in a pipeline, or under a
|
||||
// supervisor there is no one to see it or shut it down.
|
||||
return nil, ErrDeclined
|
||||
}
|
||||
ok, err := opts.Confirm(fmt.Sprintf("No LocalAI server at %s. Start one now?", opts.Endpoint))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("asking whether to start a server: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return nil, ErrDeclined
|
||||
}
|
||||
|
||||
bin := opts.Executable
|
||||
if bin == "" {
|
||||
if bin, err = os.Executable(); err != nil {
|
||||
return nil, fmt.Errorf("locating the local-ai binary: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
cmd := newServerCommand(bin, opts.Stderr)
|
||||
if err := cmd.Start(); err != nil {
|
||||
return nil, fmt.Errorf("starting a LocalAI server with %s: %w", bin, err)
|
||||
}
|
||||
|
||||
s := &StartedServer{exited: make(chan struct{}), proc: cmd.Process}
|
||||
go func() {
|
||||
s.waitErr = cmd.Wait()
|
||||
close(s.exited)
|
||||
}()
|
||||
|
||||
timeout := opts.ReadyTimeout
|
||||
if timeout <= 0 {
|
||||
timeout = defaultReadyTimeout
|
||||
}
|
||||
if err := waitReady(ctx, opts.Endpoint, timeout, s.exited); err != nil {
|
||||
if errors.Is(err, errServerExited) {
|
||||
// Safe to read: errServerExited is only returned once exited has
|
||||
// been observed closed, which happens after waitErr is written.
|
||||
err = describeExit(err, s.waitErr)
|
||||
}
|
||||
s.Stop()
|
||||
return nil, fmt.Errorf("%w. Run 'local-ai run' in another terminal to see why it did not come up", err)
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// describeExit adds what is known about how the child died to exitErr, without
|
||||
// putting os/exec's plumbing in front of the user.
|
||||
//
|
||||
// waitErr is exec.ErrWaitDelay when the child exited cleanly but something it
|
||||
// spawned still held its output pipe open past childOutputDrainDelay. The
|
||||
// sentinel's own text names the WaitDelay field, which is meaningless to a
|
||||
// user, so it is translated. Nothing is swallowed: os/exec only substitutes
|
||||
// ErrWaitDelay when the process itself exited without an error of its own (see
|
||||
// Cmd.Wait, "Report an error from the copying goroutines only if the program
|
||||
// otherwise exited normally"), so it can never stand in for an *ExitError.
|
||||
func describeExit(exitErr, waitErr error) error {
|
||||
switch {
|
||||
case waitErr == nil:
|
||||
return exitErr
|
||||
case errors.Is(waitErr, exec.ErrWaitDelay):
|
||||
return fmt.Errorf("%w, and left a subprocess of its own still running", exitErr)
|
||||
default:
|
||||
return fmt.Errorf("%w: %w", exitErr, waitErr)
|
||||
}
|
||||
}
|
||||
|
||||
// Stop terminates the server this process started, giving it a chance to shut
|
||||
// down cleanly first. It is safe to call on a nil or never-started server, and
|
||||
// safe to call more than once.
|
||||
func (s *StartedServer) Stop() {
|
||||
if s == nil || s.proc == nil {
|
||||
return
|
||||
}
|
||||
s.stopOnce.Do(func() {
|
||||
// SIGINT rather than SIGKILL: local-ai run installs its own handler and
|
||||
// needs it to unload models and stop backend subprocesses. Killing it
|
||||
// outright would strand those children.
|
||||
_ = s.proc.Signal(os.Interrupt)
|
||||
|
||||
select {
|
||||
case <-s.exited:
|
||||
case <-time.After(shutdownGrace):
|
||||
// It ignored the interrupt or wedged on the way down. The user is
|
||||
// waiting on their shell prompt, so stop being polite.
|
||||
_ = s.proc.Kill()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// waitReady polls the endpoint's /readyz until the server reports ready, the
|
||||
// budget expires, the caller gives up, or exited signals that the process we
|
||||
// are waiting on is gone. A nil exited channel means there is no process to
|
||||
// watch.
|
||||
//
|
||||
// Readiness lives on the endpoint ROOT, not under the /v1 API base URL, and it
|
||||
// answers 503 for as long as startup is still in progress.
|
||||
func waitReady(ctx context.Context, endpoint string, timeout time.Duration, exited <-chan struct{}) error {
|
||||
url := strings.TrimSuffix(endpoint, "/") + "/readyz"
|
||||
|
||||
// A real deadline rather than context.WithCancel plus a timer: the latter
|
||||
// expires as context.Canceled, which every classifier here reads as "the
|
||||
// caller gave up" rather than "the endpoint never answered".
|
||||
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
client := httpclient.NewWithTimeout(readyProbeTimeout)
|
||||
ticker := time.NewTicker(readyPollInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-exited:
|
||||
return errServerExited
|
||||
case <-waitCtx.Done():
|
||||
// Distinguish our budget from the caller's: only ours is advice
|
||||
// about the server.
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("the LocalAI server did not become ready within %s", timeout)
|
||||
case <-ticker.C:
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(waitCtx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("building the readiness request for %s: %w", url, err)
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
continue // nothing listening yet
|
||||
}
|
||||
// Drain before closing so the next poll can reuse the connection
|
||||
// instead of opening a socket every 500ms for two minutes.
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
return nil
|
||||
}
|
||||
// Anything else means startup is still in progress; keep polling.
|
||||
}
|
||||
}
|
||||
@@ -1,375 +0,0 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
// unusedPort is a loopback address nothing listens on, used wherever a spec
|
||||
// needs a readiness poll to keep failing. Port 1 is privileged, so no test
|
||||
// process could have bound it.
|
||||
const unusedPort = "http://127.0.0.1:1"
|
||||
|
||||
var _ = Describe("OfferToStart", func() {
|
||||
It("never spawns anything when there is no confirmer", func() {
|
||||
started, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: "http://127.0.0.1:59999",
|
||||
Confirm: nil,
|
||||
Stderr: io.Discard,
|
||||
Executable: "/nonexistent/binary-that-must-not-run",
|
||||
})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(errors.Is(err, ErrDeclined)).To(BeTrue(), "want ErrDeclined, got %v", err)
|
||||
Expect(started).To(BeNil())
|
||||
})
|
||||
|
||||
It("does not spawn when the user declines", func() {
|
||||
asked := false
|
||||
started, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: "http://127.0.0.1:59999",
|
||||
Confirm: func(string) (bool, error) {
|
||||
asked = true
|
||||
return false, nil
|
||||
},
|
||||
Stderr: io.Discard,
|
||||
Executable: "/nonexistent/binary-that-must-not-run",
|
||||
})
|
||||
Expect(asked).To(BeTrue(), "the user should have been asked")
|
||||
Expect(errors.Is(err, ErrDeclined)).To(BeTrue())
|
||||
Expect(started).To(BeNil())
|
||||
})
|
||||
|
||||
It("names the endpoint in the question", func() {
|
||||
var question string
|
||||
_, _ = OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: "http://example.invalid:9090",
|
||||
Confirm: func(q string) (bool, error) {
|
||||
question = q
|
||||
return false, nil
|
||||
},
|
||||
Stderr: io.Discard,
|
||||
Executable: "/nonexistent/binary-that-must-not-run",
|
||||
})
|
||||
Expect(question).To(ContainSubstring("http://example.invalid:9090"))
|
||||
})
|
||||
|
||||
It("propagates a confirmer error", func() {
|
||||
boom := errors.New("boom")
|
||||
_, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: "http://127.0.0.1:59999",
|
||||
Confirm: func(string) (bool, error) { return false, boom },
|
||||
Stderr: io.Discard,
|
||||
Executable: "/nonexistent/binary-that-must-not-run",
|
||||
})
|
||||
Expect(errors.Is(err, boom)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("reports which binary it failed to launch", func() {
|
||||
started, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: "http://127.0.0.1:59999",
|
||||
Confirm: func(string) (bool, error) { return true, nil },
|
||||
Stderr: io.Discard,
|
||||
Executable: "/nonexistent/binary-that-must-not-run",
|
||||
})
|
||||
Expect(started).To(BeNil())
|
||||
Expect(err).To(MatchError(ContainSubstring("starting a LocalAI server")))
|
||||
Expect(err).To(MatchError(ContainSubstring("/nonexistent/binary-that-must-not-run")))
|
||||
})
|
||||
|
||||
It("stops waiting as soon as the process it started exits", func() {
|
||||
// A harmless no-op binary rather than a real server: this exercises the
|
||||
// early-exit path without starting LocalAI, binding a port, or running
|
||||
// 'local-ai run'. Without early-exit detection the call would sit here
|
||||
// polling until ReadyTimeout.
|
||||
bin, lookErr := exec.LookPath("true")
|
||||
if lookErr != nil {
|
||||
Skip("no 'true' binary on PATH to stand in for a server that dies at once")
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
started, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: unusedPort,
|
||||
Confirm: func(string) (bool, error) { return true, nil },
|
||||
Stderr: io.Discard,
|
||||
Executable: bin,
|
||||
ReadyTimeout: 30 * time.Second,
|
||||
})
|
||||
Expect(started).To(BeNil())
|
||||
Expect(err).To(MatchError(ContainSubstring("exited before it became ready")))
|
||||
Expect(time.Since(start)).To(BeNumerically("<", 10*time.Second),
|
||||
"the wait should end with the process, not with the readiness budget")
|
||||
})
|
||||
|
||||
It("gives up on a child whose grandchildren still hold its output pipe", func() {
|
||||
// The real LocalAI shape: 'local-ai run' exits but a backend
|
||||
// subprocess it spawned inherited the stderr pipe and keeps it open.
|
||||
// Without cmd.WaitDelay, cmd.Wait blocks on the copy goroutine, exited
|
||||
// never closes, and the readiness wait runs out the full budget instead
|
||||
// of reporting that the server died.
|
||||
sh, lookErr := exec.LookPath("sh")
|
||||
if lookErr != nil {
|
||||
Skip("no 'sh' binary on PATH to stand in for a server with a lingering child")
|
||||
}
|
||||
|
||||
dir := GinkgoT().TempDir()
|
||||
pidFile := filepath.Join(dir, "grandchild.pid")
|
||||
script := filepath.Join(dir, "server-with-lingering-child")
|
||||
// #nosec G306 -- this has to be executable to stand in for a binary.
|
||||
Expect(os.WriteFile(script,
|
||||
[]byte("#!"+sh+"\nsleep 30 &\necho $! > "+pidFile+"\nexit 0\n"),
|
||||
0o700)).To(Succeed())
|
||||
|
||||
// Reap the grandchild whatever happens: it outlives its own parent by
|
||||
// design, so nothing else will clean it up.
|
||||
DeferCleanup(func() {
|
||||
raw, err := os.ReadFile(pidFile)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
pid, err := strconv.Atoi(strings.TrimSpace(string(raw)))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
proc, err := os.FindProcess(pid)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = proc.Kill()
|
||||
_, _ = proc.Wait()
|
||||
})
|
||||
|
||||
start := time.Now()
|
||||
started, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: unusedPort,
|
||||
Confirm: func(string) (bool, error) { return true, nil },
|
||||
Stderr: io.Discard,
|
||||
Executable: script,
|
||||
ReadyTimeout: 25 * time.Second,
|
||||
})
|
||||
elapsed := time.Since(start)
|
||||
|
||||
Expect(started).To(BeNil())
|
||||
Expect(err).To(MatchError(ContainSubstring("exited before it became ready")),
|
||||
"an unbounded cmd.Wait would report a readiness timeout instead")
|
||||
Expect(elapsed).To(BeNumerically("<", 20*time.Second),
|
||||
"the wait must be bounded by the output drain, not by the readiness budget")
|
||||
|
||||
// This is the case where cmd.Wait returns exec.ErrWaitDelay, whose own
|
||||
// text names a struct field of os/exec. Users get told what happened
|
||||
// instead.
|
||||
Expect(err).NotTo(MatchError(ContainSubstring("WaitDelay")),
|
||||
"os/exec plumbing must not reach the user")
|
||||
Expect(err).NotTo(MatchError(ContainSubstring("exec:")))
|
||||
Expect(err).To(MatchError(ContainSubstring("left a subprocess of its own still running")))
|
||||
})
|
||||
|
||||
It("reports the exit status of a server that failed outright", func() {
|
||||
// The counterpart to the case above: translating ErrWaitDelay must not
|
||||
// cost a real exit status, which is the one diagnostic worth having.
|
||||
bin, lookErr := exec.LookPath("false")
|
||||
if lookErr != nil {
|
||||
Skip("no 'false' binary on PATH to stand in for a server that fails")
|
||||
}
|
||||
|
||||
_, err := OfferToStart(context.Background(), StartOptions{
|
||||
Endpoint: unusedPort,
|
||||
Confirm: func(string) (bool, error) { return true, nil },
|
||||
Stderr: io.Discard,
|
||||
Executable: bin,
|
||||
ReadyTimeout: 30 * time.Second,
|
||||
})
|
||||
Expect(err).To(MatchError(ContainSubstring("exited before it became ready")))
|
||||
Expect(err).To(MatchError(ContainSubstring("exit status 1")))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("StartedServer.Stop", func() {
|
||||
It("is a no-op on a server that was never started", func() {
|
||||
var nilServer *StartedServer
|
||||
Expect(nilServer.Stop).NotTo(Panic())
|
||||
Expect((&StartedServer{}).Stop).NotTo(Panic())
|
||||
})
|
||||
|
||||
It("interrupts the child exactly once however often it is called", func() {
|
||||
s, proc := stoppableServer()
|
||||
|
||||
s.Stop()
|
||||
s.Stop()
|
||||
s.Stop()
|
||||
|
||||
Expect(proc.interrupts.Load()).To(Equal(int32(1)),
|
||||
"a second Stop must not signal the child again")
|
||||
Expect(proc.kills.Load()).To(BeZero(), "a child that already exited must not be killed")
|
||||
})
|
||||
|
||||
It("interrupts the child exactly once when called concurrently", func() {
|
||||
// The realistic double-Stop: a deferred Stop on the way out racing the
|
||||
// signal handler that also owns shutting the server down.
|
||||
const callers = 8
|
||||
|
||||
s, proc := stoppableServer()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(callers)
|
||||
for range callers {
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
s.Stop()
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
Expect(proc.interrupts.Load()).To(Equal(int32(1)))
|
||||
Expect(proc.kills.Load()).To(BeZero())
|
||||
})
|
||||
|
||||
It("asks the child to interrupt rather than killing it outright", func() {
|
||||
// The escalation order is the whole point of the grace period: SIGKILL
|
||||
// first would strand the backend subprocesses local-ai run owns.
|
||||
s, proc := stoppableServer()
|
||||
|
||||
s.Stop()
|
||||
|
||||
Expect(proc.lastSignal.Load()).To(Equal(os.Interrupt))
|
||||
Expect(proc.kills.Load()).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
// countingProcess stands in for the *os.Process that Stop drives, recording
|
||||
// what it was asked to do.
|
||||
type countingProcess struct {
|
||||
interrupts atomic.Int32
|
||||
kills atomic.Int32
|
||||
lastSignal atomic.Value
|
||||
}
|
||||
|
||||
func (p *countingProcess) Signal(sig os.Signal) error {
|
||||
p.interrupts.Add(1)
|
||||
p.lastSignal.Store(sig)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *countingProcess) Kill() error {
|
||||
p.kills.Add(1)
|
||||
return nil
|
||||
}
|
||||
|
||||
// stoppableServer builds a StartedServer whose child has already exited, driven
|
||||
// by a countingProcess rather than a real one. Nothing is spawned.
|
||||
func stoppableServer() (*StartedServer, *countingProcess) {
|
||||
proc := &countingProcess{}
|
||||
exited := make(chan struct{})
|
||||
close(exited)
|
||||
return &StartedServer{exited: exited, proc: proc}, proc
|
||||
}
|
||||
|
||||
var _ = Describe("newServerCommand", func() {
|
||||
It("bounds how long it will wait for the child's output pipes", func() {
|
||||
cmd := newServerCommand("/nonexistent/binary-that-must-not-run", io.Discard)
|
||||
|
||||
// An unbounded wait is the failure mode: backend subprocesses inherit
|
||||
// the child's stderr pipe and can hold it open long after the server
|
||||
// itself is gone.
|
||||
Expect(cmd.WaitDelay).To(BeNumerically(">", 0), "cmd.Wait must not be unbounded")
|
||||
Expect(cmd.WaitDelay).To(BeNumerically("<", shutdownGrace),
|
||||
"a drain longer than the shutdown grace would kill a cleanly exited server")
|
||||
})
|
||||
|
||||
It("runs the server subcommand without giving it the terminal", func() {
|
||||
cmd := newServerCommand("/nonexistent/binary-that-must-not-run", io.Discard)
|
||||
|
||||
Expect(cmd.Args).To(Equal([]string{"/nonexistent/binary-that-must-not-run", "run"}))
|
||||
Expect(cmd.Stdin).To(BeNil(), "the child must not compete with the agent for stdin")
|
||||
Expect(cmd.Stdout).NotTo(BeNil())
|
||||
Expect(cmd.Stderr).NotTo(BeNil())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("waitReady", func() {
|
||||
It("polls /readyz on the endpoint root and returns only once it answers 200", func() {
|
||||
// readyOnPoll is deliberately above 1. A handler that answers 200 to the
|
||||
// first poll cannot tell a correct implementation apart from one that
|
||||
// treats 503 as ready, because both return after a single request; the
|
||||
// poll count is what makes 503-as-ready observable.
|
||||
const readyOnPoll = 3
|
||||
|
||||
var polls atomic.Int32
|
||||
var paths atomic.Value
|
||||
paths.Store("")
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
paths.Store(r.URL.Path)
|
||||
if polls.Add(1) < readyOnPoll {
|
||||
// What LocalAI answers while startup is still in progress.
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
Expect(waitReady(context.Background(), srv.URL, 20*time.Second, nil)).To(Succeed())
|
||||
Expect(paths.Load()).To(Equal("/readyz"), "readiness lives on the endpoint root, not under /v1")
|
||||
Expect(polls.Load()).To(BeNumerically(">=", readyOnPoll),
|
||||
"503 means startup is still in progress and must never be accepted as ready")
|
||||
})
|
||||
|
||||
It("tolerates a trailing slash on the endpoint", func() {
|
||||
var path atomic.Value
|
||||
path.Store("")
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
path.Store(r.URL.Path)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
Expect(waitReady(context.Background(), srv.URL+"/", 20*time.Second, nil)).To(Succeed())
|
||||
Expect(path.Load()).To(Equal("/readyz"))
|
||||
})
|
||||
|
||||
It("reports a timeout, not a cancellation, when the budget runs out", func() {
|
||||
err := waitReady(context.Background(), unusedPort, 1200*time.Millisecond, nil)
|
||||
Expect(err).To(HaveOccurred())
|
||||
// A budget built from context.WithCancel plus a timer would surface as
|
||||
// context.Canceled, which downstream code reads as "the caller gave up"
|
||||
// and would stop classifying a hung server as unreachable.
|
||||
Expect(errors.Is(err, context.Canceled)).To(BeFalse(), "got %v", err)
|
||||
Expect(err).To(MatchError(ContainSubstring("did not become ready")))
|
||||
})
|
||||
|
||||
It("returns the caller's cancellation when the caller gives up", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
cancel()
|
||||
}()
|
||||
defer cancel()
|
||||
|
||||
err := waitReady(ctx, unusedPort, time.Minute, nil)
|
||||
Expect(errors.Is(err, context.Canceled)).To(BeTrue(), "got %v", err)
|
||||
})
|
||||
|
||||
It("gives up when the process it is waiting on has exited", func() {
|
||||
exited := make(chan struct{})
|
||||
close(exited)
|
||||
|
||||
err := waitReady(context.Background(), unusedPort, time.Minute, exited)
|
||||
Expect(err).To(MatchError(ContainSubstring("exited before it became ready")))
|
||||
})
|
||||
})
|
||||
112
core/cli/chat/session.go
Normal file
112
core/cli/chat/session.go
Normal file
@@ -0,0 +1,112 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"slices"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
chatRoleUser = "user"
|
||||
chatRoleAssistant = "assistant"
|
||||
)
|
||||
|
||||
type chatMessage struct {
|
||||
Role string
|
||||
Content string
|
||||
}
|
||||
|
||||
type chatSession struct {
|
||||
client chatClient
|
||||
model string
|
||||
models []string
|
||||
messages []chatMessage
|
||||
}
|
||||
|
||||
func newChatSession(ctx context.Context, client chatClient, requestedModel string) (*chatSession, error) {
|
||||
models, err := client.ListModels(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list models: %w", err)
|
||||
}
|
||||
|
||||
model, err := resolveChatModel(requestedModel, models)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &chatSession{
|
||||
client: client,
|
||||
model: model,
|
||||
models: models,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *chatSession) CurrentModel() string {
|
||||
return s.model
|
||||
}
|
||||
|
||||
func (s *chatSession) Models() []string {
|
||||
models := make([]string, len(s.models))
|
||||
copy(models, s.models)
|
||||
return models
|
||||
}
|
||||
|
||||
func (s *chatSession) Clear() {
|
||||
s.messages = nil
|
||||
}
|
||||
|
||||
func (s *chatSession) SwitchModel(model string) error {
|
||||
if !slices.Contains(s.models, model) {
|
||||
return fmt.Errorf("model %q is not available. Use /models to see installed models", model)
|
||||
}
|
||||
s.model = model
|
||||
s.Clear()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *chatSession) Send(ctx context.Context, prompt string, out io.Writer) error {
|
||||
s.messages = append(s.messages, chatMessage{
|
||||
Role: chatRoleUser,
|
||||
Content: prompt,
|
||||
})
|
||||
|
||||
answer, err := s.client.StreamChat(ctx, s.model, s.messages, out)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
s.messages = append(s.messages, chatMessage{
|
||||
Role: chatRoleAssistant,
|
||||
Content: answer,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func resolveChatModel(requested string, models []string) (string, error) {
|
||||
switch {
|
||||
case requested == "" && len(models) == 0:
|
||||
return "", errors.New(`no chat models are installed.
|
||||
|
||||
Install a model first, for example:
|
||||
local-ai models list
|
||||
local-ai models install <model>
|
||||
local-ai run
|
||||
|
||||
Then start a chat session:
|
||||
local-ai chat --model <model>`)
|
||||
case requested == "" && len(models) == 1:
|
||||
return models[0], nil
|
||||
case requested == "" && len(models) > 1:
|
||||
var b strings.Builder
|
||||
b.WriteString("multiple models are available; choose one with --model:\n")
|
||||
b.WriteString(formatChatModelList(models, ""))
|
||||
return "", errors.New(b.String())
|
||||
case !slices.Contains(models, requested):
|
||||
return "", fmt.Errorf("model %q is not available. Use `local-ai models list` and `local-ai models install <model>`, or pass an installed model with --model", requested)
|
||||
default:
|
||||
return requested, nil
|
||||
}
|
||||
}
|
||||
56
core/cli/chat/session_test.go
Normal file
56
core/cli/chat/session_test.go
Normal file
@@ -0,0 +1,56 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Chat session", func() {
|
||||
It("keeps model switching and message history out of the terminal adapter", func() {
|
||||
client := &fakeChatClient{
|
||||
models: []string{"alpha", "beta"},
|
||||
answer: "pong",
|
||||
}
|
||||
|
||||
session, err := newChatSession(context.Background(), client, "alpha")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(session.CurrentModel()).To(Equal("alpha"))
|
||||
|
||||
Expect(session.SwitchModel("beta")).To(Succeed())
|
||||
Expect(session.CurrentModel()).To(Equal("beta"))
|
||||
Expect(session.Send(context.Background(), "ping", io.Discard)).To(Succeed())
|
||||
|
||||
Expect(client.requests).To(HaveLen(1))
|
||||
Expect(client.requests[0].model).To(Equal("beta"))
|
||||
Expect(client.requests[0].messages).To(HaveLen(1))
|
||||
Expect(client.requests[0].messages[0].Content).To(Equal("ping"))
|
||||
})
|
||||
})
|
||||
|
||||
type fakeChatClient struct {
|
||||
models []string
|
||||
answer string
|
||||
requests []fakeChatRequest
|
||||
}
|
||||
|
||||
type fakeChatRequest struct {
|
||||
model string
|
||||
messages []chatMessage
|
||||
}
|
||||
|
||||
func (c *fakeChatClient) ListModels(context.Context) ([]string, error) {
|
||||
return c.models, nil
|
||||
}
|
||||
|
||||
func (c *fakeChatClient) StreamChat(_ context.Context, model string, messages []chatMessage, out io.Writer) (string, error) {
|
||||
copied := make([]chatMessage, len(messages))
|
||||
copy(copied, messages)
|
||||
c.requests = append(c.requests, fakeChatRequest{model: model, messages: copied})
|
||||
if _, err := io.WriteString(out, c.answer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return c.answer, nil
|
||||
}
|
||||
93
core/cli/chat/terminal.go
Normal file
93
core/cli/chat/terminal.go
Normal file
@@ -0,0 +1,93 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func runTerminalChat(ctx context.Context, session *chatSession, in io.Reader, out io.Writer) error {
|
||||
scanner := bufio.NewScanner(in)
|
||||
scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024)
|
||||
|
||||
if err := writeChat(out, "LocalAI chat (%s)\n", session.CurrentModel()); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writeChat(out, "Type /exit to quit, /clear to reset the conversation, /models to list models.\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for {
|
||||
if err := writeChat(out, "\n> "); err != nil {
|
||||
return err
|
||||
}
|
||||
if !scanner.Scan() {
|
||||
break
|
||||
}
|
||||
|
||||
prompt := strings.TrimSpace(scanner.Text())
|
||||
switch prompt {
|
||||
case "":
|
||||
continue
|
||||
case "/bye", "/exit", "/quit":
|
||||
return writeChat(out, "bye\n")
|
||||
case "/clear":
|
||||
session.Clear()
|
||||
if err := writeChat(out, "conversation cleared\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
case "/models":
|
||||
if err := printChatModels(out, session.Models(), session.CurrentModel()); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if nextModel, ok := strings.CutPrefix(prompt, "/model "); ok {
|
||||
nextModel = strings.TrimSpace(nextModel)
|
||||
if nextModel == "" {
|
||||
if err := writeChat(out, "usage: /model <name>\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := session.SwitchModel(nextModel); err != nil {
|
||||
if writeErr := writeChat(out, "%s\n", err); writeErr != nil {
|
||||
return writeErr
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := writeChat(out, "switched to %s; conversation cleared\n", session.CurrentModel()); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if err := writeChat(out, "assistant: "); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := session.Send(ctx, prompt, out); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writeChat(out, "\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return scanner.Err()
|
||||
}
|
||||
|
||||
func printChatModels(out io.Writer, models []string, current string) error {
|
||||
if len(models) == 0 {
|
||||
return writeChat(out, "no models installed\n")
|
||||
}
|
||||
return writeChat(out, "%s", formatChatModelList(models, current))
|
||||
}
|
||||
|
||||
func writeChat(out io.Writer, format string, args ...any) error {
|
||||
_, err := fmt.Fprintf(out, format, args...)
|
||||
return err
|
||||
}
|
||||
@@ -8,72 +8,18 @@ import (
|
||||
cliContext "github.com/mudler/LocalAI/core/cli/context"
|
||||
)
|
||||
|
||||
// ChatCMD runs the built-in terminal agent. Everything after the first
|
||||
// positional argument is forwarded to the agent verbatim, so its own
|
||||
// subcommands (plugin, skill, mcp) and their flags work unchanged. LocalAI's
|
||||
// own flags must therefore come first.
|
||||
type ChatCMD struct {
|
||||
Model string `short:"m" help:"Model to use. Defaults to the only model the server offers, or asks when there are several"`
|
||||
Endpoint string `env:"LOCALAI_CHAT_ENDPOINT" default:"http://127.0.0.1:8080" help:"LocalAI server endpoint. The /v1 path is added automatically when omitted"`
|
||||
APIKey string `env:"LOCALAI_API_KEY,API_KEY" help:"API key to use when the LocalAI server requires authentication"`
|
||||
ConfigDir string `env:"LOCALAI_CHAT_CONFIG_DIR" help:"Directory holding the agent's config, plugins, and skills. Defaults to ~/.config/localai/chat" type:"path"`
|
||||
TraceDir string `env:"LOCALAI_CHAT_TRACE_DIR" help:"Write a session LLM trace (NDJSON) to this directory" type:"path"`
|
||||
|
||||
CLI bool `help:"Run in plain CLI mode instead of the full-screen interface"`
|
||||
TUI bool `help:"Force the full-screen interface"`
|
||||
Height string `help:"Run as an inline drop-down of this height, e.g. '40%'"`
|
||||
Tmux bool `help:"Run in a tmux split"`
|
||||
NoTmux bool `name:"no-tmux" help:"Never use a tmux split, even inside tmux"`
|
||||
Init string `help:"Print the shell integration script for Ctrl+Space (zsh, bash, or fish)"`
|
||||
Yolo bool `env:"LOCALAI_CHAT_YOLO" help:"Auto-approve every tool call without prompting"`
|
||||
|
||||
Args []string `arg:"" optional:"" passthrough:"" help:"Arguments forwarded to the agent, e.g. 'plugin install <url>', 'skill list', 'mcp add'"`
|
||||
Model string `short:"m" help:"Model name to use. Defaults to the only model returned by the server when exactly one is available"`
|
||||
Endpoint string `env:"LOCALAI_CHAT_ENDPOINT" default:"http://127.0.0.1:8080" help:"LocalAI server endpoint. The /v1 path is added automatically when omitted"`
|
||||
APIKey string `env:"LOCALAI_API_KEY,API_KEY" help:"API key to use when the LocalAI server requires authentication"`
|
||||
}
|
||||
|
||||
func (c *ChatCMD) Run(ctx *cliContext.Context) error {
|
||||
err := chatcli.Run(context.Background(), chatcli.Options{
|
||||
Args: c.agentArgs(),
|
||||
Endpoint: c.Endpoint,
|
||||
BaseURL: chatAPIBaseURL(c.Endpoint),
|
||||
APIKey: c.APIKey,
|
||||
Model: c.Model,
|
||||
StateDir: c.ConfigDir,
|
||||
TraceDir: c.TraceDir,
|
||||
Yolo: c.Yolo,
|
||||
In: os.Stdin,
|
||||
Out: os.Stdout,
|
||||
ErrOut: os.Stderr,
|
||||
return chatcli.Run(context.Background(), chatcli.Options{
|
||||
Model: c.Model,
|
||||
BaseURL: chatAPIBaseURL(c.Endpoint),
|
||||
APIKey: c.APIKey,
|
||||
In: os.Stdin,
|
||||
Out: os.Stdout,
|
||||
})
|
||||
// The agent explains its own failures on stderr and hands back a code, so
|
||||
// carry the code out and leave the explanation to stand alone.
|
||||
if code, reported := chatcli.ExitStatus(err); reported {
|
||||
return ExitCodeError{Code: code}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// agentArgs rebuilds the argument vector the agent expects: LocalAI's mode
|
||||
// flags are declared here for discoverability and shell completion, so they
|
||||
// have to be translated back into the agent's own flag names.
|
||||
func (c *ChatCMD) agentArgs() []string {
|
||||
var args []string
|
||||
if c.CLI {
|
||||
args = append(args, "--cli")
|
||||
}
|
||||
if c.TUI {
|
||||
args = append(args, "--tui")
|
||||
}
|
||||
if c.Height != "" {
|
||||
args = append(args, "--height", c.Height)
|
||||
}
|
||||
if c.Tmux {
|
||||
args = append(args, "--tmux")
|
||||
}
|
||||
if c.NoTmux {
|
||||
args = append(args, "--no-tmux")
|
||||
}
|
||||
if c.Init != "" {
|
||||
args = append(args, "--init", c.Init)
|
||||
}
|
||||
return append(args, c.Args...)
|
||||
}
|
||||
|
||||
@@ -1,10 +1,6 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/alecthomas/kong"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
@@ -28,70 +24,4 @@ var _ = Describe("Chat command wiring", func() {
|
||||
Expect(chatAPIBaseURL("http://127.0.0.1:8080/localai")).To(Equal("http://127.0.0.1:8080/localai/v1"))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("argument parsing", func() {
|
||||
parse := func(args ...string) *ChatCMD {
|
||||
var cli struct {
|
||||
Chat ChatCMD `cmd:""`
|
||||
}
|
||||
parser, err := kong.New(&cli)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
_, err = parser.Parse(append([]string{"chat"}, args...))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
return &cli.Chat
|
||||
}
|
||||
|
||||
It("leaves Args empty for a bare invocation", func() {
|
||||
Expect(parse().Args).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("binds flags that precede the forwarded arguments", func() {
|
||||
c := parse("--endpoint", "http://host:9090", "--model", "m", "plugin", "list")
|
||||
Expect(c.Endpoint).To(Equal("http://host:9090"))
|
||||
Expect(c.Model).To(Equal("m"))
|
||||
Expect(c.Args).To(Equal([]string{"plugin", "list"}))
|
||||
})
|
||||
|
||||
It("forwards flags that follow the first positional to the agent", func() {
|
||||
c := parse("plugin", "install", "https://example.invalid/p", "--yes")
|
||||
Expect(c.Args).To(Equal([]string{"plugin", "install", "https://example.invalid/p", "--yes"}))
|
||||
})
|
||||
|
||||
It("parses its own mode flags", func() {
|
||||
c := parse("--cli")
|
||||
Expect(c.CLI).To(BeTrue())
|
||||
Expect(c.Args).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
// The agent prints its own diagnosis and hands back a status. main exits
|
||||
// with that status and prints nothing more, so the user reads one message
|
||||
// rather than an "exit status 1" stacked under it.
|
||||
Describe("ExitCodeError", func() {
|
||||
It("carries the status out", func() {
|
||||
Expect(ExitCodeError{Code: 2}.Code).To(Equal(2))
|
||||
})
|
||||
|
||||
It("is recognisable after wrapping", func() {
|
||||
var got ExitCodeError
|
||||
Expect(errors.As(fmt.Errorf("chat: %w", ExitCodeError{Code: 2}), &got)).To(BeTrue())
|
||||
Expect(got.Code).To(Equal(2))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("agentArgs", func() {
|
||||
It("translates mode flags into the agent's own flags", func() {
|
||||
c := &ChatCMD{CLI: true}
|
||||
Expect(c.agentArgs()).To(Equal([]string{"--cli"}))
|
||||
})
|
||||
|
||||
It("puts forwarded arguments after the translated flags", func() {
|
||||
c := &ChatCMD{Height: "40%", Args: []string{"plugin", "list"}}
|
||||
Expect(c.agentArgs()).To(Equal([]string{"--height", "40%", "plugin", "list"}))
|
||||
})
|
||||
|
||||
It("returns nothing for a bare invocation", func() {
|
||||
Expect((&ChatCMD{}).agentArgs()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -9,7 +9,7 @@ var CLI struct {
|
||||
cliContext.Context `embed:""`
|
||||
|
||||
Run RunCMD `cmd:"" help:"Run LocalAI, this the default command if no other command is specified. Run 'local-ai run --help' for more information" default:"withargs"`
|
||||
Chat ChatCMD `cmd:"" help:"Run the built-in terminal agent against a LocalAI server"`
|
||||
Chat ChatCMD `cmd:"" help:"Open an interactive chat session against a running LocalAI server"`
|
||||
Federated FederatedCLI `cmd:"" help:"Run LocalAI in federated mode"`
|
||||
Models ModelsCMD `cmd:"" help:"Manage LocalAI models and definitions"`
|
||||
Backends BackendsCMD `cmd:"" help:"Manage LocalAI backends and definitions"`
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
package cli
|
||||
|
||||
import "fmt"
|
||||
|
||||
// ExitCodeError is a failure a command has already reported to the user. It
|
||||
// carries nothing but the status the process should exit with, and main prints
|
||||
// nothing more for it.
|
||||
//
|
||||
// It exists for commands that hand their terminal to something that does its
|
||||
// own error reporting. Returning that subordinate's error instead would put a
|
||||
// bare "exit status 1" underneath the explanation the user has just read, and
|
||||
// returning nil would tell a script the run succeeded.
|
||||
type ExitCodeError struct{ Code int }
|
||||
|
||||
func (e ExitCodeError) Error() string { return fmt.Sprintf("exit status %d", e.Code) }
|
||||
@@ -44,7 +44,6 @@ const (
|
||||
MethodPredictStream GRPCMethod = "PredictStream"
|
||||
MethodEmbedding GRPCMethod = "Embedding"
|
||||
MethodGenerateImage GRPCMethod = "GenerateImage"
|
||||
MethodUpscaleImage GRPCMethod = "UpscaleImage"
|
||||
MethodGenerateVideo GRPCMethod = "GenerateVideo"
|
||||
MethodGenerate3D GRPCMethod = "Generate3D"
|
||||
MethodAudioTranscription GRPCMethod = "AudioTranscription"
|
||||
@@ -349,7 +348,7 @@ var BackendCapabilities = map[string]BackendCapability{
|
||||
|
||||
// --- Image/video generation backends ---
|
||||
"diffusers": {
|
||||
GRPCMethods: []GRPCMethod{MethodGenerateImage, MethodUpscaleImage, MethodGenerateVideo},
|
||||
GRPCMethods: []GRPCMethod{MethodGenerateImage, MethodGenerateVideo},
|
||||
PossibleUsecases: []string{UsecaseImage, UsecaseVideo},
|
||||
DefaultUsecases: []string{UsecaseImage},
|
||||
Description: "HuggingFace diffusers — Stable Diffusion, Flux, video generation",
|
||||
|
||||
@@ -1,117 +0,0 @@
|
||||
package config
|
||||
|
||||
// Speculative-decoding auto-defaults for the vllm-cpp backend, the safetensors
|
||||
// counterpart of the GGUF/llama.cpp hook in mtp.go.
|
||||
//
|
||||
// The two engines detect and spell the same feature differently. llama.cpp
|
||||
// reads `<arch>.nextn_predict_layers` out of the GGUF header and takes
|
||||
// `spec_type:draft-mtp` in `options:`; vllm.cpp reads `mtp_num_hidden_layers`
|
||||
// out of the checkpoint's config.json and takes vLLM's own
|
||||
// `--speculative-config` JSON, which LocalAI carries in `engine_args`. The
|
||||
// engine resolves the draft depth and the default k itself, so the config only
|
||||
// has to name the method.
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
// hfSpecConfig is the subset of a HuggingFace config.json that decides whether
|
||||
// speculative decoding can be auto-enabled.
|
||||
type hfSpecConfig struct {
|
||||
ModelType string `json:"model_type"`
|
||||
// MtpNumHiddenLayers is the MTP head depth (upstream speculative.py reads
|
||||
// it as n_predict for the qwen3_5 / qwen3_5_moe families).
|
||||
MtpNumHiddenLayers uint32 `json:"mtp_num_hidden_layers"`
|
||||
// DFlashConfig marks a z-lab DFlash DRAFT checkpoint (mask_token_id +
|
||||
// target_layer_ids). Its presence means this repo is a draft, not a
|
||||
// servable target.
|
||||
DFlashConfig json.RawMessage `json:"dflash_config"`
|
||||
// TextConfig is where multimodal checkpoints nest the language-model
|
||||
// config, and therefore the MTP depth.
|
||||
TextConfig *hfSpecConfig `json:"text_config"`
|
||||
}
|
||||
|
||||
// parseHFSpecConfig decodes the speculative-relevant subset of a config.json.
|
||||
// A document that does not parse yields nothing rather than an error: detection
|
||||
// is best-effort and must never break an import.
|
||||
func parseHFSpecConfig(configJSON []byte) (hfSpecConfig, bool) {
|
||||
if len(configJSON) == 0 {
|
||||
return hfSpecConfig{}, false
|
||||
}
|
||||
var c hfSpecConfig
|
||||
if err := json.Unmarshal(configJSON, &c); err != nil {
|
||||
xlog.Debug("[vllm-spec] config.json did not parse; skipping detection", "error", err)
|
||||
return hfSpecConfig{}, false
|
||||
}
|
||||
return c, true
|
||||
}
|
||||
|
||||
// IsDFlashDraftConfig reports whether a HuggingFace config.json describes a
|
||||
// DFlash DRAFT checkpoint. Unlike MTP - whose head ships inside the target
|
||||
// checkpoint's `mtp.*` tensors - a DFlash draft is its own repo that can only
|
||||
// run paired with a target it verifies against, so it must never be configured
|
||||
// as a standalone model.
|
||||
func IsDFlashDraftConfig(configJSON []byte) bool {
|
||||
c, ok := parseHFSpecConfig(configJSON)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return len(c.DFlashConfig) > 0 ||
|
||||
(c.TextConfig != nil && len(c.TextConfig.DFlashConfig) > 0)
|
||||
}
|
||||
|
||||
// HasSafetensorsMTPHead reports whether a HuggingFace config.json declares a
|
||||
// self-speculating Multi-Token Prediction head, returning its depth. The depth
|
||||
// is informational: vllm.cpp resolves n_predict and the default
|
||||
// num_speculative_tokens from the checkpoint itself.
|
||||
//
|
||||
// DFlash drafts are excluded for the same reason `gemma4-assistant` GGUFs are
|
||||
// excluded from the llama.cpp hook: they carry head metadata but cannot
|
||||
// self-speculate.
|
||||
//
|
||||
// NOTE this is a safetensors-only signal. vllm.cpp rejects an MTP config over a
|
||||
// GGUF source, because the `mtp.*` draft tensors only exist in the safetensors
|
||||
// checkpoint - so the GGUF import path must not use this.
|
||||
func HasSafetensorsMTPHead(configJSON []byte) (uint32, bool) {
|
||||
c, ok := parseHFSpecConfig(configJSON)
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
if IsDFlashDraftConfig(configJSON) {
|
||||
return 0, false
|
||||
}
|
||||
n := c.MtpNumHiddenLayers
|
||||
if n == 0 && c.TextConfig != nil {
|
||||
n = c.TextConfig.MtpNumHiddenLayers
|
||||
}
|
||||
return n, n > 0
|
||||
}
|
||||
|
||||
// ApplyVLLMSpeculativeDefaults enables MTP speculative decoding in cfg's
|
||||
// engine_args when nothing is configured there yet. It is a no-op when the user
|
||||
// already set a speculative_config, so an explicit choice (a different method,
|
||||
// an explicit k, a DFlash draft) is never clobbered.
|
||||
//
|
||||
// `layers` is the detected head depth and is only used for the diagnostic log
|
||||
// line - the engine derives the real k from the checkpoint.
|
||||
func ApplyVLLMSpeculativeDefaults(cfg *ModelConfig, layers uint32) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
if _, set := cfg.EngineArgs["speculative_config"]; set {
|
||||
xlog.Debug("[vllm-spec] MTP head detected but speculative_config already configured; leaving user choice intact",
|
||||
"name", cfg.Name, "mtp_num_hidden_layers", layers)
|
||||
return
|
||||
}
|
||||
if cfg.EngineArgs == nil {
|
||||
cfg.EngineArgs = map[string]any{}
|
||||
}
|
||||
// Only the method: vllm.cpp defaults num_speculative_tokens to the
|
||||
// checkpoint's own n_predict (speculative.py:865-875), which is the right
|
||||
// value far more reliably than anything guessable here.
|
||||
cfg.EngineArgs["speculative_config"] = map[string]any{"method": "mtp"}
|
||||
xlog.Info("[vllm-spec] MTP head detected; enabling mtp speculative decoding",
|
||||
"name", cfg.Name, "mtp_num_hidden_layers", layers)
|
||||
}
|
||||
@@ -1,117 +0,0 @@
|
||||
package config_test
|
||||
|
||||
import (
|
||||
. "github.com/mudler/LocalAI/core/config"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("vllm-cpp speculative-decoding auto-defaults", func() {
|
||||
Context("HasSafetensorsMTPHead", func() {
|
||||
It("detects a top-level mtp_num_hidden_layers", func() {
|
||||
n, ok := HasSafetensorsMTPHead([]byte(`{
|
||||
"model_type": "qwen3_5_moe",
|
||||
"mtp_num_hidden_layers": 1
|
||||
}`))
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(n).To(Equal(uint32(1)))
|
||||
})
|
||||
|
||||
It("detects the head nested under text_config", func() {
|
||||
// Multimodal checkpoints nest the language-model config, which is
|
||||
// where the MTP depth lives (mirrors the engine's own resolution
|
||||
// off config.raw text_config).
|
||||
n, ok := HasSafetensorsMTPHead([]byte(`{
|
||||
"model_type": "qwen3_5_moe",
|
||||
"text_config": {"mtp_num_hidden_layers": 2}
|
||||
}`))
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(n).To(Equal(uint32(2)))
|
||||
})
|
||||
|
||||
It("reports no head when the key is absent", func() {
|
||||
n, ok := HasSafetensorsMTPHead([]byte(`{"model_type": "llama"}`))
|
||||
Expect(ok).To(BeFalse())
|
||||
Expect(n).To(BeZero())
|
||||
})
|
||||
|
||||
It("reports no head for a zero depth", func() {
|
||||
_, ok := HasSafetensorsMTPHead([]byte(`{"mtp_num_hidden_layers": 0}`))
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
|
||||
It("ignores a DFlash draft checkpoint", func() {
|
||||
// A DFlash draft is a SEPARATE checkpoint that cannot serve alone:
|
||||
// it needs a target to verify against. Same exclusion the GGUF path
|
||||
// makes for gemma4-assistant drafts.
|
||||
_, ok := HasSafetensorsMTPHead([]byte(`{
|
||||
"model_type": "qwen3_dflash",
|
||||
"mtp_num_hidden_layers": 1,
|
||||
"dflash_config": {"mask_token_id": 151666, "target_layer_ids": [0, 1]}
|
||||
}`))
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
|
||||
It("reports no head on unparseable JSON", func() {
|
||||
_, ok := HasSafetensorsMTPHead([]byte(`{not json`))
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
|
||||
It("reports no head on empty input", func() {
|
||||
_, ok := HasSafetensorsMTPHead(nil)
|
||||
Expect(ok).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
Context("IsDFlashDraftConfig", func() {
|
||||
It("recognises a draft by its dflash_config block", func() {
|
||||
Expect(IsDFlashDraftConfig([]byte(`{
|
||||
"dflash_config": {"mask_token_id": 151666, "target_layer_ids": [0]}
|
||||
}`))).To(BeTrue())
|
||||
})
|
||||
|
||||
It("does not flag an ordinary checkpoint", func() {
|
||||
Expect(IsDFlashDraftConfig([]byte(`{"model_type": "qwen3_5_moe"}`))).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
Context("ApplyVLLMSpeculativeDefaults", func() {
|
||||
It("writes the mtp method into engine_args", func() {
|
||||
cfg := &ModelConfig{Name: "qwen"}
|
||||
ApplyVLLMSpeculativeDefaults(cfg, 1)
|
||||
Expect(cfg.EngineArgs).To(HaveKey("speculative_config"))
|
||||
spec, ok := cfg.EngineArgs["speculative_config"].(map[string]any)
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(spec["method"]).To(Equal("mtp"))
|
||||
})
|
||||
|
||||
It("leaves an existing speculative_config alone", func() {
|
||||
cfg := &ModelConfig{
|
||||
Name: "qwen",
|
||||
LLMConfig: LLMConfig{
|
||||
EngineArgs: map[string]any{
|
||||
"speculative_config": map[string]any{"method": "ngram", "num_speculative_tokens": 4},
|
||||
},
|
||||
},
|
||||
}
|
||||
ApplyVLLMSpeculativeDefaults(cfg, 1)
|
||||
spec := cfg.EngineArgs["speculative_config"].(map[string]any)
|
||||
Expect(spec["method"]).To(Equal("ngram"))
|
||||
})
|
||||
|
||||
It("preserves unrelated engine_args keys", func() {
|
||||
cfg := &ModelConfig{
|
||||
Name: "qwen",
|
||||
LLMConfig: LLMConfig{EngineArgs: map[string]any{"max_num_seqs": 32}},
|
||||
}
|
||||
ApplyVLLMSpeculativeDefaults(cfg, 1)
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("max_num_seqs", 32))
|
||||
Expect(cfg.EngineArgs).To(HaveKey("speculative_config"))
|
||||
})
|
||||
|
||||
It("tolerates a nil config", func() {
|
||||
Expect(func() { ApplyVLLMSpeculativeDefaults(nil, 1) }).ToNot(Panic())
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -1,211 +0,0 @@
|
||||
package gallery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
"github.com/mudler/LocalAI/pkg/vram"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
// EstimateInput builds the VRAM estimator's input from a gallery entry.
|
||||
//
|
||||
// It lives here rather than beside the HTTP handler because two callers need
|
||||
// it: the handler answering one model, and the warmer below answering all of
|
||||
// them ahead of time.
|
||||
func EstimateInput(m *GalleryModel) vram.ModelEstimateInput {
|
||||
var input vram.ModelEstimateInput
|
||||
input.Size = m.Size
|
||||
if repoID := extractHFRepo(m.Overrides, m.URLs); repoID != "" {
|
||||
input.HFRepo = repoID
|
||||
}
|
||||
for _, f := range m.AdditionalFiles {
|
||||
if vram.IsWeightFile(f.URI) {
|
||||
input.Files = append(input.Files, vram.FileInput{URI: f.URI, Size: 0})
|
||||
}
|
||||
}
|
||||
return input
|
||||
}
|
||||
|
||||
// extractHFRepo finds a HuggingFace repo ID in a model's overrides or URLs.
|
||||
func extractHFRepo(overrides map[string]any, urls []string) string {
|
||||
if overrides != nil {
|
||||
if params, ok := overrides["parameters"].(map[string]any); ok {
|
||||
if modelRef, ok := params["model"].(string); ok {
|
||||
if repoID, ok := vram.ExtractHFRepoID(modelRef); ok {
|
||||
return repoID
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, u := range urls {
|
||||
if repoID, ok := vram.ExtractHFRepoID(u); ok {
|
||||
return repoID
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// EstimateWarmConfig bounds the background warm-up.
|
||||
type EstimateWarmConfig struct {
|
||||
// Limit is how many gallery entries to warm, in gallery order. Zero
|
||||
// disables warming entirely. The order matters: it is the order the UI
|
||||
// lists them in, so the entries a user sees first are warmed first.
|
||||
Limit int
|
||||
// Concurrency is how many estimates run at once. Each one can be a remote
|
||||
// probe, so this is deliberately small: the point is to be finished before
|
||||
// anybody looks, not to saturate the link or the upstream.
|
||||
Concurrency int
|
||||
// Contexts are the context lengths to estimate at. These want to match what
|
||||
// the UI asks for, or the warmed entry is not the one it reads.
|
||||
Contexts []uint32
|
||||
}
|
||||
|
||||
// DefaultEstimateWarmConfig is what the server uses unless told otherwise.
|
||||
//
|
||||
// The limit is a deliberate compromise. Warming the whole gallery would be
|
||||
// thousands of remote probes on every boot, which is rude to the upstream and
|
||||
// slow to finish; warming nothing leaves the first page of the model gallery
|
||||
// paying two seconds per row. A few hundred covers what anyone browses in a
|
||||
// sitting, and everything past it still warms itself on first view.
|
||||
var DefaultEstimateWarmConfig = EstimateWarmConfig{
|
||||
Limit: 300,
|
||||
Concurrency: 4,
|
||||
Contexts: []uint32{8192, 16384, 32768, 65536, 131072, 262144},
|
||||
}
|
||||
|
||||
// WarmEstimateCache fills the gallery's derived caches in the background.
|
||||
//
|
||||
// Two things are warmed, and they are the same cost wearing different hats.
|
||||
// An estimate for an entry the server has never seen costs a network probe of
|
||||
// its weight files, and describing an entry's variants costs one probe per
|
||||
// build it offers. The UI asks for an estimate per row and a variant
|
||||
// description per model opened, so without this the first visitor pays for
|
||||
// both: ten seconds of a page filling in its own sizes, then another second
|
||||
// and a half the first time they click anything.
|
||||
//
|
||||
// Both land in the same caches underneath, which is why one pass covers them.
|
||||
//
|
||||
// It returns immediately; the work happens on its own goroutine and stops when
|
||||
// ctx is done. Failures are logged at debug and otherwise ignored: a warm-up
|
||||
// that cannot reach an upstream must never stop the server from starting, and
|
||||
// the entry it failed on simply stays cold.
|
||||
func WarmEstimateCache(ctx context.Context, galleries []config.Gallery, systemState *system.SystemState, cfg EstimateWarmConfig) {
|
||||
if cfg.Limit <= 0 || cfg.Concurrency <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
started := time.Now()
|
||||
|
||||
models, err := AvailableGalleryModelsCached(galleries, systemState)
|
||||
if err != nil {
|
||||
xlog.Debug("VRAM estimate warm-up skipped, gallery unavailable", "error", err)
|
||||
return
|
||||
}
|
||||
if len(models) > cfg.Limit {
|
||||
models = models[:cfg.Limit]
|
||||
}
|
||||
if len(models) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// The host gate the variant picker resolves against. Derived once: it
|
||||
// describes this machine, not this entry, and HostResolveEnv reads the
|
||||
// system state to build it.
|
||||
env := HostResolveEnv(ctx, systemState)
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
cursor = make(chan *GalleryModel)
|
||||
warmed int
|
||||
warmedVariants int
|
||||
mu sync.Mutex
|
||||
)
|
||||
|
||||
for i := 0; i < cfg.Concurrency; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for m := range cursor {
|
||||
// Per entry, not for the run: one unreachable weight file
|
||||
// must not hold a worker for the whole warm-up.
|
||||
entryCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
|
||||
input := EstimateInput(m)
|
||||
if len(input.Files) > 0 || input.HFRepo != "" || input.Size != "" {
|
||||
if _, err := vram.EstimateModelMultiContext(entryCtx, input, cfg.Contexts); err != nil {
|
||||
xlog.Debug("VRAM estimate warm-up failed for entry", "model", m.GetName(), "error", err)
|
||||
} else {
|
||||
mu.Lock()
|
||||
warmed++
|
||||
mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
// Describing variants probes each build the entry offers.
|
||||
// An entry that declares none costs nothing here, so this is
|
||||
// gated rather than attempted and discarded.
|
||||
if m.HasVariants() {
|
||||
if _, err := DescribeVariants(models, m, env); err != nil {
|
||||
xlog.Debug("variant warm-up failed for entry", "model", m.GetName(), "error", err)
|
||||
} else {
|
||||
mu.Lock()
|
||||
warmedVariants++
|
||||
mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
cancel()
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
feed:
|
||||
for _, m := range models {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
break feed
|
||||
case cursor <- m:
|
||||
}
|
||||
}
|
||||
close(cursor)
|
||||
wg.Wait()
|
||||
|
||||
if ctx.Err() != nil {
|
||||
xlog.Debug("gallery warm-up stopped", "estimates", warmed, "variants", warmedVariants)
|
||||
return
|
||||
}
|
||||
xlog.Info("gallery caches warmed", "estimates", warmed, "variants", warmedVariants, "of", len(models), "took", time.Since(started).Round(time.Second))
|
||||
}()
|
||||
}
|
||||
|
||||
// EstimateWarmConfigFromEnv reads the warm-up bounds from the environment,
|
||||
// falling back to the defaults.
|
||||
//
|
||||
// LOCALAI_VRAM_WARM_LIMIT entries to warm; 0 disables the warm-up
|
||||
// LOCALAI_VRAM_WARM_CONCURRENCY estimates in flight at once
|
||||
//
|
||||
// Env rather than a flag because it is an operational tuning knob, not part of
|
||||
// what the server does: an air-gapped host wants it off, and a host behind a
|
||||
// slow link wants it slower, and neither is a decision the CLI should carry.
|
||||
func EstimateWarmConfigFromEnv() EstimateWarmConfig {
|
||||
cfg := DefaultEstimateWarmConfig
|
||||
if v, ok := os.LookupEnv("LOCALAI_VRAM_WARM_LIMIT"); ok {
|
||||
if n, err := strconv.Atoi(strings.TrimSpace(v)); err == nil && n >= 0 {
|
||||
cfg.Limit = n
|
||||
}
|
||||
}
|
||||
if v, ok := os.LookupEnv("LOCALAI_VRAM_WARM_CONCURRENCY"); ok {
|
||||
if n, err := strconv.Atoi(strings.TrimSpace(v)); err == nil && n > 0 {
|
||||
cfg.Concurrency = n
|
||||
}
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
@@ -1,115 +0,0 @@
|
||||
package gallery_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
)
|
||||
|
||||
var _ = Describe("VRAM estimate warm-up", func() {
|
||||
var state *system.SystemState
|
||||
|
||||
BeforeEach(func() {
|
||||
dir, err := os.MkdirTemp("", "warm")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(func() { os.RemoveAll(dir) })
|
||||
state, err = system.GetSystemState(system.WithModelPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
gallery.ResetGalleryModelCache()
|
||||
DeferCleanup(gallery.ResetGalleryModelCache)
|
||||
})
|
||||
|
||||
It("does nothing when disabled, and returns without blocking", func() {
|
||||
cfg := gallery.DefaultEstimateWarmConfig
|
||||
cfg.Limit = 0
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
gallery.WarmEstimateCache(context.Background(), []config.Gallery{}, state, cfg)
|
||||
}()
|
||||
Eventually(done, "1s").Should(BeClosed())
|
||||
})
|
||||
|
||||
It("returns immediately even when there is work to do", func() {
|
||||
// The caller is a server still starting up: warming must never be on
|
||||
// the path to listening.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
gallery.WarmEstimateCache(context.Background(), []config.Gallery{}, state, gallery.DefaultEstimateWarmConfig)
|
||||
}()
|
||||
Eventually(done, "1s").Should(BeClosed())
|
||||
})
|
||||
|
||||
It("stops when its context is cancelled", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
gallery.WarmEstimateCache(ctx, []config.Gallery{}, state, gallery.DefaultEstimateWarmConfig)
|
||||
cancel()
|
||||
// Nothing to assert beyond not hanging or panicking: an aborted warm-up
|
||||
// leaves entries cold, which is the state they were already in.
|
||||
Consistently(func() bool { return true }, "100ms").Should(BeTrue())
|
||||
})
|
||||
|
||||
Describe("configuration from the environment", func() {
|
||||
AfterEach(func() {
|
||||
os.Unsetenv("LOCALAI_VRAM_WARM_LIMIT")
|
||||
os.Unsetenv("LOCALAI_VRAM_WARM_CONCURRENCY")
|
||||
})
|
||||
|
||||
It("falls back to the defaults", func() {
|
||||
cfg := gallery.EstimateWarmConfigFromEnv()
|
||||
Expect(cfg.Limit).To(Equal(gallery.DefaultEstimateWarmConfig.Limit))
|
||||
Expect(cfg.Concurrency).To(Equal(gallery.DefaultEstimateWarmConfig.Concurrency))
|
||||
})
|
||||
|
||||
It("lets an operator turn it off entirely", func() {
|
||||
os.Setenv("LOCALAI_VRAM_WARM_LIMIT", "0")
|
||||
Expect(gallery.EstimateWarmConfigFromEnv().Limit).To(BeZero())
|
||||
})
|
||||
|
||||
It("lets an operator slow it down", func() {
|
||||
os.Setenv("LOCALAI_VRAM_WARM_CONCURRENCY", "1")
|
||||
Expect(gallery.EstimateWarmConfigFromEnv().Concurrency).To(Equal(1))
|
||||
})
|
||||
|
||||
It("ignores values that are not usable", func() {
|
||||
os.Setenv("LOCALAI_VRAM_WARM_LIMIT", "not-a-number")
|
||||
os.Setenv("LOCALAI_VRAM_WARM_CONCURRENCY", "0")
|
||||
cfg := gallery.EstimateWarmConfigFromEnv()
|
||||
Expect(cfg.Limit).To(Equal(gallery.DefaultEstimateWarmConfig.Limit))
|
||||
// Zero workers would be a warm-up that never runs while looking
|
||||
// enabled, so it keeps the default rather than honouring it.
|
||||
Expect(cfg.Concurrency).To(Equal(gallery.DefaultEstimateWarmConfig.Concurrency))
|
||||
})
|
||||
})
|
||||
|
||||
It("warms variant descriptions as well as estimates", func() {
|
||||
// Both are the same cost wearing different hats - a probe of an entry's
|
||||
// weight files - and both land in the same caches, so a warm-up that
|
||||
// covered only one would leave the first click paying for the other.
|
||||
// Asserted through the shared config rather than by observing network
|
||||
// calls: the gallery here is empty by design.
|
||||
Expect(gallery.DefaultEstimateWarmConfig.Limit).To(BeNumerically(">", 0))
|
||||
})
|
||||
|
||||
It("keeps the estimate contexts the UI actually asks for", func() {
|
||||
// A warmed entry at the wrong context lengths is a cache the gallery
|
||||
// never reads, so this pins them together.
|
||||
Expect(gallery.DefaultEstimateWarmConfig.Contexts).To(ContainElements(
|
||||
uint32(8192), uint32(16384), uint32(32768), uint32(65536), uint32(131072), uint32(262144),
|
||||
))
|
||||
})
|
||||
|
||||
It("bounds concurrency so a warm-up cannot saturate the link", func() {
|
||||
Expect(gallery.DefaultEstimateWarmConfig.Concurrency).To(BeNumerically("<=", 8))
|
||||
Expect(gallery.DefaultEstimateWarmConfig.Concurrency).To(BeNumerically(">", 0))
|
||||
})
|
||||
|
||||
})
|
||||
@@ -325,32 +325,10 @@ func AvailableGalleryModels(galleries []config.Gallery, systemState *system.Syst
|
||||
var (
|
||||
availableModelsMu sync.RWMutex
|
||||
availableModelsCache GalleryElements[*GalleryModel]
|
||||
// Whether a load has happened, tracked apart from the slice itself. A
|
||||
// gallery that legitimately holds nothing caches as an empty (often nil)
|
||||
// slice, and testing the slice for nil read that as "never loaded": every
|
||||
// call then took the blocking path and bumped the generation, which is the
|
||||
// same cache-defeating loop the refresh interval exists to stop.
|
||||
availableModelsLoaded bool
|
||||
refreshing atomic.Bool
|
||||
galleryGeneration atomic.Uint64
|
||||
lastRefreshUnixNano atomic.Int64
|
||||
refreshing atomic.Bool
|
||||
galleryGeneration atomic.Uint64
|
||||
)
|
||||
|
||||
// How often the cached model list may be refreshed from upstream.
|
||||
//
|
||||
// This is a floor on refresh frequency, not a TTL: the cache is served
|
||||
// regardless, and this only decides how often a background re-fetch is worth
|
||||
// starting. It matters far more than it looks, because a refresh bumps
|
||||
// galleryGeneration, and that invalidates every VRAM estimate cache in
|
||||
// pkg/vram. Refreshing on every call therefore kept those caches permanently
|
||||
// cold: the gallery listing is one request but the UI asks for one VRAM
|
||||
// estimate per row, so a single page view triggered dozens of refreshes and
|
||||
// every estimate paid full price for a remote probe it had already made.
|
||||
//
|
||||
// A package variable rather than a constant so tests can drive refreshes
|
||||
// without waiting.
|
||||
var GalleryRefreshInterval = 5 * time.Minute
|
||||
|
||||
// GalleryGeneration returns a counter that increments each time the gallery
|
||||
// model list is refreshed from upstream. VRAM estimation caches use this to
|
||||
// invalidate entries when the gallery data changes.
|
||||
@@ -374,11 +352,7 @@ func ResetGalleryModelCache() {
|
||||
}
|
||||
availableModelsMu.Lock()
|
||||
availableModelsCache = nil
|
||||
availableModelsLoaded = false
|
||||
availableModelsMu.Unlock()
|
||||
// Also clear the refresh stamp, or a suite that reset the cache would find
|
||||
// the next refresh throttled by the previous spec's clock.
|
||||
lastRefreshUnixNano.Store(0)
|
||||
}
|
||||
|
||||
// AvailableGalleryModelsCached returns gallery models from an in-memory cache.
|
||||
@@ -389,10 +363,9 @@ func ResetGalleryModelCache() {
|
||||
func AvailableGalleryModelsCached(galleries []config.Gallery, systemState *system.SystemState) (GalleryElements[*GalleryModel], error) {
|
||||
availableModelsMu.RLock()
|
||||
cached := availableModelsCache
|
||||
loaded := availableModelsLoaded
|
||||
availableModelsMu.RUnlock()
|
||||
|
||||
if loaded {
|
||||
if cached != nil {
|
||||
// Refresh installed status under write lock to avoid races with
|
||||
// concurrent readers and the background refresh goroutine.
|
||||
availableModelsMu.Lock()
|
||||
@@ -414,10 +387,8 @@ func AvailableGalleryModelsCached(galleries []config.Gallery, systemState *syste
|
||||
|
||||
availableModelsMu.Lock()
|
||||
availableModelsCache = models
|
||||
availableModelsLoaded = true
|
||||
galleryGeneration.Add(1)
|
||||
availableModelsMu.Unlock()
|
||||
lastRefreshUnixNano.Store(time.Now().UnixNano())
|
||||
|
||||
return models, nil
|
||||
}
|
||||
@@ -426,18 +397,9 @@ func AvailableGalleryModelsCached(galleries []config.Gallery, systemState *syste
|
||||
// gallery model cache. Only one refresh runs at a time; concurrent calls
|
||||
// are no-ops.
|
||||
func triggerGalleryRefresh(galleries []config.Gallery, systemState *system.SystemState) {
|
||||
if GalleryRefreshInterval > 0 {
|
||||
last := lastRefreshUnixNano.Load()
|
||||
if last != 0 && time.Since(time.Unix(0, last)) < GalleryRefreshInterval {
|
||||
return
|
||||
}
|
||||
}
|
||||
if !refreshing.CompareAndSwap(false, true) {
|
||||
return
|
||||
}
|
||||
// Stamped before the fetch rather than after, so a slow upstream cannot
|
||||
// let a queue of callers each start their own refresh behind this one.
|
||||
lastRefreshUnixNano.Store(time.Now().UnixNano())
|
||||
go func() {
|
||||
defer refreshing.Store(false)
|
||||
models, err := AvailableGalleryModels(galleries, systemState)
|
||||
@@ -446,37 +408,12 @@ func triggerGalleryRefresh(galleries []config.Gallery, systemState *system.Syste
|
||||
return
|
||||
}
|
||||
availableModelsMu.Lock()
|
||||
changed := !sameModelSet(availableModelsCache, models)
|
||||
availableModelsCache = models
|
||||
availableModelsLoaded = true
|
||||
// Only a real change invalidates the VRAM caches. An unchanged gallery
|
||||
// re-fetched on schedule must not throw away work that is still valid,
|
||||
// which is the difference between an estimate costing nothing and
|
||||
// costing a network round trip.
|
||||
if changed {
|
||||
galleryGeneration.Add(1)
|
||||
}
|
||||
galleryGeneration.Add(1)
|
||||
availableModelsMu.Unlock()
|
||||
}()
|
||||
}
|
||||
|
||||
// sameModelSet reports whether two model lists describe the same gallery, for
|
||||
// the purpose of deciding whether derived caches are still valid. Names and
|
||||
// order are enough: a change to an entry's files or size arrives with a new
|
||||
// gallery index, and comparing every field on every entry would cost more than
|
||||
// the caches save.
|
||||
func sameModelSet(a, b GalleryElements[*GalleryModel]) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := range a {
|
||||
if a[i].GetName() != b[i].GetName() {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// List available backends
|
||||
func AvailableBackends(galleries []config.Gallery, systemState *system.SystemState) (GalleryElements[*GalleryBackend], error) {
|
||||
return availableBackendsWithFilter(galleries, systemState, func(backend *GalleryBackend) bool {
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
package gallery_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
)
|
||||
|
||||
// The gallery generation counter is what every VRAM estimate cache keys on, so
|
||||
// how often it moves decides whether those caches are worth having. Refreshing
|
||||
// on every call kept them permanently cold: one page of the model gallery asks
|
||||
// for a VRAM estimate per row, and each of those requests re-read the gallery,
|
||||
// triggering a refresh that invalidated the estimate the previous row had just
|
||||
// paid a network round trip for.
|
||||
var _ = Describe("Gallery refresh throttling", func() {
|
||||
var (
|
||||
tmp *system.SystemState
|
||||
galleries []config.Gallery
|
||||
origInterval time.Duration
|
||||
)
|
||||
|
||||
BeforeEach(func() {
|
||||
dir, err := os.MkdirTemp("", "gallery-throttle")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(func() { os.RemoveAll(dir) })
|
||||
|
||||
tmp, err = system.GetSystemState(system.WithModelPath(dir))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
// No upstream: the list comes back empty, which is all this needs. What
|
||||
// is under test is how often a refresh is started, not what it returns.
|
||||
galleries = []config.Gallery{}
|
||||
origInterval = gallery.GalleryRefreshInterval
|
||||
gallery.ResetGalleryModelCache()
|
||||
})
|
||||
|
||||
AfterEach(func() {
|
||||
gallery.GalleryRefreshInterval = origInterval
|
||||
gallery.ResetGalleryModelCache()
|
||||
})
|
||||
|
||||
It("does not bump the generation once per call", func() {
|
||||
gallery.GalleryRefreshInterval = time.Hour
|
||||
|
||||
_, err := gallery.AvailableGalleryModelsCached(galleries, tmp)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
start := gallery.GalleryGeneration()
|
||||
|
||||
// Stands in for one page view: many callers in quick succession.
|
||||
for i := 0; i < 30; i++ {
|
||||
_, err := gallery.AvailableGalleryModelsCached(galleries, tmp)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
}
|
||||
// Let any refresh that did start finish, so this cannot pass by racing.
|
||||
Eventually(func() uint64 { return gallery.GalleryGeneration() }, "2s", "50ms").
|
||||
Should(Equal(start))
|
||||
})
|
||||
|
||||
It("still refreshes once the interval has passed", func() {
|
||||
gallery.GalleryRefreshInterval = time.Millisecond
|
||||
|
||||
_, err := gallery.AvailableGalleryModelsCached(galleries, tmp)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
_, err = gallery.AvailableGalleryModelsCached(galleries, tmp)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
// An empty gallery refreshing to an empty gallery is unchanged, so the
|
||||
// generation must hold: only a real change may invalidate the caches.
|
||||
Consistently(func() uint64 { return gallery.GalleryGeneration() }, "300ms", "50ms").
|
||||
Should(Equal(gallery.GalleryGeneration()))
|
||||
})
|
||||
})
|
||||
@@ -298,15 +298,7 @@ func (i *LlamaCPPImporter) Import(details Details) (gallery.ModelConfig, error)
|
||||
// imported configs already carry spec_type:draft-mtp before the model is
|
||||
// ever loaded - users see it in the YAML preview rather than discovering
|
||||
// it after the first start.
|
||||
//
|
||||
// vllm-cpp is excluded on both counts: `spec_type:*` are llama.cpp option
|
||||
// keys it does not read, and vllm.cpp rejects an MTP config over a GGUF
|
||||
// source outright (the `mtp.*` draft tensors exist only in the safetensors
|
||||
// checkpoint). Its MTP auto-config runs in the vllm importer instead, over
|
||||
// the safetensors config.json.
|
||||
if backend != "vllm-cpp" {
|
||||
maybeApplyMTPDefaults(&modelConfig, details, &cfg)
|
||||
}
|
||||
maybeApplyMTPDefaults(&modelConfig, details, &cfg)
|
||||
|
||||
data, err := yaml.Marshal(modelConfig)
|
||||
if err != nil {
|
||||
|
||||
@@ -3,7 +3,6 @@ package importers
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
@@ -32,7 +31,7 @@ func (i *MLXImporter) Match(details Details) bool {
|
||||
}
|
||||
|
||||
b, ok := preferencesMap["backend"].(string)
|
||||
if ok && slices.Contains([]string{"mlx", "mlx-vlm", "mlx-audio"}, b) {
|
||||
if ok && b == "mlx" || b == "mlx-vlm" {
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -72,32 +71,19 @@ func (i *MLXImporter) Import(details Details) (gallery.ModelConfig, error) {
|
||||
// (issue #10269). Send them to the mlx-vlm backend, which applies the
|
||||
// processor-aware chat template.
|
||||
backend := "mlx"
|
||||
usecases := []string{config.UsecaseChat}
|
||||
useTokenizerTemplate := true
|
||||
if details.HuggingFace != nil {
|
||||
switch details.HuggingFace.PipelineTag {
|
||||
case "image-text-to-text":
|
||||
backend = "mlx-vlm"
|
||||
case "text-to-speech":
|
||||
backend = "mlx-audio"
|
||||
usecases = []string{config.UsecaseTTS}
|
||||
useTokenizerTemplate = false
|
||||
}
|
||||
if details.HuggingFace != nil && details.HuggingFace.PipelineTag == "image-text-to-text" {
|
||||
backend = "mlx-vlm"
|
||||
}
|
||||
// An explicit backend preference always wins.
|
||||
b, ok := preferencesMap["backend"].(string)
|
||||
if ok {
|
||||
backend = b
|
||||
if backend == "mlx-audio" {
|
||||
usecases = []string{config.UsecaseTTS}
|
||||
useTokenizerTemplate = false
|
||||
}
|
||||
}
|
||||
|
||||
modelConfig := config.ModelConfig{
|
||||
Name: name,
|
||||
Description: description,
|
||||
KnownUsecaseStrings: usecases,
|
||||
KnownUsecaseStrings: []string{config.UsecaseChat},
|
||||
Backend: backend,
|
||||
PredictionOptions: schema.PredictionOptions{
|
||||
BasicModelRequest: schema.BasicModelRequest{
|
||||
@@ -105,7 +91,7 @@ func (i *MLXImporter) Import(details Details) (gallery.ModelConfig, error) {
|
||||
},
|
||||
},
|
||||
TemplateConfig: config.TemplateConfig{
|
||||
UseTokenizerTemplate: useTokenizerTemplate,
|
||||
UseTokenizerTemplate: true,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -48,16 +48,6 @@ var _ = Describe("MLXImporter", func() {
|
||||
Expect(result).To(BeTrue())
|
||||
})
|
||||
|
||||
It("should match when backend preference is mlx-audio", func() {
|
||||
preferences := json.RawMessage(`{"backend": "mlx-audio"}`)
|
||||
details := importers.Details{
|
||||
URI: "https://example.com/model",
|
||||
Preferences: preferences,
|
||||
}
|
||||
|
||||
Expect(importer.Match(details)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("should not match when URI does not contain mlx-community/ and no backend preference", func() {
|
||||
details := importers.Details{
|
||||
URI: "https://huggingface.co/other-org/test-model",
|
||||
@@ -133,21 +123,6 @@ var _ = Describe("MLXImporter", func() {
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("backend: mlx-vlm"))
|
||||
})
|
||||
|
||||
It("should configure explicit mlx-audio imports for text-to-speech", func() {
|
||||
preferences := json.RawMessage(`{"backend": "mlx-audio"}`)
|
||||
details := importers.Details{
|
||||
URI: "https://huggingface.co/mlx-community/Kokoro-82M-4bit",
|
||||
Preferences: preferences,
|
||||
}
|
||||
|
||||
modelConfig, err := importer.Import(details)
|
||||
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("backend: mlx-audio"))
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("- tts"))
|
||||
Expect(modelConfig.ConfigFile).ToNot(ContainSubstring("use_tokenizer_template: true"))
|
||||
})
|
||||
|
||||
It("should auto-route vision-language models to the mlx-vlm backend", func() {
|
||||
// gemma-4 E4B and similar VLMs declare pipeline_tag
|
||||
// "image-text-to-text" on HuggingFace. The text-only mlx-lm
|
||||
@@ -168,23 +143,6 @@ var _ = Describe("MLXImporter", func() {
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("backend: mlx-vlm"))
|
||||
})
|
||||
|
||||
It("should auto-route text-to-speech models to the mlx-audio backend", func() {
|
||||
details := importers.Details{
|
||||
URI: "https://huggingface.co/mlx-community/Kokoro-82M-4bit",
|
||||
HuggingFace: &hfapi.ModelDetails{
|
||||
ModelID: "mlx-community/Kokoro-82M-4bit",
|
||||
PipelineTag: "text-to-speech",
|
||||
},
|
||||
}
|
||||
|
||||
modelConfig, err := importer.Import(details)
|
||||
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("backend: mlx-audio"))
|
||||
Expect(modelConfig.ConfigFile).To(ContainSubstring("- tts"))
|
||||
Expect(modelConfig.ConfigFile).ToNot(ContainSubstring("use_tokenizer_template: true"))
|
||||
})
|
||||
|
||||
It("should keep text-only models on the plain mlx backend", func() {
|
||||
details := importers.Details{
|
||||
URI: "https://huggingface.co/mlx-community/Llama-3.2-1B-Instruct-4bit",
|
||||
|
||||
@@ -1,21 +1,13 @@
|
||||
package importers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/pkg/downloader"
|
||||
"github.com/mudler/LocalAI/pkg/httpclient"
|
||||
"github.com/mudler/xlog"
|
||||
"go.yaml.in/yaml/v2"
|
||||
)
|
||||
|
||||
@@ -115,12 +107,6 @@ func (i *VLLMImporter) Import(details Details) (gallery.ModelConfig, error) {
|
||||
// vllm python backend, so use_tokenizer_template carries over), but
|
||||
// tool/reasoning parsing is the engine's own autoparser pipeline -
|
||||
// the vllm-python tool_parser/reasoning_parser options don't apply.
|
||||
//
|
||||
// Auto-detect a Multi-Token Prediction head, the safetensors analogue
|
||||
// of the llama-cpp importer's GGUF hook, so a freshly imported
|
||||
// Qwen3.5 / Qwen3.6 config already carries speculative decoding in its
|
||||
// engine_args instead of leaving the throughput on the table.
|
||||
maybeApplyVLLMSpeculativeDefaults(&modelConfig, details)
|
||||
} else {
|
||||
// Auto-detect tool_parser and reasoning_parser for known model families.
|
||||
// Surfacing them in the generated YAML lets users see and edit the choices.
|
||||
@@ -146,89 +132,3 @@ func (i *VLLMImporter) Import(details Details) (gallery.ModelConfig, error) {
|
||||
ConfigFile: string(data),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// maxSpecConfigProbeBytes caps the config.json body we read. Real ones are a
|
||||
// few KB; the cap keeps a hostile or mislabelled URL from streaming into the
|
||||
// importer.
|
||||
const maxSpecConfigProbeBytes = 1 << 20 // 1 MiB
|
||||
|
||||
// specConfigProbeTimeout bounds the config.json fetch. Detection is an
|
||||
// optimisation, so it must never hold an import open for long.
|
||||
const specConfigProbeTimeout = 30 * time.Second
|
||||
|
||||
// specConfigFetcher is the seam the config.json probe goes through, so tests can
|
||||
// drive the whole import path without a network round trip.
|
||||
var specConfigFetcher = fetchProbeBody
|
||||
|
||||
// maybeApplyVLLMSpeculativeDefaults fetches the repository's config.json and,
|
||||
// when it declares a Multi-Token Prediction head, enables MTP speculative
|
||||
// decoding in the emitted engine_args. This is the safetensors counterpart of
|
||||
// the llama-cpp importer's GGUF header probe.
|
||||
//
|
||||
// Every failure is non-fatal and logged at debug: a network blip, a private
|
||||
// repo, or a config.json this doesn't understand must leave the import working
|
||||
// exactly as it did before, just without the speculative default.
|
||||
func maybeApplyVLLMSpeculativeDefaults(modelConfig *config.ModelConfig, details Details) {
|
||||
probeURL := vllmSpecProbeURL(details)
|
||||
if probeURL == "" {
|
||||
return
|
||||
}
|
||||
|
||||
body, err := specConfigFetcher(probeURL)
|
||||
if err != nil {
|
||||
xlog.Debug("[vllm-spec-importer] could not read config.json for MTP detection", "uri", probeURL, "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
applySpecFromConfigJSON(modelConfig, body, details.URI)
|
||||
}
|
||||
|
||||
// applySpecFromConfigJSON is the decision half of the probe, split out so it can
|
||||
// be exercised without a network round trip.
|
||||
func applySpecFromConfigJSON(modelConfig *config.ModelConfig, body []byte, uri string) {
|
||||
if config.IsDFlashDraftConfig(body) {
|
||||
// A DFlash draft cannot serve on its own - it only proposes tokens for
|
||||
// a target model to verify. Say so rather than emitting a config that
|
||||
// would fail at load.
|
||||
xlog.Warn("[vllm-spec-importer] this repository is a DFlash DRAFT checkpoint, not a servable model; "+
|
||||
"import the TARGET model and point engine_args.speculative_config at this repo "+
|
||||
`({"method":"dflash","model":"<this repo>"})`, "uri", uri)
|
||||
return
|
||||
}
|
||||
|
||||
n, ok := config.HasSafetensorsMTPHead(body)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
config.ApplyVLLMSpeculativeDefaults(modelConfig, n)
|
||||
}
|
||||
|
||||
// vllmSpecProbeURL returns the HTTP(S) URL of the repository's config.json, or
|
||||
// "" when the import isn't backed by a HuggingFace repo we can fetch from (a
|
||||
// local directory import, an OCI artifact, ...).
|
||||
func vllmSpecProbeURL(details Details) string {
|
||||
if details.HuggingFace == nil || details.HuggingFace.ModelID == "" {
|
||||
return ""
|
||||
}
|
||||
return resolveHTTPProbe(downloader.HuggingFacePrefix + details.HuggingFace.ModelID + "/config.json")
|
||||
}
|
||||
|
||||
// fetchProbeBody GETs a small remote JSON document under a short timeout.
|
||||
func fetchProbeBody(url string) ([]byte, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), specConfigProbeTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := httpclient.NewWithTimeout(specConfigProbeTimeout).Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("unexpected status %d", resp.StatusCode)
|
||||
}
|
||||
return io.ReadAll(io.LimitReader(resp.Body, maxSpecConfigProbeBytes))
|
||||
}
|
||||
|
||||
@@ -1,118 +0,0 @@
|
||||
package importers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
hfapi "github.com/mudler/LocalAI/pkg/huggingface-api"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("vllm-cpp speculative auto-config (importer)", func() {
|
||||
Context("applySpecFromConfigJSON", func() {
|
||||
It("enables mtp when the checkpoint declares an MTP head", func() {
|
||||
cfg := &config.ModelConfig{Name: "qwen3.5"}
|
||||
applySpecFromConfigJSON(cfg, []byte(`{
|
||||
"model_type": "qwen3_5_moe",
|
||||
"mtp_num_hidden_layers": 1
|
||||
}`), "huggingface://Qwen/Qwen3.5-A3B")
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("speculative_config",
|
||||
map[string]any{"method": "mtp"}))
|
||||
})
|
||||
|
||||
It("leaves a plain checkpoint untouched", func() {
|
||||
cfg := &config.ModelConfig{Name: "llama"}
|
||||
applySpecFromConfigJSON(cfg, []byte(`{"model_type": "llama"}`), "huggingface://meta/llama")
|
||||
Expect(cfg.EngineArgs).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("refuses to configure a DFlash draft as a servable model", func() {
|
||||
// The draft only proposes tokens; configuring it standalone would
|
||||
// produce a model that cannot load.
|
||||
cfg := &config.ModelConfig{Name: "dflash-draft"}
|
||||
applySpecFromConfigJSON(cfg, []byte(`{
|
||||
"model_type": "qwen3_dflash",
|
||||
"dflash_config": {"mask_token_id": 151666, "target_layer_ids": [0, 1]}
|
||||
}`), "huggingface://z-lab/Qwen3.6-27B-DFlash")
|
||||
Expect(cfg.EngineArgs).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("survives a config.json it cannot parse", func() {
|
||||
cfg := &config.ModelConfig{Name: "weird"}
|
||||
Expect(func() {
|
||||
applySpecFromConfigJSON(cfg, []byte(`<html>404</html>`), "huggingface://a/b")
|
||||
}).ToNot(Panic())
|
||||
Expect(cfg.EngineArgs).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
Context("Import over a repository with an MTP head", func() {
|
||||
var restore func()
|
||||
|
||||
BeforeEach(func() {
|
||||
original := specConfigFetcher
|
||||
restore = func() { specConfigFetcher = original }
|
||||
})
|
||||
AfterEach(func() { restore() })
|
||||
|
||||
importWith := func(backend, configJSON string) string {
|
||||
specConfigFetcher = func(string) ([]byte, error) {
|
||||
return []byte(configJSON), nil
|
||||
}
|
||||
importer := &VLLMImporter{}
|
||||
out, err := importer.Import(Details{
|
||||
URI: "huggingface://Qwen/Qwen3.5-A3B",
|
||||
Preferences: json.RawMessage(`{"backend": "` + backend + `"}`),
|
||||
HuggingFace: &hfapi.ModelDetails{ModelID: "Qwen/Qwen3.5-A3B"},
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
return out.ConfigFile
|
||||
}
|
||||
|
||||
It("emits engine_args.speculative_config for vllm-cpp", func() {
|
||||
yaml := importWith("vllm-cpp", `{"model_type":"qwen3_5_moe","mtp_num_hidden_layers":1}`)
|
||||
Expect(yaml).To(ContainSubstring("engine_args:"))
|
||||
Expect(yaml).To(ContainSubstring("speculative_config:"))
|
||||
Expect(yaml).To(ContainSubstring("method: mtp"))
|
||||
})
|
||||
|
||||
It("emits nothing speculative for the python vllm backend", func() {
|
||||
// The python backend has its own speculative surface and its own
|
||||
// version-dependent MTP support; this hook is vllm-cpp only.
|
||||
yaml := importWith("vllm", `{"model_type":"qwen3_5_moe","mtp_num_hidden_layers":1}`)
|
||||
Expect(yaml).NotTo(ContainSubstring("speculative_config"))
|
||||
})
|
||||
|
||||
It("emits nothing speculative when the probe fails", func() {
|
||||
specConfigFetcher = func(string) ([]byte, error) {
|
||||
return nil, errors.New("network down")
|
||||
}
|
||||
importer := &VLLMImporter{}
|
||||
out, err := importer.Import(Details{
|
||||
URI: "huggingface://Qwen/Qwen3.5-A3B",
|
||||
Preferences: json.RawMessage(`{"backend": "vllm-cpp"}`),
|
||||
HuggingFace: &hfapi.ModelDetails{ModelID: "Qwen/Qwen3.5-A3B"},
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(out.ConfigFile).NotTo(ContainSubstring("speculative_config"))
|
||||
})
|
||||
})
|
||||
|
||||
Context("vllmSpecProbeURL", func() {
|
||||
It("resolves the repository's config.json to an HTTPS URL", func() {
|
||||
url := vllmSpecProbeURL(Details{
|
||||
URI: "huggingface://Qwen/Qwen3.5-A3B",
|
||||
HuggingFace: &hfapi.ModelDetails{ModelID: "Qwen/Qwen3.5-A3B"},
|
||||
})
|
||||
Expect(url).To(ContainSubstring("Qwen/Qwen3.5-A3B"))
|
||||
Expect(url).To(HaveSuffix("config.json"))
|
||||
Expect(url).To(HavePrefix("https://"))
|
||||
})
|
||||
|
||||
It("skips the probe when there is no HuggingFace repo behind the import", func() {
|
||||
Expect(vllmSpecProbeURL(Details{URI: "/models/local-dir"})).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -39,8 +39,6 @@ var RouteFeatureRegistry = []RouteFeature{
|
||||
{"POST", "/images/generations", FeatureImages},
|
||||
{"POST", "/v1/images/inpainting", FeatureImages},
|
||||
{"POST", "/images/inpainting", FeatureImages},
|
||||
{"POST", "/v1/images/upscale", FeatureImages},
|
||||
{"POST", "/images/upscale", FeatureImages},
|
||||
|
||||
// Audio transcription
|
||||
{"POST", "/v1/audio/transcriptions", FeatureAudioTranscription},
|
||||
@@ -118,10 +116,6 @@ var RouteFeatureRegistry = []RouteFeature{
|
||||
// Rerank
|
||||
{"POST", "/v1/rerank", FeatureRerank},
|
||||
|
||||
// Moderation
|
||||
{"POST", "/v1/moderations", FeatureModeration},
|
||||
{"POST", "/moderations", FeatureModeration},
|
||||
|
||||
// Stores
|
||||
{"POST", "/stores/set", FeatureStores},
|
||||
{"POST", "/stores/delete", FeatureStores},
|
||||
@@ -197,7 +191,6 @@ func APIFeatureMetas() []FeatureMeta {
|
||||
{FeatureEmbeddings, "Embeddings", true},
|
||||
{FeatureSound, "Sound Generation", true},
|
||||
{FeatureRealtime, "Realtime", true},
|
||||
{FeatureModeration, "Moderation", true},
|
||||
{FeatureRerank, "Rerank", true},
|
||||
{FeatureTokenize, "Tokenize", true},
|
||||
{FeatureMCP, "MCP", true},
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
package auth_test
|
||||
|
||||
import (
|
||||
. "github.com/mudler/LocalAI/core/http/auth"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Moderation feature registration", func() {
|
||||
It("registers both moderation routes as default-on API features", func() {
|
||||
Expect(APIFeatures).To(ContainElement(FeatureModeration))
|
||||
|
||||
patterns := []string{}
|
||||
for _, route := range RouteFeatureRegistry {
|
||||
if route.Feature == FeatureModeration {
|
||||
patterns = append(patterns, route.Pattern)
|
||||
}
|
||||
}
|
||||
Expect(patterns).To(ConsistOf("/v1/moderations", "/moderations"))
|
||||
|
||||
metas := APIFeatureMetas()
|
||||
Expect(metas).To(ContainElement(FeatureMeta{Key: FeatureModeration, Label: "Moderation", DefaultValue: true}))
|
||||
})
|
||||
})
|
||||
@@ -59,14 +59,10 @@ func ok(c echo.Context) error {
|
||||
func newAuthTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo {
|
||||
e := echo.New()
|
||||
e.Use(auth.Middleware(db, appConfig))
|
||||
if db != nil {
|
||||
e.Use(auth.RequireRouteFeature(db))
|
||||
}
|
||||
|
||||
// API routes (require auth)
|
||||
e.GET("/v1/models", ok)
|
||||
e.POST("/v1/chat/completions", ok)
|
||||
e.POST("/v1/moderations", ok)
|
||||
e.GET("/api/settings", ok)
|
||||
e.POST("/api/settings", ok)
|
||||
|
||||
@@ -85,14 +81,10 @@ func newAuthTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo
|
||||
func newAdminTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo {
|
||||
e := echo.New()
|
||||
e.Use(auth.Middleware(db, appConfig))
|
||||
if db != nil {
|
||||
e.Use(auth.RequireRouteFeature(db))
|
||||
}
|
||||
|
||||
// Regular routes
|
||||
e.GET("/v1/models", ok)
|
||||
e.POST("/v1/chat/completions", ok)
|
||||
e.POST("/v1/moderations", ok)
|
||||
|
||||
// Admin-only routes
|
||||
adminMw := auth.RequireAdmin()
|
||||
|
||||
@@ -91,19 +91,6 @@ var _ = Describe("Auth Middleware", func() {
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
})
|
||||
|
||||
It("allows authenticated users to call moderation by default", func() {
|
||||
sessionID := createTestSession(db, user.ID)
|
||||
rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID))
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
})
|
||||
|
||||
It("blocks moderation when the user's feature is disabled", func() {
|
||||
Expect(auth.UpdateUserPermissions(db, user.ID, auth.PermissionMap{auth.FeatureModeration: false})).To(Succeed())
|
||||
sessionID := createTestSession(db, user.ID)
|
||||
rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID))
|
||||
Expect(rec.Code).To(Equal(http.StatusForbidden))
|
||||
})
|
||||
|
||||
It("allows requests with valid session as Bearer token", func() {
|
||||
sessionID := createTestSession(db, user.ID)
|
||||
rec := doRequest(app, http.MethodGet, "/v1/models", withBearerToken(sessionID))
|
||||
@@ -169,11 +156,6 @@ var _ = Describe("Auth Middleware", func() {
|
||||
Expect(rec.Code).To(Equal(http.StatusUnauthorized))
|
||||
})
|
||||
|
||||
It("returns 401 for unauthenticated moderation requests", func() {
|
||||
rec := doRequest(app, http.MethodPost, "/v1/moderations")
|
||||
Expect(rec.Code).To(Equal(http.StatusUnauthorized))
|
||||
})
|
||||
|
||||
It("returns 401 for unauthenticated 3D generation requests", func() {
|
||||
rec := doRequest(app, http.MethodPost, "/3d/generations")
|
||||
Expect(rec.Code).To(Equal(http.StatusUnauthorized))
|
||||
|
||||
@@ -51,7 +51,6 @@ const (
|
||||
FeatureEmbeddings = "embeddings"
|
||||
FeatureSound = "sound"
|
||||
FeatureRealtime = "realtime"
|
||||
FeatureModeration = "moderation"
|
||||
FeatureRerank = "rerank"
|
||||
FeatureTokenize = "tokenize"
|
||||
FeatureMCP = "mcp"
|
||||
@@ -76,7 +75,7 @@ var APIFeatures = []string{
|
||||
FeatureChat, FeatureImages, FeatureAudioSpeech, FeatureAudioTranscription,
|
||||
FeatureAudioDiarization, FeatureAudioClassification,
|
||||
FeatureVAD, FeatureDetection, FeatureVideo, Feature3D, FeatureEmbeddings, FeatureSound,
|
||||
FeatureRealtime, FeatureModeration, FeatureRerank, FeatureTokenize, FeatureMCP, FeatureStores,
|
||||
FeatureRealtime, FeatureRerank, FeatureTokenize, FeatureMCP, FeatureStores,
|
||||
FeatureFaceRecognition, FeatureVoiceRecognition, FeatureAudioTransform,
|
||||
FeaturePIIFilter,
|
||||
}
|
||||
|
||||
@@ -30,12 +30,6 @@ var instructionDefs = []instructionDef{
|
||||
Tags: []string{"inference", "embeddings"},
|
||||
Intro: "Set \"stream\": true for SSE streaming. Supports tool/function calling when the model config has function templates configured.",
|
||||
},
|
||||
{
|
||||
Name: "moderation",
|
||||
Description: "OpenAI-compatible text moderation using a local completion model",
|
||||
Tags: []string{"moderation"},
|
||||
Intro: "POST /v1/moderations accepts a text string or array plus a LocalAI completion model. LocalAI constrains the model to the OpenAI moderation category schema and returns one result per input. Multimodal moderation inputs are not yet supported.",
|
||||
},
|
||||
{
|
||||
Name: "audio",
|
||||
Description: "Text-to-speech, voice activity detection, transcription, speaker diarization, sound classification, and sound generation",
|
||||
|
||||
@@ -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(18))
|
||||
|
||||
// Verify each instruction has required fields and correct URL format
|
||||
for _, s := range instructions {
|
||||
@@ -69,7 +69,6 @@ var _ = Describe("API Instructions Endpoints", func() {
|
||||
|
||||
Expect(names).To(ContainElements(
|
||||
"chat-inference",
|
||||
"moderation",
|
||||
"config-management",
|
||||
"model-management",
|
||||
"monitoring",
|
||||
|
||||
@@ -38,7 +38,6 @@ var knownPrefOnlyBackends = []schema.KnownBackend{
|
||||
{Name: "whisperx", Modality: "asr", AutoDetect: false, Description: "WhisperX transcription (preference-only)"},
|
||||
{Name: "crispasr", Modality: "asr", AutoDetect: false, Description: "CrispASR multi-architecture transcription (preference-only)"},
|
||||
// TTS
|
||||
{Name: "mlx-audio", Modality: "tts", AutoDetect: false, Description: "MLX-Audio text-to-speech models (auto-detected; pref-only fallback)"},
|
||||
{Name: "kokoros", Modality: "tts", AutoDetect: false, Description: "Kokoros TTS (preference-only)"},
|
||||
{Name: "qwen-tts", Modality: "tts", AutoDetect: false, Description: "Qwen TTS (preference-only)"},
|
||||
{Name: "qwen3-tts-cpp", Modality: "tts", AutoDetect: false, Description: "Qwen3 TTS C++ (preference-only)"},
|
||||
|
||||
@@ -152,7 +152,6 @@ var _ = Describe("Backend Endpoints", func() {
|
||||
expectPrefOnly("tinygrad", "text")
|
||||
expectPrefOnly("trl", "text")
|
||||
expectPrefOnly("mlx-vlm", "text")
|
||||
expectPrefOnly("mlx-audio", "tts")
|
||||
expectPrefOnly("whisperx", "asr")
|
||||
expectPrefOnly("crispasr", "asr")
|
||||
expectPrefOnly("kokoros", "tts")
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
// @Tags monitoring
|
||||
// @Success 200 {object} schema.SystemInformationResponse "Response"
|
||||
// @Router /system [get]
|
||||
func SystemInformations(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) echo.HandlerFunc {
|
||||
func SystemInformations(ml *model.ModelLoader, appConfig *config.ApplicationConfig) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
availableBackends := []string{}
|
||||
loadedModels := ml.ListLoadedModels()
|
||||
@@ -25,14 +25,7 @@ func SystemInformations(cl *config.ModelConfigLoader, ml *model.ModelLoader, app
|
||||
|
||||
sysmodels := []schema.SysInfoModel{}
|
||||
for _, m := range loadedModels {
|
||||
entry := schema.SysInfoModel{ID: m.ID}
|
||||
// The loader tracks only the ID. Which engine is serving a model is
|
||||
// the first thing an operator wants beside its name, and it is one
|
||||
// config lookup away.
|
||||
if cfg, ok := cl.GetModelConfig(m.ID); ok {
|
||||
entry.Backend = cfg.Backend
|
||||
}
|
||||
sysmodels = append(sysmodels, entry)
|
||||
sysmodels = append(sysmodels, schema.SysInfoModel{ID: m.ID})
|
||||
}
|
||||
return c.JSON(200,
|
||||
schema.SystemInformationResponse{
|
||||
|
||||
@@ -3,7 +3,6 @@ package localai
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
@@ -86,35 +85,6 @@ func GetAPITracesEndpoint() echo.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// GetAPITracesSummaryEndpoint returns counted totals over a recent window
|
||||
// @Summary Summarize recent API traces
|
||||
// @Description Returns request, failure and latency totals over a recent window, plus a bucketed series for sparklines. Exists so callers wanting three numbers do not have to fetch the whole trace list and count it themselves.
|
||||
// @Tags monitoring
|
||||
// @Produce json
|
||||
// @Param hours query int false "Window in hours (default 24, max 168)"
|
||||
// @Success 200 {object} middleware.TraceSummary "Counted trace totals"
|
||||
// @Router /api/traces/summary [get]
|
||||
func GetAPITracesSummaryEndpoint() echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
hours := 24
|
||||
if raw := c.QueryParam("hours"); raw != "" {
|
||||
if v, err := strconv.Atoi(raw); err == nil && v > 0 {
|
||||
hours = v
|
||||
}
|
||||
}
|
||||
// A week is plenty for a dashboard, and the trace buffer is bounded
|
||||
// anyway; an unbounded window would just scan the whole buffer.
|
||||
if hours > 168 {
|
||||
hours = 168
|
||||
}
|
||||
return c.JSON(http.StatusOK, middleware.GetTracesSummary(time.Duration(hours)*time.Hour, traceSummaryBuckets))
|
||||
}
|
||||
}
|
||||
|
||||
// Enough columns for a sparkline to show a shape, few enough that each one
|
||||
// still holds a meaningful count on a quiet installation.
|
||||
const traceSummaryBuckets = 12
|
||||
|
||||
// GetAPITraceEndpoint returns a single API trace with its full payload
|
||||
// @Summary Get one API trace
|
||||
// @Description Returns a single captured API exchange, including the request and response bodies omitted from the list response
|
||||
|
||||
@@ -84,22 +84,6 @@ func (stubClient) ListNodes(_ context.Context) ([]localaitools.Node, error) {
|
||||
return []localaitools.Node{}, nil
|
||||
}
|
||||
|
||||
func (stubClient) ListScheduling(_ context.Context) ([]localaitools.ModelSchedulingConfig, error) {
|
||||
return []localaitools.ModelSchedulingConfig{}, nil
|
||||
}
|
||||
|
||||
func (stubClient) GetScheduling(_ context.Context, _ string) (*localaitools.ModelSchedulingConfig, error) {
|
||||
return &localaitools.ModelSchedulingConfig{}, nil
|
||||
}
|
||||
|
||||
func (stubClient) SetScheduling(_ context.Context, _ localaitools.SetSchedulingRequest) (*localaitools.ModelSchedulingConfig, error) {
|
||||
return &localaitools.ModelSchedulingConfig{}, nil
|
||||
}
|
||||
|
||||
func (stubClient) DeleteScheduling(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (stubClient) SetNodeVRAMBudget(_ context.Context, _, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,190 +0,0 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/templates"
|
||||
"github.com/mudler/LocalAI/pkg/functions"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
)
|
||||
|
||||
var moderationCategories = []string{
|
||||
"harassment",
|
||||
"harassment/threatening",
|
||||
"hate",
|
||||
"hate/threatening",
|
||||
"illicit",
|
||||
"illicit/violent",
|
||||
"self-harm",
|
||||
"self-harm/intent",
|
||||
"self-harm/instructions",
|
||||
"sexual",
|
||||
"sexual/minors",
|
||||
"violence",
|
||||
"violence/graphic",
|
||||
}
|
||||
|
||||
type moderationGenerator func(context.Context, string, *config.ModelConfig) (string, backend.TokenUsage, error)
|
||||
|
||||
type generatedModeration struct {
|
||||
Categories map[string]bool `json:"categories"`
|
||||
CategoryScores map[string]float64 `json:"category_scores"`
|
||||
}
|
||||
|
||||
// ModerationEndpoint implements the text input subset of OpenAI's moderation
|
||||
// API using any LocalAI completion model and constrained JSON generation.
|
||||
// @Summary Classify text for potentially harmful content.
|
||||
// @Tags moderation
|
||||
// @Param request body schema.ModerationRequest true "query params"
|
||||
// @Success 200 {object} schema.ModerationResponse "Response"
|
||||
// @Router /v1/moderations [post]
|
||||
func ModerationEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig) echo.HandlerFunc {
|
||||
return moderationEndpoint(func(ctx context.Context, input string, cfg *config.ModelConfig) (string, backend.TokenUsage, error) {
|
||||
prompt := moderationPrompt(input)
|
||||
var messages schema.Messages
|
||||
if cfg.TemplateConfig.UseTokenizerTemplate {
|
||||
messages = schema.Messages{{Role: "user", Content: prompt}}
|
||||
prompt = ""
|
||||
} else if evaluator != nil {
|
||||
if rendered, err := evaluator.EvaluateTemplateForPrompt(templates.CompletionPromptTemplate, *cfg, templates.PromptTemplateData{Input: prompt, SystemPrompt: cfg.SystemPrompt}); err == nil {
|
||||
prompt = rendered
|
||||
}
|
||||
}
|
||||
|
||||
predict, err := backend.ModelInferenceFunc(ctx, prompt, messages, nil, nil, nil, ml, cfg, cl, appConfig, nil, "", "", nil, nil, nil, nil)
|
||||
if err != nil {
|
||||
return "", backend.TokenUsage{}, err
|
||||
}
|
||||
response, err := predict()
|
||||
return response.Response, response.Usage, err
|
||||
})
|
||||
}
|
||||
|
||||
func moderationEndpoint(generate moderationGenerator) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.ModerationRequest)
|
||||
if !ok || input == nil {
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "invalid moderation request")
|
||||
}
|
||||
if len(input.Input) == 0 {
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "input must contain at least one text string")
|
||||
}
|
||||
if generate == nil {
|
||||
return echo.NewHTTPError(http.StatusInternalServerError, "moderation generator is unavailable")
|
||||
}
|
||||
|
||||
modelConfig, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
|
||||
if !ok || modelConfig == nil {
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "moderation model configuration is unavailable")
|
||||
}
|
||||
|
||||
grammar, err := moderationGrammar()
|
||||
if err != nil {
|
||||
return echo.NewHTTPError(http.StatusInternalServerError, "failed to build moderation grammar").SetInternal(err)
|
||||
}
|
||||
cfg := *modelConfig
|
||||
cfg.Grammar = grammar
|
||||
maxTokens := 512
|
||||
cfg.Maxtokens = &maxTokens
|
||||
|
||||
results := make([]schema.ModerationResult, 0, len(input.Input))
|
||||
promptTokens, completionTokens := 0, 0
|
||||
for _, text := range input.Input {
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "input strings must not be empty")
|
||||
}
|
||||
raw, usage, err := generate(c.Request().Context(), text, &cfg)
|
||||
if err != nil {
|
||||
return echo.NewHTTPError(http.StatusInternalServerError, "moderation inference failed").SetInternal(err)
|
||||
}
|
||||
result, err := parseModerationResult(raw)
|
||||
if err != nil {
|
||||
return echo.NewHTTPError(http.StatusInternalServerError, "invalid moderation result: "+err.Error()).SetInternal(err)
|
||||
}
|
||||
results = append(results, result)
|
||||
promptTokens += usage.Prompt
|
||||
completionTokens += usage.Completion
|
||||
}
|
||||
|
||||
middleware.StampUsage(c, input.Model, promptTokens, completionTokens)
|
||||
return c.JSON(http.StatusOK, schema.ModerationResponse{
|
||||
ID: "modr-" + uuid.NewString(),
|
||||
Model: input.Model,
|
||||
Results: results,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func moderationPrompt(input string) string {
|
||||
encoded, _ := json.Marshal(input)
|
||||
return `Classify the supplied text for safety moderation. Return only the JSON object required by the grammar.
|
||||
For each category, set the boolean to true only when the text matches that category. Give each category a confidence score from 0 to 1.
|
||||
The categories are harassment, harassment/threatening, hate, hate/threatening, illicit, illicit/violent, self-harm, self-harm/intent, self-harm/instructions, sexual, sexual/minors, violence, and violence/graphic.
|
||||
Text to classify: ` + string(encoded)
|
||||
}
|
||||
|
||||
func moderationGrammar() (string, error) {
|
||||
boolProperties := map[string]any{}
|
||||
scoreProperties := map[string]any{}
|
||||
for _, category := range moderationCategories {
|
||||
boolProperties[category] = map[string]any{"type": "boolean"}
|
||||
scoreProperties[category] = map[string]any{"type": "number"}
|
||||
}
|
||||
structure := functions.JSONFunctionStructure{AnyOf: []functions.Item{{
|
||||
Type: "object",
|
||||
Properties: map[string]any{
|
||||
"categories": map[string]any{
|
||||
"type": "object",
|
||||
"properties": boolProperties,
|
||||
"required": moderationCategories,
|
||||
"additionalProperties": false,
|
||||
},
|
||||
"category_scores": map[string]any{
|
||||
"type": "object",
|
||||
"properties": scoreProperties,
|
||||
"required": moderationCategories,
|
||||
"additionalProperties": false,
|
||||
},
|
||||
},
|
||||
}}}
|
||||
return structure.Grammar()
|
||||
}
|
||||
|
||||
func parseModerationResult(raw string) (schema.ModerationResult, error) {
|
||||
var generated generatedModeration
|
||||
if err := json.Unmarshal([]byte(strings.TrimSpace(raw)), &generated); err != nil {
|
||||
return schema.ModerationResult{}, err
|
||||
}
|
||||
|
||||
result := schema.ModerationResult{
|
||||
Categories: make(map[string]bool, len(moderationCategories)),
|
||||
CategoryScores: make(map[string]float64, len(moderationCategories)),
|
||||
CategoryAppliedInputTypes: make(map[string][]string, len(moderationCategories)),
|
||||
}
|
||||
for _, category := range moderationCategories {
|
||||
flagged, exists := generated.Categories[category]
|
||||
if !exists {
|
||||
return schema.ModerationResult{}, fmt.Errorf("missing category %q", category)
|
||||
}
|
||||
score, exists := generated.CategoryScores[category]
|
||||
if !exists || math.IsNaN(score) || math.IsInf(score, 0) || score < 0 || score > 1 {
|
||||
return schema.ModerationResult{}, fmt.Errorf("category %q has an invalid score", category)
|
||||
}
|
||||
result.Categories[category] = flagged
|
||||
result.CategoryScores[category] = score
|
||||
result.CategoryAppliedInputTypes[category] = []string{"text"}
|
||||
result.Flagged = result.Flagged || flagged
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -1,105 +0,0 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"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"
|
||||
)
|
||||
|
||||
var _ = Describe("Moderations endpoint", func() {
|
||||
It("classifies each text input and returns the OpenAI response shape", func() {
|
||||
inputs := []string{}
|
||||
generate := func(_ context.Context, input string, cfg *config.ModelConfig) (string, backend.TokenUsage, error) {
|
||||
inputs = append(inputs, input)
|
||||
Expect(cfg.Grammar).To(ContainSubstring("harassment"))
|
||||
return `{
|
||||
"categories":{"harassment":true,"harassment/threatening":false,"hate":false,"hate/threatening":false,"illicit":false,"illicit/violent":false,"self-harm":false,"self-harm/intent":false,"self-harm/instructions":false,"sexual":false,"sexual/minors":false,"violence":false,"violence/graphic":false},
|
||||
"category_scores":{"harassment":0.9,"harassment/threatening":0.1,"hate":0,"hate/threatening":0,"illicit":0,"illicit/violent":0,"self-harm":0,"self-harm/intent":0,"self-harm/instructions":0,"sexual":0,"sexual/minors":0,"violence":0,"violence/graphic":0}
|
||||
}`, backend.TokenUsage{Prompt: 12, Completion: 8}, nil
|
||||
}
|
||||
|
||||
e := echo.New()
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/moderations", strings.NewReader(`{"model":"guard","input":["first","second"]}`))
|
||||
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
rec := httptest.NewRecorder()
|
||||
ctx := e.NewContext(req, rec)
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.ModerationRequest{
|
||||
BasicModelRequest: schema.BasicModelRequest{Model: "guard"},
|
||||
Input: schema.ModerationInput{"first", "second"},
|
||||
})
|
||||
modelConfig := &config.ModelConfig{Name: "guard"}
|
||||
modelConfig.Model = "guard.gguf"
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, modelConfig)
|
||||
|
||||
Expect(moderationEndpoint(generate)(ctx)).To(Succeed())
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
Expect(inputs).To(Equal([]string{"first", "second"}))
|
||||
|
||||
var response schema.ModerationResponse
|
||||
Expect(json.Unmarshal(rec.Body.Bytes(), &response)).To(Succeed())
|
||||
Expect(response.ID).To(HavePrefix("modr-"))
|
||||
Expect(response.Model).To(Equal("guard"))
|
||||
Expect(response.Results).To(HaveLen(2))
|
||||
Expect(response.Results[0].Flagged).To(BeTrue())
|
||||
Expect(response.Results[0].Categories["harassment"]).To(BeTrue())
|
||||
Expect(response.Results[0].CategoryAppliedInputTypes["harassment"]).To(Equal([]string{"text"}))
|
||||
})
|
||||
|
||||
It("rejects an empty input list", func() {
|
||||
e := echo.New()
|
||||
ctx := e.NewContext(httptest.NewRequest(http.MethodPost, "/v1/moderations", nil), httptest.NewRecorder())
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.ModerationRequest{
|
||||
BasicModelRequest: schema.BasicModelRequest{Model: "guard"},
|
||||
})
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Name: "guard"})
|
||||
|
||||
err := moderationEndpoint(nil)(ctx)
|
||||
Expect(err).To(MatchError(ContainSubstring("input must contain at least one text string")))
|
||||
Expect(err.(*echo.HTTPError).Code).To(Equal(http.StatusBadRequest))
|
||||
})
|
||||
|
||||
It("surfaces malformed classifier output without returning a partial result", func() {
|
||||
generate := func(context.Context, string, *config.ModelConfig) (string, backend.TokenUsage, error) {
|
||||
return "not-json", backend.TokenUsage{}, nil
|
||||
}
|
||||
e := echo.New()
|
||||
ctx := e.NewContext(httptest.NewRequest(http.MethodPost, "/v1/moderations", nil), httptest.NewRecorder())
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.ModerationRequest{
|
||||
BasicModelRequest: schema.BasicModelRequest{Model: "guard"},
|
||||
Input: schema.ModerationInput{"text"},
|
||||
})
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Name: "guard"})
|
||||
|
||||
err := moderationEndpoint(generate)(ctx)
|
||||
Expect(err).To(MatchError(ContainSubstring("invalid moderation result")))
|
||||
Expect(err.(*echo.HTTPError).Code).To(Equal(http.StatusInternalServerError))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Moderation input", func() {
|
||||
DescribeTable("accepts OpenAI text input forms",
|
||||
func(body string, expected schema.ModerationInput) {
|
||||
var req schema.ModerationRequest
|
||||
Expect(json.Unmarshal([]byte(body), &req)).To(Succeed())
|
||||
Expect(req.Input).To(Equal(expected))
|
||||
},
|
||||
Entry("single text", `{"input":"hello"}`, schema.ModerationInput{"hello"}),
|
||||
Entry("text array", `{"input":["hello","world"]}`, schema.ModerationInput{"hello", "world"}),
|
||||
)
|
||||
|
||||
It("rejects multimodal input in the text-only MVP", func() {
|
||||
var req schema.ModerationRequest
|
||||
err := json.Unmarshal([]byte(`{"input":[{"type":"image_url","image_url":{"url":"https://example.com/a.png"}}]}`), &req)
|
||||
Expect(err).To(MatchError(ContainSubstring("text string or array of text strings")))
|
||||
})
|
||||
})
|
||||
@@ -1,134 +0,0 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/xlog"
|
||||
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
model "github.com/mudler/LocalAI/pkg/model"
|
||||
)
|
||||
|
||||
// UpscaleEndpoint handles POST /v1/images/upscale
|
||||
//
|
||||
// @Summary Image upscaling
|
||||
// @Description Upscale an image using a specified model (e.g. stable-diffusion-x4-upscaler). Accepts multipart/form-data.
|
||||
// @Tags images
|
||||
// @Accept multipart/form-data
|
||||
// @Produce application/json
|
||||
// @Param model formData string true "Upscaler model identifier (e.g. stable-diffusion-x4-upscaler)"
|
||||
// @Param image formData file true "Input image file"
|
||||
// @Param scale formData int false "Upscale factor: 2 or 4 (default 2)"
|
||||
// @Success 200 {object} schema.OpenAIResponse
|
||||
// @Failure 400 {object} map[string]string
|
||||
// @Failure 500 {object} map[string]string
|
||||
// @Router /v1/images/upscale [post]
|
||||
func UpscaleEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
modelName := c.FormValue("model")
|
||||
scaleStr := c.FormValue("scale")
|
||||
|
||||
if modelName == "" {
|
||||
xlog.Error("Upscale Endpoint - missing model")
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "missing model")
|
||||
}
|
||||
|
||||
scale := 2
|
||||
if scaleStr != "" {
|
||||
v, err := strconv.Atoi(scaleStr)
|
||||
if err != nil || (v != 2 && v != 4) {
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "scale must be 2 or 4")
|
||||
}
|
||||
scale = v
|
||||
}
|
||||
|
||||
// Read uploaded image
|
||||
imageFile, err := c.FormFile("image")
|
||||
if err != nil {
|
||||
xlog.Error("Upscale Endpoint - missing image file", "error", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "missing image file")
|
||||
}
|
||||
|
||||
imgSrc, err := imageFile.Open()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer imgSrc.Close()
|
||||
imgBytes, err := io.ReadAll(imgSrc)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get model config from middleware context
|
||||
cfg, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
|
||||
if !ok || cfg == nil {
|
||||
xlog.Error("Upscale Endpoint - model config not found in context")
|
||||
return echo.ErrBadRequest
|
||||
}
|
||||
|
||||
tmpDir := filepath.Join(appConfig.GeneratedContentDir, "images")
|
||||
if err := os.MkdirAll(tmpDir, 0750); err != nil {
|
||||
return echo.NewHTTPError(http.StatusInternalServerError, "failed to prepare storage")
|
||||
}
|
||||
|
||||
// Write input image to a temp file
|
||||
srcTmp, err := os.CreateTemp(tmpDir, "upscale_src_")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := srcTmp.Write(imgBytes); err != nil {
|
||||
_ = srcTmp.Close()
|
||||
_ = os.Remove(srcTmp.Name())
|
||||
return err
|
||||
}
|
||||
if err := srcTmp.Close(); err != nil {
|
||||
xlog.Warn("Upscale Endpoint - failed to close src temp file", "error", err)
|
||||
}
|
||||
srcPath := srcTmp.Name()
|
||||
defer os.Remove(srcPath)
|
||||
|
||||
// Prepare output file path
|
||||
id := uuid.New().String()
|
||||
dstPath := filepath.Join(tmpDir, fmt.Sprintf("upscale_%s.png", id))
|
||||
|
||||
fn, err := backend.ImageUpscaleFunc(c.Request().Context(), srcPath, dstPath, scale, ml, *cfg, appConfig)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fn(); err != nil {
|
||||
_ = os.Remove(dstPath)
|
||||
return err
|
||||
}
|
||||
|
||||
baseURL := middleware.BaseURL(c)
|
||||
imgURL, err := url.JoinPath(baseURL, "generated-images", filepath.Base(dstPath))
|
||||
if err != nil {
|
||||
_ = os.Remove(dstPath)
|
||||
return err
|
||||
}
|
||||
|
||||
created := int(time.Now().Unix())
|
||||
resp := &schema.OpenAIResponse{
|
||||
ID: id,
|
||||
Created: created,
|
||||
Data: []schema.Item{{URL: imgURL}},
|
||||
Usage: &schema.OpenAIUsage{
|
||||
InputTokensDetails: &schema.InputTokensDetails{},
|
||||
},
|
||||
}
|
||||
|
||||
return c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
}
|
||||
@@ -1,89 +0,0 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
model "github.com/mudler/LocalAI/pkg/model"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Image upscaling", func() {
|
||||
var (
|
||||
appConfig *config.ApplicationConfig
|
||||
tmpDir string
|
||||
)
|
||||
|
||||
BeforeEach(func() {
|
||||
var err error
|
||||
tmpDir, err = os.MkdirTemp("", "upscale")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
appConfig = config.NewApplicationConfig(config.WithGeneratedContentDir(tmpDir))
|
||||
})
|
||||
|
||||
AfterEach(func() {
|
||||
Expect(os.RemoveAll(tmpDir)).To(Succeed())
|
||||
})
|
||||
|
||||
It("stores the result in the directory served by /generated-images", func() {
|
||||
original := backend.ImageUpscaleFunc
|
||||
backend.ImageUpscaleFunc = func(_ context.Context, _, dst string, scale int, _ *model.ModelLoader, _ config.ModelConfig, _ *config.ApplicationConfig) (func() error, error) {
|
||||
Expect(scale).To(Equal(4))
|
||||
return func() error {
|
||||
return os.WriteFile(dst, []byte("PNGDATA"), 0o644)
|
||||
}, nil
|
||||
}
|
||||
DeferCleanup(func() { backend.ImageUpscaleFunc = original })
|
||||
|
||||
req, _ := makeMultipartRequest(
|
||||
map[string]string{"model": "stable-diffusion-x4-upscaler", "scale": "4"},
|
||||
map[string][]byte{"image": []byte("IMAGEDATA")},
|
||||
)
|
||||
rec := httptest.NewRecorder()
|
||||
ctx := echo.New().NewContext(req, rec)
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Backend: "diffusers"})
|
||||
|
||||
Expect(UpscaleEndpoint(nil, nil, appConfig)(ctx)).To(Succeed())
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
|
||||
var response schema.OpenAIResponse
|
||||
Expect(json.Unmarshal(rec.Body.Bytes(), &response)).To(Succeed())
|
||||
Expect(response.Data).To(HaveLen(1))
|
||||
Expect(response.Data[0].URL).To(ContainSubstring("/generated-images/upscale_"))
|
||||
|
||||
filename := filepath.Base(response.Data[0].URL)
|
||||
contents, err := os.ReadFile(filepath.Join(tmpDir, "images", filename))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(contents).To(Equal([]byte("PNGDATA")))
|
||||
})
|
||||
|
||||
It("rejects unsupported scale factors", func() {
|
||||
req, _ := makeMultipartRequest(
|
||||
map[string]string{"model": "stable-diffusion-x4-upscaler", "scale": "3"},
|
||||
map[string][]byte{"image": []byte("IMAGEDATA")},
|
||||
)
|
||||
rec := httptest.NewRecorder()
|
||||
ctx := echo.New().NewContext(req, rec)
|
||||
ctx.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Backend: "diffusers"})
|
||||
|
||||
err := UpscaleEndpoint(nil, nil, appConfig)(ctx)
|
||||
var httpErr *echo.HTTPError
|
||||
Expect(err).To(MatchError(ContainSubstring("scale must be 2 or 4")))
|
||||
Expect(err).To(BeAssignableToTypeOf(httpErr))
|
||||
httpErr = err.(*echo.HTTPError)
|
||||
Expect(httpErr.Code).To(Equal(http.StatusBadRequest))
|
||||
Expect(httpErr.Message).To(Equal("scale must be 2 or 4"))
|
||||
Expect(bytes.TrimSpace(rec.Body.Bytes())).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -1,109 +0,0 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"math"
|
||||
"slices"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TraceSummary is the counted view of the trace buffer.
|
||||
//
|
||||
// It exists so a caller that wants "how many, how many failed, how slow" does
|
||||
// not have to fetch every exchange and count them in the browser. The Operate
|
||||
// overview needs exactly those three numbers, and the trace list is capped in
|
||||
// the thousands, so shipping it across the wire to produce a single integer is
|
||||
// waste that grows with the buffer.
|
||||
type TraceSummary struct {
|
||||
Total int `json:"total"`
|
||||
Errors int `json:"errors"`
|
||||
P95Millis int64 `json:"p95_ms"`
|
||||
WindowHours int `json:"window_hours"`
|
||||
Buckets []TraceBucket `json:"buckets"`
|
||||
}
|
||||
|
||||
// TraceBucket is one column of a sparkline: oldest first, so the series reads
|
||||
// left to right the way a chart is drawn.
|
||||
type TraceBucket struct {
|
||||
Start time.Time `json:"start"`
|
||||
Count int `json:"count"`
|
||||
Errors int `json:"errors"`
|
||||
}
|
||||
|
||||
// GetTracesSummary counts the buffered exchanges over the given window.
|
||||
func GetTracesSummary(window time.Duration, buckets int) TraceSummary {
|
||||
return summarize(GetTraces(), window, buckets)
|
||||
}
|
||||
|
||||
func summarize(traces []APIExchange, window time.Duration, buckets int) TraceSummary {
|
||||
if buckets < 1 {
|
||||
buckets = 1
|
||||
}
|
||||
now := time.Now()
|
||||
cutoff := now.Add(-window)
|
||||
|
||||
summary := TraceSummary{
|
||||
WindowHours: int(window.Hours()),
|
||||
// Never nil: a nil slice serialises as null and breaks .map() on the
|
||||
// other side, which is a silent runtime error rather than an empty chart.
|
||||
Buckets: make([]TraceBucket, buckets),
|
||||
}
|
||||
|
||||
bucketWidth := window / time.Duration(buckets)
|
||||
for i := range summary.Buckets {
|
||||
summary.Buckets[i].Start = cutoff.Add(time.Duration(i) * bucketWidth)
|
||||
}
|
||||
|
||||
durations := make([]time.Duration, 0, len(traces))
|
||||
for _, t := range traces {
|
||||
if t.Timestamp.Before(cutoff) {
|
||||
continue
|
||||
}
|
||||
summary.Total++
|
||||
failed := isFailure(t)
|
||||
if failed {
|
||||
summary.Errors++
|
||||
}
|
||||
durations = append(durations, t.Duration)
|
||||
|
||||
// Clamp rather than skip: a request timestamped a hair in the future
|
||||
// (clock skew, or arriving mid-call) still belongs in the newest column.
|
||||
idx := int(t.Timestamp.Sub(cutoff) / bucketWidth)
|
||||
if idx >= buckets {
|
||||
idx = buckets - 1
|
||||
}
|
||||
if idx < 0 {
|
||||
idx = 0
|
||||
}
|
||||
summary.Buckets[idx].Count++
|
||||
if failed {
|
||||
summary.Buckets[idx].Errors++
|
||||
}
|
||||
}
|
||||
|
||||
summary.P95Millis = percentileMillis(durations, 0.95)
|
||||
return summary
|
||||
}
|
||||
|
||||
// A 4xx is the caller getting it wrong, which is not the installation being
|
||||
// unhealthy. Only 5xx and a transport-level error count against the runtime.
|
||||
func isFailure(t APIExchange) bool {
|
||||
return t.Error != "" || t.Response.Status >= 500
|
||||
}
|
||||
|
||||
func percentileMillis(durations []time.Duration, p float64) int64 {
|
||||
if len(durations) == 0 {
|
||||
return 0
|
||||
}
|
||||
slices.Sort(durations)
|
||||
// Nearest-rank: the smallest value at or above the pth percentile.
|
||||
rank := int(math.Ceil(p*float64(len(durations)))) - 1
|
||||
if rank < 0 {
|
||||
rank = 0
|
||||
}
|
||||
if rank >= len(durations) {
|
||||
rank = len(durations) - 1
|
||||
}
|
||||
return durations[rank].Milliseconds()
|
||||
}
|
||||
@@ -1,79 +0,0 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("API trace summary", func() {
|
||||
exchange := func(age time.Duration, status int, dur time.Duration) APIExchange {
|
||||
return APIExchange{
|
||||
Timestamp: time.Now().Add(-age),
|
||||
Duration: dur,
|
||||
Response: APIExchangeResponse{Status: status},
|
||||
}
|
||||
}
|
||||
|
||||
It("counts only what falls inside the window", func() {
|
||||
traces := []APIExchange{
|
||||
exchange(1*time.Hour, 200, 10*time.Millisecond),
|
||||
exchange(2*time.Hour, 200, 10*time.Millisecond),
|
||||
// Older than the window: must not be counted at all.
|
||||
exchange(48*time.Hour, 500, 10*time.Millisecond),
|
||||
}
|
||||
s := summarize(traces, 24*time.Hour, 6)
|
||||
Expect(s.Total).To(Equal(2))
|
||||
Expect(s.Errors).To(BeZero())
|
||||
})
|
||||
|
||||
It("treats 5xx and a transport error as failures, but not 4xx", func() {
|
||||
traces := []APIExchange{
|
||||
exchange(time.Minute, 500, time.Millisecond),
|
||||
exchange(time.Minute, 503, time.Millisecond),
|
||||
// A client sending a bad request is not the server failing.
|
||||
exchange(time.Minute, 404, time.Millisecond),
|
||||
exchange(time.Minute, 200, time.Millisecond),
|
||||
}
|
||||
traces[3].Error = "connection reset"
|
||||
|
||||
s := summarize(traces, 24*time.Hour, 6)
|
||||
Expect(s.Total).To(Equal(4))
|
||||
Expect(s.Errors).To(Equal(3))
|
||||
})
|
||||
|
||||
It("reports p95 as a real percentile rather than the slowest request", func() {
|
||||
traces := make([]APIExchange, 0, 100)
|
||||
for i := 1; i <= 100; i++ {
|
||||
traces = append(traces, exchange(time.Minute, 200, time.Duration(i)*time.Millisecond))
|
||||
}
|
||||
s := summarize(traces, 24*time.Hour, 6)
|
||||
// 95th of 1..100ms, not the 100ms max.
|
||||
Expect(s.P95Millis).To(BeNumerically("~", 95, 1))
|
||||
})
|
||||
|
||||
It("buckets oldest-first so a sparkline reads left to right", func() {
|
||||
traces := []APIExchange{
|
||||
exchange(30*time.Minute, 200, time.Millisecond),
|
||||
exchange(30*time.Minute, 200, time.Millisecond),
|
||||
exchange(5*time.Hour, 200, time.Millisecond),
|
||||
}
|
||||
s := summarize(traces, 6*time.Hour, 6)
|
||||
Expect(s.Buckets).To(HaveLen(6))
|
||||
Expect(s.Buckets[0].Count).To(Equal(1), "the 5h-old request lands in the first bucket")
|
||||
Expect(s.Buckets[5].Count).To(Equal(2), "the recent pair lands in the last")
|
||||
})
|
||||
|
||||
It("returns an empty, non-nil summary when nothing has been traced", func() {
|
||||
s := summarize(nil, 24*time.Hour, 6)
|
||||
Expect(s.Total).To(BeZero())
|
||||
Expect(s.Errors).To(BeZero())
|
||||
Expect(s.P95Millis).To(BeZero())
|
||||
// A nil slice serialises as null and breaks .map() in the browser.
|
||||
Expect(s.Buckets).NotTo(BeNil())
|
||||
Expect(s.Buckets).To(HaveLen(6))
|
||||
})
|
||||
})
|
||||
@@ -43,40 +43,6 @@ test('lists live operations and cancels one from a labelled button', async ({ pa
|
||||
expect(cancelledPath).toBe('/api/operations/job-gemma/cancel')
|
||||
})
|
||||
|
||||
test('pauses a model download without invoking destructive cancel', async ({ page }) => {
|
||||
await stub(page, {
|
||||
operations: [{
|
||||
id: 'gemma-3-27b-it',
|
||||
name: 'gemma-3-27b-it',
|
||||
jobID: 'job-gemma',
|
||||
progress: 22,
|
||||
taskType: 'installation',
|
||||
isBackend: false,
|
||||
isQueued: false,
|
||||
isDeletion: false,
|
||||
cancellable: true,
|
||||
phase: 'downloading',
|
||||
}],
|
||||
})
|
||||
|
||||
const requests = []
|
||||
await page.route('**/api/operations/job-gemma/pause', (route) => {
|
||||
requests.push(new URL(route.request().url()).pathname)
|
||||
return route.fulfill({ contentType: 'application/json', body: '{}' })
|
||||
})
|
||||
await page.route('**/api/operations/job-gemma/cancel', (route) => {
|
||||
requests.push(new URL(route.request().url()).pathname)
|
||||
return route.fulfill({ contentType: 'application/json', body: '{}' })
|
||||
})
|
||||
|
||||
await page.goto('/app/activity')
|
||||
|
||||
const card = page.locator('.operation-card').filter({ hasText: 'gemma-3-27b-it' })
|
||||
await card.locator('.operation-card__pause').click()
|
||||
|
||||
await expect.poll(() => requests).toEqual(['/api/operations/job-gemma/pause'])
|
||||
})
|
||||
|
||||
test('separates an unacknowledged failure from the record', async ({ page }) => {
|
||||
await stub(page, {
|
||||
operations: [{
|
||||
|
||||
@@ -5,9 +5,7 @@ test.describe('Admin console', () => {
|
||||
await page.goto('/app/backends')
|
||||
const rail = page.locator('.console-rail')
|
||||
await expect(rail).toBeVisible()
|
||||
// Four groups since the overview landed: Inference folded into Runtime
|
||||
// (both are "the runtime right now"), Access and System into Administration.
|
||||
for (const group of ['Runtime', 'Cluster', 'Observability', 'Administration']) {
|
||||
for (const group of ['Inference', 'Cluster', 'Observability', 'Access', 'System']) {
|
||||
await expect(rail.locator('.console-group-title', { hasText: group })).toBeVisible()
|
||||
}
|
||||
})
|
||||
|
||||
@@ -69,9 +69,9 @@ test.describe('Manage - alias badge', () => {
|
||||
|
||||
test('renders a read-only alias -> target badge on aliased rows', async ({ page }) => {
|
||||
await page.goto('/app/manage')
|
||||
// The badge moved off the row and into the pane: it is a fact about the
|
||||
// model, and the rail line is spent on state.
|
||||
await page.locator('[data-entity="gpt-4"]').click()
|
||||
await expect(page.locator('.table')).toBeVisible({ timeout: 10_000 })
|
||||
|
||||
// The aliased row shows the target; the plain model row does not.
|
||||
await expect(page.getByText('alias -> fast-llm')).toBeVisible({ timeout: 10_000 })
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// Backends admin page (src/pages/Backends.jsx).
|
||||
const PANE = '[data-testid="backends-pane"]'
|
||||
const railItem = (page, name) => page.locator(`[data-entity="${name}"]`)
|
||||
|
||||
test.describe('Backends management page', () => {
|
||||
test.beforeEach(async ({ page }) => {
|
||||
await page.goto('/app/backends')
|
||||
@@ -52,14 +49,11 @@ test.describe('Backends management page - Markdown descriptions', () => {
|
||||
})
|
||||
})
|
||||
await page.goto('/app/backends')
|
||||
// Rendered means the rail has entries. The old gate waited on a column
|
||||
// header, and there are no columns now.
|
||||
await expect(railItem(page, 'markdown-backend')).toBeVisible({ timeout: 10_000 })
|
||||
await expect(page.locator('th', { hasText: 'Description' })).toBeVisible({ timeout: 10_000 })
|
||||
})
|
||||
|
||||
test('the pane lede shows the description as clean text, not raw Markdown', async ({ page }) => {
|
||||
await railItem(page, 'markdown-backend').click()
|
||||
const cell = page.locator('.detail-pane__lede')
|
||||
test('table cell shows the description as clean text, not raw Markdown', async ({ page }) => {
|
||||
const cell = page.locator('tr', { hasText: 'markdown-backend' }).locator('span[title]', { hasText: 'InsightFace' })
|
||||
|
||||
await expect(cell).toHaveText(STRIPPED_DESCRIPTION)
|
||||
// The syntax itself must be gone, not merely rendered somewhere.
|
||||
@@ -71,77 +65,15 @@ test.describe('Backends management page - Markdown descriptions', () => {
|
||||
await expect(cell.locator('h1')).toHaveCount(0)
|
||||
})
|
||||
|
||||
test("the lede's tooltip carries the stripped text, not raw Markdown", async ({ page }) => {
|
||||
await railItem(page, 'markdown-backend').click()
|
||||
await expect(page.locator('.detail-pane__lede')).toHaveAttribute('title', STRIPPED_DESCRIPTION)
|
||||
test('title tooltip carries the stripped text, not raw Markdown', async ({ page }) => {
|
||||
const cell = page.locator('tr', { hasText: 'markdown-backend' }).locator('span[title]', { hasText: 'InsightFace' })
|
||||
|
||||
await expect(cell).toHaveAttribute('title', STRIPPED_DESCRIPTION)
|
||||
})
|
||||
|
||||
test('a backend with no description renders no lede rather than a blank one', async ({ page }) => {
|
||||
// The table needed a placeholder because an empty cell in a grid of full
|
||||
// ones reads as a fault. The pane has no grid to keep aligned, so it omits
|
||||
// the line - but must never print "undefined".
|
||||
await railItem(page, 'plain-backend').click()
|
||||
await expect(page.locator(PANE)).toContainText('plain-backend')
|
||||
await expect(page.locator('.detail-pane__lede')).toHaveCount(0)
|
||||
await expect(page.locator(PANE)).not.toContainText('undefined')
|
||||
})
|
||||
})
|
||||
|
||||
test.describe('Backends gallery - split view', () => {
|
||||
test.beforeEach(async ({ page }) => {
|
||||
await page.route('**/api/backends*', (route) => {
|
||||
route.fulfill({
|
||||
contentType: 'application/json',
|
||||
body: JSON.stringify({
|
||||
backends: [
|
||||
{ name: 'llama-cpp', description: 'GGUF inference', installed: true, version: '1.52.0', license: 'MIT', tags: ['chat'] },
|
||||
{ name: 'whisper', description: 'Speech to text', installed: true, version: '1.8.2', license: 'MIT', tags: ['transcript'] },
|
||||
{ name: 'diffusers', description: 'Image generation', installed: false, license: 'Apache-2.0', tags: ['image'] },
|
||||
],
|
||||
}),
|
||||
})
|
||||
})
|
||||
await page.goto('/app/backends')
|
||||
await expect(railItem(page, 'llama-cpp')).toBeVisible({ timeout: 10_000 })
|
||||
})
|
||||
|
||||
test('the gallery renders no table', async ({ page }) => {
|
||||
await expect(page.locator('[data-testid="backends"]')).toBeVisible()
|
||||
await expect(page.locator('table thead th')).toHaveCount(0)
|
||||
})
|
||||
|
||||
test('with nothing selected the pane describes the host', async ({ page }) => {
|
||||
await expect(page.locator(PANE)).toContainText('This host')
|
||||
await expect(page.locator('[data-testid="backends-back"]')).toHaveCount(0)
|
||||
})
|
||||
|
||||
test('choosing a backend turns the pane into its detail, and back returns', async ({ page }) => {
|
||||
await railItem(page, 'llama-cpp').click()
|
||||
await expect(page.locator(PANE)).toContainText('llama-cpp')
|
||||
await expect(page.locator(PANE)).toContainText('MIT')
|
||||
await expect(page.locator(PANE)).not.toContainText('This host')
|
||||
|
||||
await page.locator('[data-testid="backends-back"]').click()
|
||||
await expect(page.locator(PANE)).toContainText('This host')
|
||||
})
|
||||
|
||||
test('the selection lives in the URL and survives a reload', async ({ page }) => {
|
||||
await railItem(page, 'whisper').click()
|
||||
await expect(page).toHaveURL(/[?&]backend=whisper/)
|
||||
await page.reload()
|
||||
await expect(railItem(page, 'whisper')).toBeVisible({ timeout: 10_000 })
|
||||
await expect(page.locator('[data-testid="backends-back"]')).toBeVisible()
|
||||
})
|
||||
|
||||
|
||||
test('the rail groups while browsing and flattens on a query', async ({ page }) => {
|
||||
await expect(page.locator('[data-testid^="backends-rail-group-"]').first()).toBeVisible()
|
||||
await page.locator('input[placeholder*="Search backends"]').fill('llama')
|
||||
await expect(page.locator('[data-testid^="backends-rail-group-"]')).toHaveCount(0)
|
||||
})
|
||||
|
||||
test('an installed backend states its version, an absent one says so', async ({ page }) => {
|
||||
await expect(railItem(page, 'llama-cpp')).toContainText('v1.52.0')
|
||||
await expect(railItem(page, 'diffusers')).toContainText('not installed')
|
||||
test('a backend with no description still shows the placeholder', async ({ page }) => {
|
||||
const row = page.locator('tr', { hasText: 'plain-backend' })
|
||||
|
||||
await expect(row.locator('span[title=""]')).toHaveText('-')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,23 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// A notice is a hairline with a coloured left edge, not a filled panel. A tint
|
||||
// makes every notice shout at the weight of an error, which is how notices stop
|
||||
// being read — and it is the same treatment the Operate overview uses for the
|
||||
// rows that want a decision.
|
||||
|
||||
test('the backends notice is an edge, not a filled card', async ({ page }) => {
|
||||
// The upgrade banner is the notice worth pinning, so make one exist.
|
||||
await page.route('**/api/backends/upgrades', route => route.fulfill({
|
||||
json: { 'llama-cpp': { backend_name: 'llama-cpp', installed_version: '0.9.4', available_version: '0.9.7' } },
|
||||
}))
|
||||
await page.goto('/app/backends')
|
||||
const notice = page.locator('.bk-notice', { hasText: /update/i }).first()
|
||||
await expect(notice).toBeVisible()
|
||||
const s = await notice.evaluate(el => {
|
||||
const cs = getComputedStyle(el)
|
||||
return { bg: cs.backgroundColor, left: parseFloat(cs.borderLeftWidth), top: parseFloat(cs.borderTopWidth) }
|
||||
})
|
||||
expect(s.bg).toMatch(/rgba\(0, 0, 0, 0\)|transparent/)
|
||||
expect(s.left).toBeGreaterThanOrEqual(3)
|
||||
expect(s.top).toBeLessThanOrEqual(1)
|
||||
})
|
||||
@@ -1,55 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// Chat reads as a transcript rather than a bubble thread (mock 04).
|
||||
|
||||
const CHAT = {
|
||||
chats: [{
|
||||
id: 'c1', name: 'Transcript', model: 'mock-model',
|
||||
history: [
|
||||
{ role: 'user', content: 'Which backends do I have?' },
|
||||
{ role: 'assistant', content: 'Seven are installed.' },
|
||||
],
|
||||
}],
|
||||
activeChatId: 'c1',
|
||||
}
|
||||
|
||||
test.describe('Chat transcript', () => {
|
||||
test.beforeEach(async ({ page }) => {
|
||||
await page.addInitScript(chat => {
|
||||
localStorage.setItem('localai_chats_data', JSON.stringify(chat))
|
||||
}, CHAT)
|
||||
await page.goto('/app/chat')
|
||||
})
|
||||
|
||||
test('neither role is a filled, rounded bubble', async ({ page }) => {
|
||||
const user = page.locator('.chat-message-user .chat-message-content').first()
|
||||
await expect(user).toBeVisible()
|
||||
const cs = await user.evaluate(el => {
|
||||
const s = getComputedStyle(el)
|
||||
return { radius: s.borderTopLeftRadius, shadow: s.boxShadow }
|
||||
})
|
||||
// A rounded filled bubble carries the speaker in shape and side; a
|
||||
// transcript carries it in words, which survives being read aloud.
|
||||
expect(cs.radius).toBe('0px')
|
||||
expect(cs.shadow).toBe('none')
|
||||
})
|
||||
|
||||
test('both turns run full width in one column, not left and right', async ({ page }) => {
|
||||
const user = page.locator('.chat-message-user').first()
|
||||
const assistant = page.locator('.chat-message-assistant').first()
|
||||
const [u, a] = [await user.boundingBox(), await assistant.boundingBox()]
|
||||
expect(Math.abs(u.x - a.x)).toBeLessThan(2)
|
||||
})
|
||||
|
||||
test('every turn says who is speaking', async ({ page }) => {
|
||||
await expect(page.locator('.chat-message-user .chat-message-model')).toHaveText('You')
|
||||
await expect(page.locator('.chat-message-assistant .chat-message-model').first())
|
||||
.toHaveText('mock-model')
|
||||
})
|
||||
|
||||
test('turns are separated by a rule', async ({ page }) => {
|
||||
const border = await page.locator('.chat-message').first()
|
||||
.evaluate(el => getComputedStyle(el).borderBottomStyle)
|
||||
expect(border).toBe('solid')
|
||||
})
|
||||
})
|
||||
@@ -1,50 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// A standing guard against the two defects an earlier automated edit left
|
||||
// scattered through the pages: icons stripped of their fa-* class (which render
|
||||
// nothing at all), and controls left with the user agent's own chrome, which is
|
||||
// a pale grey button on a dark ground.
|
||||
const ROUTES = [
|
||||
'/app', '/app/chat', '/app/models', '/app/studio', '/app/talk',
|
||||
'/app/agents', '/app/skills', '/app/collections', '/app/agent-jobs',
|
||||
'/app/fine-tune', '/app/quantize', '/app/face', '/app/voice',
|
||||
'/app/manage', '/app/backends', '/app/activity', '/app/operate',
|
||||
'/app/settings', '/app/traces', '/app/usage', '/app/nodes', '/app/p2p',
|
||||
'/app/voice-library', '/app/voice-library/new', '/app/account',
|
||||
]
|
||||
|
||||
test('no page renders a dead icon or a default-chrome control', async ({ page }) => {
|
||||
// One test walks every route, so its budget has to scale with the list rather
|
||||
// than sit on Playwright's per-test default of 30s. At 25 routes that default
|
||||
// allows ~1.2s per navigation, which holds on a developer machine and does
|
||||
// not on a loaded CI runner: the suite went red on the commit that added this
|
||||
// spec, timing out mid-loop at waitForTimeout rather than at any single goto,
|
||||
// which is what cumulative slowness looks like as opposed to one hung route.
|
||||
// Six seconds a route absorbs a slow runner and still fails promptly if a
|
||||
// route really does hang.
|
||||
test.setTimeout(ROUTES.length * 6_000)
|
||||
|
||||
const findings = []
|
||||
for (const route of ROUTES) {
|
||||
await page.goto(route)
|
||||
await page.waitForTimeout(400)
|
||||
const found = await page.evaluate(() => {
|
||||
const out = []
|
||||
for (const el of document.querySelectorAll('button, a')) {
|
||||
if (el.getBoundingClientRect().width === 0) continue
|
||||
const cs = getComputedStyle(el)
|
||||
if (cs.borderTopStyle === 'outset' || cs.backgroundColor === 'rgb(239, 239, 239)') {
|
||||
out.push(`default-chrome: "${(el.textContent || '').trim().slice(0, 24)}" [${el.className}]`)
|
||||
}
|
||||
}
|
||||
for (const i of document.querySelectorAll('i')) {
|
||||
if (!/\bfa-/.test((i.className || '').toString())) {
|
||||
out.push(`dead-icon: [${i.className}]`)
|
||||
}
|
||||
}
|
||||
return [...new Set(out)]
|
||||
})
|
||||
for (const f of found) findings.push(`${route} — ${f}`)
|
||||
}
|
||||
expect(findings).toEqual([])
|
||||
})
|
||||
@@ -1,101 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// Small-screen behaviour of the Operate console and the dashboard stat cards.
|
||||
//
|
||||
// Both defects here are about a narrow viewport but neither is only a narrow
|
||||
// viewport problem: the stat cards were being laid out by the wrong rule at
|
||||
// every width, and the rail's height was never bounded.
|
||||
|
||||
test.describe('Operate console on a narrow screen', () => {
|
||||
test('expanding the rail leaves the page still on screen', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 390, height: 800 })
|
||||
await page.goto('/app/manage')
|
||||
|
||||
const toggle = page.locator('.console-rail-toggle')
|
||||
await expect(toggle).toBeVisible()
|
||||
await toggle.click()
|
||||
await expect(page.locator('.console-rail-groups')).toBeVisible()
|
||||
|
||||
// Thirteen destinations in one column is taller than a phone. If opening
|
||||
// the menu pushes the page's own heading past the fold, the menu has
|
||||
// replaced the page instead of annotating it.
|
||||
// Manage titles itself with .view-bar__title rather than .page-title.
|
||||
const heading = page.locator('.page-title, .view-bar__title').first()
|
||||
const box = await heading.boundingBox()
|
||||
expect(box).not.toBeNull()
|
||||
expect(box.y).toBeLessThan(800)
|
||||
})
|
||||
|
||||
test('the rail scrolls internally rather than growing without bound', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 390, height: 800 })
|
||||
await page.goto('/app/manage')
|
||||
await page.locator('.console-rail-toggle').click()
|
||||
|
||||
const groups = page.locator('.console-rail-groups')
|
||||
await expect(groups).toBeVisible()
|
||||
const height = await groups.evaluate(el => el.getBoundingClientRect().height)
|
||||
expect(height).toBeLessThan(800)
|
||||
})
|
||||
})
|
||||
|
||||
test.describe('Headline figures', () => {
|
||||
// Host used shadowed StatCards; it now shares the Operate overview's hairline
|
||||
// figure strip, so the guard is that its labels stay legible, not that it
|
||||
// keeps a card gap.
|
||||
for (const width of [768, 1024]) {
|
||||
test(`Host figure labels are not clipped at ${width}px`, async ({ page }) => {
|
||||
await page.setViewportSize({ width, height: 1000 })
|
||||
await page.goto('/app/manage')
|
||||
const labels = page.locator('.stat-strip__label')
|
||||
await expect(labels.first()).toBeVisible()
|
||||
const clipped = await labels.evaluateAll(els =>
|
||||
els.filter(el => el.scrollWidth > el.clientWidth + 1).map(el => el.textContent))
|
||||
expect(clipped).toEqual([])
|
||||
})
|
||||
}
|
||||
|
||||
test('a Host figure routes into the thing it counts', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1280, height: 1000 })
|
||||
await page.goto('/app/manage')
|
||||
const cell = page.locator('.stat-strip__cell').first()
|
||||
await expect(cell).toBeVisible()
|
||||
// A count is worth more when it is also the way to what it counted.
|
||||
await expect(cell).toHaveJSProperty('tagName', 'BUTTON')
|
||||
})
|
||||
|
||||
test('the figure strip keeps its height inside the flex column', async ({ page }) => {
|
||||
// .page--app is a flex column whose split view takes flex:1, so a child
|
||||
// with no intrinsic minimum gets shrunk to nothing. This strip did exactly
|
||||
// that and rendered 2px tall with four invisible cells.
|
||||
await page.setViewportSize({ width: 1440, height: 900 })
|
||||
await page.goto('/app/manage')
|
||||
const strip = page.locator('.manage-summary')
|
||||
await expect(strip).toBeVisible()
|
||||
const h = await strip.evaluate(el => el.getBoundingClientRect().height)
|
||||
expect(h).toBeGreaterThan(40)
|
||||
})
|
||||
})
|
||||
|
||||
test.describe('Headline figure contrast', () => {
|
||||
test('every figure is legible against the cell it sits on', async ({ page }) => {
|
||||
// A <button> does not inherit colour, so a value with no tone rule fell
|
||||
// back to the UA's `buttontext` — pure black on the dark ground, invisible.
|
||||
await page.setViewportSize({ width: 1440, height: 950 })
|
||||
await page.goto('/app/manage')
|
||||
const bad = await page.locator('.stat-strip__value').evaluateAll(els => els
|
||||
.map(el => ({ text: el.textContent, color: getComputedStyle(el).color }))
|
||||
.filter(v => v.color === 'rgb(0, 0, 0)'))
|
||||
expect(bad).toEqual([])
|
||||
})
|
||||
|
||||
test('the strip keeps its top margin against the shared shorthand', async ({ page }) => {
|
||||
// `.stat-strip` declares `margin: 0 0 ...` later in the file, which was
|
||||
// silently resetting this element's top margin and leaving it flush
|
||||
// against the resources panel above it.
|
||||
await page.setViewportSize({ width: 1440, height: 950 })
|
||||
await page.goto('/app/manage')
|
||||
const top = await page.locator('.manage-summary')
|
||||
.evaluate(el => parseFloat(getComputedStyle(el).marginTop))
|
||||
expect(top).toBeGreaterThan(12)
|
||||
})
|
||||
})
|
||||
@@ -1,72 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// The split view is meant to scroll inside itself. It is easy to regress into
|
||||
// scrolling the document instead, because the shell's height rules are floors
|
||||
// (min-height: 100dvh) rather than ceilings, so any tall pane silently grows
|
||||
// the whole column and takes the rail with it.
|
||||
// A description long enough that the detail pane must overflow, which is the
|
||||
// only condition under which the bug shows.
|
||||
const LONG = Array.from({ length: 60 }, (_, i) =>
|
||||
`Paragraph ${i + 1}. This entry carries a long description so the detail pane has more content than the viewport can hold.`,
|
||||
).join('\n\n')
|
||||
|
||||
const MOCK = {
|
||||
models: [
|
||||
{ name: 'long-model', description: LONG, backend: 'llama-cpp', installed: false, tags: ['llm'] },
|
||||
{ name: 'short-model', description: 'Short.', backend: 'llama-cpp', installed: false, tags: ['llm'] },
|
||||
],
|
||||
allBackends: ['llama-cpp'], allTags: ['llm'],
|
||||
availableModels: 2, installedModels: 0, totalPages: 1, currentPage: 1,
|
||||
}
|
||||
|
||||
test.describe('Discover - the view scrolls, not the page', () => {
|
||||
test.beforeEach(async ({ page }) => {
|
||||
await page.route('**/api/models*', (route) =>
|
||||
route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK) }))
|
||||
})
|
||||
|
||||
test('a long detail scrolls the pane and leaves the page height alone', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1400, height: 900 })
|
||||
await page.goto('/app/models')
|
||||
await expect(page.locator('[data-testid="discover-rail-item"]').first()).toBeVisible({ timeout: 10_000 })
|
||||
|
||||
const pageHeight = () => page.evaluate(() => document.documentElement.scrollHeight)
|
||||
const railHeight = () => page.evaluate(
|
||||
() => document.querySelector('.entity-rail')?.getBoundingClientRect().height,
|
||||
)
|
||||
|
||||
const beforePage = await pageHeight()
|
||||
const beforeRail = await railHeight()
|
||||
|
||||
await page.locator('[data-testid="discover-rail-item"]').first().click()
|
||||
await expect(page.locator('[data-testid="discover-back"]')).toBeVisible()
|
||||
|
||||
// Selecting something must not make the document taller, and must not
|
||||
// stretch the rail to match the pane.
|
||||
expect(await pageHeight()).toBe(beforePage)
|
||||
// Sub-pixel: layout can settle a fraction differently without the rail
|
||||
// having grown. A pixel of tolerance keeps this about the bug it guards.
|
||||
expect(Math.abs((await railHeight()) - beforeRail)).toBeLessThan(1)
|
||||
|
||||
// The pane is the thing that scrolls.
|
||||
const paneOverflows = await page.evaluate(() => {
|
||||
const el = document.querySelector('.split-view__pane')
|
||||
return el ? getComputedStyle(el).overflowY : null
|
||||
})
|
||||
expect(paneOverflows).toBe('auto')
|
||||
})
|
||||
|
||||
test('stacked below the breakpoint it scrolls with the document again', async ({ page }) => {
|
||||
// Pinning the height when the columns stack would trap both halves in short
|
||||
// scrollers, so the constraint is lifted there on purpose.
|
||||
await page.setViewportSize({ width: 700, height: 800 })
|
||||
await page.goto('/app/models')
|
||||
await expect(page.locator('[data-testid="discover-rail-item"]').first()).toBeVisible({ timeout: 10_000 })
|
||||
|
||||
const overflow = await page.evaluate(() => {
|
||||
const el = document.querySelector('.split-view__pane')
|
||||
return el ? getComputedStyle(el).overflowY : null
|
||||
})
|
||||
expect(overflow).toBe('visible')
|
||||
})
|
||||
})
|
||||
@@ -1,52 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// Searching triggers a refetch. The search box lives in the rail column, so if
|
||||
// a refetch unmounts the view it takes the field you are typing into with it,
|
||||
// dropping focus and the caret. That is what this guards.
|
||||
const MOCK = {
|
||||
models: [
|
||||
{ name: 'alpha-model', description: 'a', backend: 'llama-cpp', installed: false, tags: ['llm'] },
|
||||
{ name: 'beta-model', description: 'b', backend: 'llama-cpp', installed: false, tags: ['llm'] },
|
||||
],
|
||||
allBackends: ['llama-cpp'], allTags: ['llm'],
|
||||
availableModels: 2, installedModels: 0, totalPages: 1, currentPage: 1,
|
||||
}
|
||||
|
||||
test.describe('Discover - searching keeps the view', () => {
|
||||
test('a refetch keeps the search box, its focus and its value', async ({ page }) => {
|
||||
let calls = 0
|
||||
await page.route('**/api/models*', async (route) => {
|
||||
calls += 1
|
||||
// Slow the refetch so the loading window is real and observable.
|
||||
if (calls > 1) await new Promise((r) => setTimeout(r, 600))
|
||||
await route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK) })
|
||||
})
|
||||
|
||||
await page.goto('/app/models')
|
||||
const search = page.locator('.filter-bar-group__search input')
|
||||
await expect(search).toBeVisible({ timeout: 10_000 })
|
||||
|
||||
await search.click()
|
||||
await search.fill('alpha')
|
||||
|
||||
// Mid-refetch: the field is still mounted, still focused, still holding
|
||||
// what was typed, and the rail is marked busy rather than replaced.
|
||||
await expect(search).toBeFocused()
|
||||
await expect(search).toHaveValue('alpha')
|
||||
await expect(page.locator('.entity-rail')).toBeVisible()
|
||||
|
||||
await page.waitForTimeout(900)
|
||||
await expect(search).toBeFocused()
|
||||
await expect(search).toHaveValue('alpha')
|
||||
})
|
||||
|
||||
test('the first load still shows a skeleton, not an empty shell', async ({ page }) => {
|
||||
// Nothing to keep on a cold start, so the skeleton is still right there.
|
||||
await page.route('**/api/models*', async (route) => {
|
||||
await new Promise((r) => setTimeout(r, 800))
|
||||
await route.fulfill({ contentType: 'application/json', body: JSON.stringify(MOCK) })
|
||||
})
|
||||
await page.goto('/app/models')
|
||||
await expect(page.getByTestId('gallery-loader')).toBeVisible({ timeout: 5_000 })
|
||||
})
|
||||
})
|
||||
@@ -1,103 +0,0 @@
|
||||
import { test, expect } from './coverage-fixtures.js'
|
||||
|
||||
// Home's resident-model list and the app footer.
|
||||
|
||||
const SYS_INFO = {
|
||||
backends: ['llama-cpp'],
|
||||
loaded_models: [
|
||||
{ id: 'qwen3-8b-instruct', backend: 'llama-cpp' },
|
||||
{ id: 'parakeet-tdt-0.6b' },
|
||||
],
|
||||
}
|
||||
|
||||
async function mockLoaded(page) {
|
||||
await page.route('**/system', route => route.fulfill({ json: SYS_INFO }))
|
||||
await page.route('**/v1/models', route =>
|
||||
route.fulfill({ json: { data: [{ id: 'qwen3-8b-instruct' }, { id: 'parakeet-tdt-0.6b' }] } }))
|
||||
}
|
||||
|
||||
test.describe('Home resident models', () => {
|
||||
test('resident models read as lanes, not status chips', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
const lanes = page.locator('.home-loaded .lane')
|
||||
await expect(lanes).toHaveCount(2)
|
||||
// Model ids are identifiers, so they are set in mono like every other
|
||||
// identifier in the app.
|
||||
const family = await lanes.first().locator('.lane__name').evaluate(
|
||||
el => getComputedStyle(el).fontFamily.toLowerCase())
|
||||
expect(family).toMatch(/mono|consol|menlo/)
|
||||
})
|
||||
|
||||
test('each lane keeps its stop control', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
const lane = page.locator('.home-loaded .lane').first()
|
||||
await expect(lane.getByRole('button', { name: /stop/i })).toBeVisible()
|
||||
})
|
||||
|
||||
test('the header reports how many are resident as a figure', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
const stat = page.locator('[data-testid="home-stat-loaded"]')
|
||||
await expect(stat).toBeVisible()
|
||||
await expect(stat).toContainText('2')
|
||||
// Digits that sit in a column need to line up.
|
||||
const numeric = await stat.locator('.home-stat__value').evaluate(
|
||||
el => getComputedStyle(el).fontVariantNumeric)
|
||||
expect(numeric).toContain('tabular-nums')
|
||||
})
|
||||
|
||||
test('a resident model names the engine serving it', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
// Lanes are sorted by id, so target by content rather than position.
|
||||
const qwen = page.locator('.home-loaded .lane', { hasText: 'qwen3-8b-instruct' })
|
||||
await expect(qwen).toContainText('llama-cpp')
|
||||
})
|
||||
|
||||
test('a model without a config shows no engine rather than a guess', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
// parakeet has no backend in the payload; the column stays blank.
|
||||
const parakeet = page.locator('.home-loaded .lane', { hasText: 'parakeet-tdt-0.6b' })
|
||||
await expect(parakeet).not.toContainText('llama-cpp')
|
||||
})
|
||||
|
||||
test('jump-back-in offers the three places worth returning to', async ({ page }) => {
|
||||
await mockLoaded(page)
|
||||
await page.goto('/app')
|
||||
const lanes = page.locator('.lanes--jump .lane')
|
||||
await expect(lanes).toHaveCount(3)
|
||||
await expect(lanes.first()).toContainText('Discover')
|
||||
})
|
||||
|
||||
test('nothing resident still says so', async ({ page }) => {
|
||||
await page.route('**/system', route =>
|
||||
route.fulfill({ json: { backends: ['llama-cpp'], loaded_models: [] } }))
|
||||
await page.route('**/v1/models', route => route.fulfill({ json: { data: [{ id: 'a-model' }] } }))
|
||||
await page.goto('/app')
|
||||
await expect(page.locator('.home-loaded-empty')).toBeVisible()
|
||||
await expect(page.locator('.home-loaded .lane')).toHaveCount(0)
|
||||
})
|
||||
})
|
||||
|
||||
test.describe('App footer', () => {
|
||||
test('is one line, not three stacked rows', async ({ page }) => {
|
||||
await page.goto('/app')
|
||||
const footer = page.locator('.app-footer')
|
||||
await expect(footer).toBeVisible()
|
||||
// Three centred rows of chrome cost more vertical space than the content
|
||||
// they sit under is usually worth.
|
||||
const height = await footer.evaluate(el => el.getBoundingClientRect().height)
|
||||
expect(height).toBeLessThan(56)
|
||||
})
|
||||
|
||||
test('keeps every link it had', async ({ page }) => {
|
||||
await page.goto('/app')
|
||||
const footer = page.locator('.app-footer')
|
||||
for (const name of [/github/i, /documentation/i, /author/i]) {
|
||||
await expect(footer.getByRole('link', { name })).toBeVisible()
|
||||
}
|
||||
})
|
||||
})
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user